110 lines
2.9 KiB
Go
110 lines
2.9 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"log"
|
|
"net/http"
|
|
"os"
|
|
"os/signal"
|
|
"syscall"
|
|
"time"
|
|
|
|
"git.misaka.ren/M1saka/token_thief/config"
|
|
"git.misaka.ren/M1saka/token_thief/db"
|
|
"git.misaka.ren/M1saka/token_thief/logger"
|
|
"git.misaka.ren/M1saka/token_thief/proxy"
|
|
)
|
|
|
|
func main() {
|
|
log.SetFlags(log.LstdFlags | log.Lmicroseconds)
|
|
if err := run(); err != nil {
|
|
log.Printf("[main] fatal: %v", err)
|
|
os.Exit(1)
|
|
}
|
|
}
|
|
|
|
func run() error {
|
|
cfg, err := config.Load()
|
|
if err != nil {
|
|
return fmt.Errorf("config: %w", err)
|
|
}
|
|
if _, err := db.ClickHouseOptions(cfg.ClickHouseURL); err != nil {
|
|
return err
|
|
}
|
|
|
|
filter, err := config.LoadFilter(cfg.FilterFile)
|
|
if err != nil {
|
|
return fmt.Errorf("load filter: %w", err)
|
|
}
|
|
log.Printf("[main] filter mode=%s patterns=%d", filter.Mode, len(filter.Patterns))
|
|
|
|
rootCtx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
|
|
pool := db.NewPool(rootCtx, cfg.ClickHouseURL, cfg.DBReconnectInterval)
|
|
|
|
queue := logger.NewQueue(pool, cfg.LogQueueSize, cfg.LogBatchSize, cfg.LogWorkers, cfg.LogBatchInterval, cfg.LogQueueBytes)
|
|
queue.Start(rootCtx)
|
|
|
|
h := proxy.NewWithOptions(cfg.UpstreamURL, filter, queue, cfg.MaxBodyBytes, proxy.Options{
|
|
UpstreamTimeout: cfg.UpstreamTimeout,
|
|
ResponseTimeout: cfg.UpstreamResponseTimeout,
|
|
SSEIdleTimeout: cfg.UpstreamStreamIdleTimeout,
|
|
UpstreamTLSInsecureSkipVerify: cfg.UpstreamTLSInsecureSkipVerify,
|
|
TrustedProxies: cfg.TrustedProxies,
|
|
})
|
|
|
|
srv := &http.Server{
|
|
Addr: cfg.ListenAddr,
|
|
Handler: h,
|
|
ReadHeaderTimeout: 30 * time.Second,
|
|
ReadTimeout: cfg.ReadTimeout,
|
|
WriteTimeout: cfg.WriteTimeout,
|
|
IdleTimeout: cfg.IdleTimeout,
|
|
}
|
|
|
|
serverErr := make(chan error, 1)
|
|
go func() {
|
|
log.Printf("[main] listening on %s, upstream=%s", cfg.ListenAddr, cfg.UpstreamURL)
|
|
serverErr <- srv.ListenAndServe()
|
|
}()
|
|
|
|
stop := make(chan os.Signal, 1)
|
|
signal.Notify(stop, syscall.SIGINT, syscall.SIGTERM)
|
|
var runErr error
|
|
select {
|
|
case <-stop:
|
|
log.Printf("[main] shutdown signal received")
|
|
case err := <-serverErr:
|
|
if !errors.Is(err, http.ErrServerClosed) {
|
|
runErr = fmt.Errorf("server: %w", err)
|
|
}
|
|
}
|
|
signal.Stop(stop)
|
|
|
|
// HTTP、升级连接和日志排空共享同一个关闭总预算。
|
|
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 30*time.Second)
|
|
defer shutdownCancel()
|
|
if err := srv.Shutdown(shutdownCtx); err != nil {
|
|
log.Printf("[main] http shutdown: %v", err)
|
|
}
|
|
if err := h.Shutdown(shutdownCtx); err != nil {
|
|
log.Printf("[main] upgraded connection shutdown: %v", err)
|
|
}
|
|
queueStopped := true
|
|
if err := queue.Stop(shutdownCtx); err != nil {
|
|
log.Printf("[main] queue shutdown: %v", err)
|
|
queueStopped = false
|
|
}
|
|
cancel()
|
|
if queueStopped {
|
|
pool.Close()
|
|
} else {
|
|
log.Printf("[main] skip db close while queue workers are still exiting")
|
|
}
|
|
log.Printf("[main] bye")
|
|
return runErr
|
|
}
|