Files
chorus/portal/main.go
T

208 lines
6.3 KiB
Go

package main
import (
"context"
"crypto/rand"
"errors"
"fmt"
"log"
"net/http"
"os"
"os/signal"
"syscall"
"time"
"git.ilapage.cn/OPC/chorus/internal/config"
"git.ilapage.cn/OPC/chorus/internal/core/queue"
platformcrypto "git.ilapage.cn/OPC/chorus/internal/platform/crypto"
safehttp "git.ilapage.cn/OPC/chorus/internal/platform/http"
platformstorage "git.ilapage.cn/OPC/chorus/internal/platform/storage"
"git.ilapage.cn/OPC/chorus/portal/auth"
"git.ilapage.cn/OPC/chorus/portal/handler"
"git.ilapage.cn/OPC/chorus/portal/service"
"git.ilapage.cn/OPC/chorus/portal/session"
"git.ilapage.cn/OPC/chorus/portal/worker"
"github.com/gin-gonic/gin"
"gorm.io/driver/mysql"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
func main() {
if err := run(); err != nil {
log.Printf("chorus portal stopped: %v", err)
os.Exit(1)
}
}
func run() error {
cfg, err := config.Load()
if err != nil {
return err
}
if err := cfg.ValidateProduction(); err != nil {
return err
}
if cfg.Environment == config.Production {
gin.SetMode(gin.ReleaseMode)
}
dbLogger := logger.New(log.New(os.Stderr, "chorus db: ", log.LstdFlags), logger.Config{LogLevel: logger.Error, IgnoreRecordNotFoundError: true, ParameterizedQueries: true})
db, err := gorm.Open(mysql.Open(cfg.DBDSN), &gorm.Config{Logger: dbLogger})
if err != nil {
return errors.New("connect portal database")
}
sqlDB, err := db.DB()
if err != nil {
return errors.New("configure portal database")
}
defer sqlDB.Close()
storage, err := platformstorage.NewLocal(platformstorage.Config{Root: cfg.StorageRoot, MaxObjectBytes: cfg.MaxImageBytes, MaxImagePixels: cfg.MaxImagePixels, ThumbnailMaxSide: 256, AllowedImageMIME: map[string]bool{"image/png": true, "image/jpeg": true, "image/webp": true}})
if err != nil {
return err
}
queueRepository, err := queue.NewMySQLRepository(db)
if err != nil {
return err
}
queueController := queue.NewController(queueRepository)
keyMaterial := []byte(cfg.MasterKey)
if len(keyMaterial) == 0 {
keyMaterial = make([]byte, 32)
if _, err := rand.Read(keyMaterial); err != nil {
return errors.New("generate development provider key")
}
}
keyRing, err := platformcrypto.NewKeyRing("primary", map[string][]byte{"primary": keyMaterial})
for index := range keyMaterial {
keyMaterial[index] = 0
}
if err != nil {
return errors.New("configure provider key ring")
}
httpClient, err := safehttp.New(safehttp.Config{Timeout: cfg.ProviderHTTPTimeout, MaxRedirects: 3})
if err != nil {
return err
}
providerFactory, err := worker.NewOpenAIFactory(safehttp.NewProviderClient(httpClient), cfg.ProviderMaxResponseBytes)
if err != nil {
return err
}
catalog, err := worker.NewGORMCatalog(db, keyRing)
if err != nil {
return err
}
runtime, err := worker.NewGORMRuntime(db)
if err != nil {
return err
}
workerStorage, err := worker.NewLocalStorage(storage)
if err != nil {
return err
}
backgroundWorker, err := worker.New(worker.Config{
Owner: fmt.Sprintf("portal-%d", os.Getpid()),
LeaseDuration: cfg.WorkerLeaseDuration,
PollInterval: cfg.WorkerPollInterval,
}, queueController, queueController, catalog, providerFactory, runtime, workerStorage)
if err != nil {
return err
}
runWorker := backgroundWorker.Run
if cfg.TestDisableWorker {
runWorker = func(ctx context.Context) error {
<-ctx.Done()
return nil
}
}
sessions, err := session.New([]byte(cfg.SessionKey), cfg.SessionTTL, cfg.Environment == config.Production)
if err != nil {
return err
}
authService, err := auth.NewService(db, cfg.LoginAttempts, cfg.LoginWindow)
if err != nil {
return err
}
generationService, err := service.New(db, queueRepository, storage, service.Config{MaxPromptBytes: cfg.MaxPromptBytes, MaxImages: cfg.MaxImages, MaxImageBytes: cfg.MaxImageBytes, MaxUploadBytes: cfg.MaxUploadBytes, MaxImagePixels: cfg.MaxImagePixels, HistoryLimit: cfg.HistoryLimit})
if err != nil {
return err
}
router, err := handler.NewRouter(sessions, authService, generationService, cfg.MaxUploadBytes)
if err != nil {
return err
}
server := &http.Server{Addr: cfg.ListenAddress, Handler: router, ReadHeaderTimeout: 5 * time.Second, ReadTimeout: 30 * time.Second, WriteTimeout: 30 * time.Second, IdleTimeout: 60 * time.Second}
shutdown, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stop()
return runComponents(shutdown, func() error {
log.Printf("chorus portal listening on %s", cfg.ListenAddress)
return server.ListenAndServe()
}, func(ctx context.Context) error {
return server.Shutdown(ctx)
}, runWorker)
}
type componentResult struct {
name string
err error
}
func runComponents(ctx context.Context, runHTTP func() error, shutdownHTTP func(context.Context) error, runWorker func(context.Context) error) error {
workCtx, cancel := context.WithCancel(context.Background())
defer cancel()
results := make(chan componentResult, 2)
go func() { results <- componentResult{name: "http", err: runHTTP()} }()
go func() { results <- componentResult{name: "worker", err: runWorker(workCtx)} }()
var first componentResult
hasFirst := false
select {
case <-ctx.Done():
case first = <-results:
hasFirst = true
}
cancel()
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 15*time.Second)
shutdownErr := shutdownHTTP(shutdownCtx)
shutdownCancel()
remaining := 2
if hasFirst {
remaining--
}
var componentErr error
if hasFirst {
componentErr = unexpectedComponentError(first)
}
for remaining > 0 {
result := <-results
remaining--
if err := completedComponentError(result); componentErr == nil && err != nil {
componentErr = err
}
}
if shutdownErr != nil {
return fmt.Errorf("shutdown portal HTTP server: %w", shutdownErr)
}
return componentErr
}
func unexpectedComponentError(result componentResult) error {
if result.name == "http" && errors.Is(result.err, http.ErrServerClosed) {
return errors.New("portal HTTP server stopped unexpectedly")
}
if result.err == nil {
return fmt.Errorf("portal %s stopped unexpectedly", result.name)
}
return fmt.Errorf("portal %s failed: %w", result.name, result.err)
}
func completedComponentError(result componentResult) error {
if result.name == "http" && errors.Is(result.err, http.ErrServerClosed) {
return nil
}
if result.err != nil {
return fmt.Errorf("portal %s failed during shutdown: %w", result.name, result.err)
}
return nil
}