208 lines
6.3 KiB
Go
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
|
|
}
|