From d0e31b198f135adb230e2f70ed366684dca8ff55 Mon Sep 17 00:00:00 2001 From: Sakurasan <26715255+Sakurasan@users.noreply.github.com> Date: Sat, 15 Aug 2026 21:05:02 +0800 Subject: [PATCH] feat(server): API relay gateway backend M0-M4 Gin + GORM + pure-Go SQLite. Users/auth (JWT), API key management with quotas, proxy gateway with weighted channel failover and health checks, usage/billing ledger, cross-protocol conversion (Anthropic Messages / OpenAI Chat Completions / OpenAI Responses), and channel/model admin API. Channels declare native API formats and auto-convert the rest. Co-Authored-By: Claude Sonnet 5 --- server/cmd/server/main.go | 188 +++++++++ server/go.mod | 58 +++ server/go.sum | 124 ++++++ server/internal/admin/channels.go | 329 +++++++++++++++ server/internal/admin/config.go | 57 +++ server/internal/admin/handler.go | 57 +++ server/internal/admin/models.go | 171 ++++++++ server/internal/admin/usage.go | 222 ++++++++++ server/internal/admin/users.go | 140 +++++++ server/internal/apikey/handler.go | 99 +++++ server/internal/apikey/service.go | 160 +++++++ server/internal/billing/service.go | 117 ++++++ server/internal/channel/health.go | 53 +++ server/internal/channel/ops.go | 220 ++++++++++ server/internal/channel/service.go | 156 +++++++ server/internal/config/config.go | 143 +++++++ server/internal/pkg/crypto/crypto.go | 64 +++ server/internal/pkg/httpx/httpx.go | 42 ++ server/internal/pkg/jwt/jwt.go | 66 +++ server/internal/pkg/password/password.go | 62 +++ server/internal/pkg/rand/rand.go | 21 + server/internal/pkg/ratelimit/ratelimit.go | 73 ++++ server/internal/proxy/claude/types.go | 50 +++ server/internal/proxy/convert/canonical.go | 196 +++++++++ server/internal/proxy/convert/chat_claude.go | 269 ++++++++++++ server/internal/proxy/convert/convert.go | 357 ++++++++++++++++ server/internal/proxy/convert/responses.go | 161 +++++++ server/internal/proxy/convert/stream.go | 416 +++++++++++++++++++ server/internal/proxy/convert_wrappers.go | 17 + server/internal/proxy/gateway.go | 414 ++++++++++++++++++ server/internal/proxy/helpers.go | 248 +++++++++++ server/internal/proxy/models.go | 34 ++ server/internal/proxy/openai/types.go | 87 ++++ server/internal/proxy/responses/types.go | 42 ++ server/internal/proxy/stream.go | 201 +++++++++ server/internal/proxy/stream/stream.go | 79 ++++ server/internal/proxy/usage.go | 48 +++ server/internal/recharge/handler.go | 84 ++++ server/internal/store/model.go | 196 +++++++++ server/internal/store/store.go | 44 ++ server/internal/usage/handler.go | 193 +++++++++ server/internal/usage/service.go | 91 ++++ server/internal/user/handler.go | 107 +++++ server/internal/user/middleware.go | 74 ++++ server/internal/user/service.go | 192 +++++++++ 45 files changed, 6222 insertions(+) create mode 100644 server/cmd/server/main.go create mode 100644 server/go.mod create mode 100644 server/go.sum create mode 100644 server/internal/admin/channels.go create mode 100644 server/internal/admin/config.go create mode 100644 server/internal/admin/handler.go create mode 100644 server/internal/admin/models.go create mode 100644 server/internal/admin/usage.go create mode 100644 server/internal/admin/users.go create mode 100644 server/internal/apikey/handler.go create mode 100644 server/internal/apikey/service.go create mode 100644 server/internal/billing/service.go create mode 100644 server/internal/channel/health.go create mode 100644 server/internal/channel/ops.go create mode 100644 server/internal/channel/service.go create mode 100644 server/internal/config/config.go create mode 100644 server/internal/pkg/crypto/crypto.go create mode 100644 server/internal/pkg/httpx/httpx.go create mode 100644 server/internal/pkg/jwt/jwt.go create mode 100644 server/internal/pkg/password/password.go create mode 100644 server/internal/pkg/rand/rand.go create mode 100644 server/internal/pkg/ratelimit/ratelimit.go create mode 100644 server/internal/proxy/claude/types.go create mode 100644 server/internal/proxy/convert/canonical.go create mode 100644 server/internal/proxy/convert/chat_claude.go create mode 100644 server/internal/proxy/convert/convert.go create mode 100644 server/internal/proxy/convert/responses.go create mode 100644 server/internal/proxy/convert/stream.go create mode 100644 server/internal/proxy/convert_wrappers.go create mode 100644 server/internal/proxy/gateway.go create mode 100644 server/internal/proxy/helpers.go create mode 100644 server/internal/proxy/models.go create mode 100644 server/internal/proxy/openai/types.go create mode 100644 server/internal/proxy/responses/types.go create mode 100644 server/internal/proxy/stream.go create mode 100644 server/internal/proxy/stream/stream.go create mode 100644 server/internal/proxy/usage.go create mode 100644 server/internal/recharge/handler.go create mode 100644 server/internal/store/model.go create mode 100644 server/internal/store/store.go create mode 100644 server/internal/usage/handler.go create mode 100644 server/internal/usage/service.go create mode 100644 server/internal/user/handler.go create mode 100644 server/internal/user/middleware.go create mode 100644 server/internal/user/service.go diff --git a/server/cmd/server/main.go b/server/cmd/server/main.go new file mode 100644 index 0000000..e598731 --- /dev/null +++ b/server/cmd/server/main.go @@ -0,0 +1,188 @@ +package main + +import ( + "log" + "net/http" + "os" + "time" + + "github.com/gin-gonic/gin" + "github.com/shopspring/decimal" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/admin" + "openteam/server/internal/apikey" + "openteam/server/internal/billing" + "openteam/server/internal/channel" + "openteam/server/internal/config" + "openteam/server/internal/pkg/password" + "openteam/server/internal/proxy" + "openteam/server/internal/recharge" + "openteam/server/internal/store" + "openteam/server/internal/usage" + "openteam/server/internal/user" +) + +func main() { + cfg := config.Load() + + zapCfg := zap.NewProductionConfig() + if cfg.Debug { + zapCfg.Level = zap.NewAtomicLevelAt(zap.DebugLevel) + } + logger, err := zapCfg.Build() + if err != nil { + log.Fatalf("init logger: %v", err) + } + defer logger.Sync() + logger.Info("starting openteam server", zap.String("env", cfg.Env), zap.String("db", cfg.DB.Driver)) + + db, err := store.Open(cfg.DB.Driver, cfg.DB.DSN, cfg.Debug) + if err != nil { + logger.Fatal("open database", zap.Error(err)) + } + if err := store.Migrate(db); err != nil { + logger.Fatal("migrate database", zap.Error(err)) + } + seedAdmin(db, cfg, logger) + + // Services. + userSvc := user.NewService(db, cfg, logger) + keySvc := apikey.NewService(db, logger) + billSvc := billing.NewService(db, logger) + usageSvc := usage.NewService(db, logger) + channelSvc := channel.NewService(db, cfg, logger) + gw := proxy.NewGateway(db, cfg, logger, channelSvc, billSvc, usageSvc, keySvc) + + // Background tasks. + channelSvc.StartHealthCheck() + go periodicRecoverCooldown(channelSvc, cfg) + + // Router. + gin.SetMode(gin.ReleaseMode) + if cfg.Debug { + gin.SetMode(gin.DebugMode) + } + r := gin.New() + r.Use(gin.Logger(), gin.Recovery(), corsMiddleware(cfg)) + + // Public proxy endpoints (Bearer API key auth handled inside the gateway). + for _, route := range proxy.Routes { + r.POST(route.UpstreamPath, gw.Handle(route)) + } + modelsHandler := proxy.NewModelsHandler(db, gw) + r.GET("/v1/models", modelsHandler.List) + + // Management API. + api := r.Group("/api/v1") + { + auth := api.Group("/auth") + auth.POST("/register", userH(userSvc, cfg, logger).Register) + auth.POST("/login", userH(userSvc, cfg, logger).Login) + auth.POST("/refresh", userH(userSvc, cfg, logger).Refresh) + auth.POST("/logout", userH(userSvc, cfg, logger).Logout) + auth.GET("/me", userSvc.Middleware(cfg.Auth.JWTSecret), userH(userSvc, cfg, logger).Me) + + userGroup := api.Group("", userSvc.Middleware(cfg.Auth.JWTSecret)) + userGroup.GET("/user/profile", userH(userSvc, cfg, logger).Profile) + userGroup.GET("/user/balance", userH(userSvc, cfg, logger).Balance) + + keyH := apikey.NewHandler(keySvc, logger) + userGroup.GET("/keys", keyH.List) + userGroup.POST("/keys", keyH.Create) + userGroup.PATCH("/keys/:id", keyH.Update) + userGroup.DELETE("/keys/:id", keyH.Delete) + + usageH := usage.NewHandler(db, logger) + userGroup.GET("/usage/summary", usageH.Summary) + userGroup.GET("/usage/stats", usageH.Stats) + userGroup.GET("/usage/logs", usageH.Logs) + + rechargeH := recharge.NewHandler(db, logger) + userGroup.POST("/recharges", rechargeH.Create) + userGroup.GET("/recharges", rechargeH.List) + + admin.Register(api.Group("/admin", userSvc.Middleware(cfg.Auth.JWTSecret)), db, channelSvc, billSvc, logger) + } + + srv := &http.Server{ + Addr: ":" + cfg.HTTPPort, + Handler: r, + ReadHeaderTimeout: 10 * time.Second, + } + logger.Info("listening", zap.String("addr", srv.Addr)) + if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed { + logger.Fatal("server error", zap.Error(err)) + } +} + +func userH(svc *user.Service, cfg *config.Config, logger *zap.Logger) *user.Handler { + return user.NewHandler(svc, cfg, logger) +} + +// seedAdmin creates the initial admin account from env vars, if configured. +func seedAdmin(db *gorm.DB, cfg *config.Config, logger *zap.Logger) { + username := os.Getenv("ADMIN_USERNAME") + adminPass := os.Getenv("ADMIN_PASSWORD") + if username == "" || adminPass == "" { + return + } + var count int64 + db.Model(&store.User{}).Where("role = ?", "admin").Count(&count) + if count > 0 { + return + } + hash, err := password.Hash(adminPass) + if err != nil { + logger.Warn("seed admin: hash failed", zap.Error(err)) + return + } + admin := &store.User{ + Username: username, + Email: username + "@admin.local", + PasswordHash: hash, + Role: "admin", + Balance: decimalZero(), + Status: "active", + } + if err := db.Create(admin).Error; err != nil { + logger.Warn("seed admin: create failed", zap.Error(err)) + return + } + logger.Info("seeded admin account", zap.String("username", username)) +} + +func decimalZero() decimal.Decimal { + return decimal.NewFromInt(0) +} + +func periodicRecoverCooldown(ch *channel.Service, cfg *config.Config) { + ticker := time.NewTicker(cfg.HealthCheck.Interval) + defer ticker.Stop() + for range ticker.C { + ch.RecoverCooldown() + } +} + +func corsMiddleware(cfg *config.Config) gin.HandlerFunc { + return func(c *gin.Context) { + origin := c.GetHeader("Origin") + allowed := "*" + for _, o := range cfg.AllowOrigins { + if o == "*" || o == origin { + allowed = o + break + } + } + c.Header("Access-Control-Allow-Origin", allowed) + c.Header("Access-Control-Allow-Credentials", "true") + c.Header("Access-Control-Allow-Methods", "GET, POST, PUT, PATCH, DELETE, OPTIONS") + c.Header("Access-Control-Allow-Headers", "Content-Type, Authorization") + if c.Request.Method == http.MethodOptions { + c.AbortWithStatus(http.StatusNoContent) + return + } + c.Next() + } +} diff --git a/server/go.mod b/server/go.mod new file mode 100644 index 0000000..3293cbe --- /dev/null +++ b/server/go.mod @@ -0,0 +1,58 @@ +module openteam/server + +go 1.26.5 + +require ( + github.com/bytedance/gopkg v0.1.3 // indirect + github.com/bytedance/sonic v1.15.0 // indirect + github.com/bytedance/sonic/loader v0.5.0 // indirect + github.com/cloudwego/base64x v0.1.6 // indirect + github.com/dustin/go-humanize v1.0.1 // indirect + github.com/gabriel-vasile/mimetype v1.4.12 // indirect + github.com/gin-contrib/sse v1.1.0 // indirect + github.com/gin-gonic/gin v1.12.0 // indirect + github.com/glebarez/go-sqlite v1.21.2 // indirect + github.com/glebarez/sqlite v1.11.0 // indirect + github.com/go-playground/locales v0.14.1 // indirect + github.com/go-playground/universal-translator v0.18.1 // indirect + github.com/go-playground/validator/v10 v10.30.1 // indirect + github.com/goccy/go-json v0.10.5 // indirect + github.com/goccy/go-yaml v1.19.2 // indirect + github.com/golang-jwt/jwt/v5 v5.3.1 // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/pgx/v5 v5.10.0 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect + github.com/jinzhu/inflection v1.0.0 // indirect + github.com/jinzhu/now v1.1.5 // indirect + github.com/json-iterator/go v1.1.12 // indirect + github.com/klauspost/cpuid/v2 v2.3.0 // indirect + github.com/leodido/go-urn v1.4.0 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect + github.com/modern-go/reflect2 v1.0.2 // indirect + github.com/pelletier/go-toml/v2 v2.2.4 // indirect + github.com/quic-go/qpack v0.6.0 // indirect + github.com/quic-go/quic-go v0.59.0 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect + github.com/shopspring/decimal v1.4.0 // indirect + github.com/twitchyliquid64/golang-asm v0.15.1 // indirect + github.com/ugorji/go/codec v1.3.1 // indirect + go.mongodb.org/mongo-driver/v2 v2.5.0 // indirect + go.uber.org/multierr v1.10.0 // indirect + go.uber.org/zap v1.28.0 // indirect + golang.org/x/arch v0.22.0 // indirect + golang.org/x/crypto v0.55.0 // indirect + golang.org/x/net v0.57.0 // indirect + golang.org/x/sync v0.22.0 // indirect + golang.org/x/sys v0.47.0 // indirect + golang.org/x/text v0.41.0 // indirect + google.golang.org/protobuf v1.36.10 // indirect + gorm.io/driver/postgres v1.6.2 // indirect + gorm.io/gorm v1.31.2 // indirect + modernc.org/libc v1.22.5 // indirect + modernc.org/mathutil v1.5.0 // indirect + modernc.org/memory v1.5.0 // indirect + modernc.org/sqlite v1.23.1 // indirect +) diff --git a/server/go.sum b/server/go.sum new file mode 100644 index 0000000..e7f8776 --- /dev/null +++ b/server/go.sum @@ -0,0 +1,124 @@ +github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M= +github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM= +github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uSE= +github.com/bytedance/sonic v1.15.0/go.mod h1:tFkWrPz0/CUCLEF4ri4UkHekCIcdnkqXw9VduqpJh0k= +github.com/bytedance/sonic/loader v0.5.0 h1:gXH3KVnatgY7loH5/TkeVyXPfESoqSBSBEiDd5VjlgE= +github.com/bytedance/sonic/loader v0.5.0/go.mod h1:AR4NYCk5DdzZizZ5djGqQ92eEhCCcdf5x77udYiSJRo= +github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M= +github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= +github.com/gabriel-vasile/mimetype v1.4.12 h1:e9hWvmLYvtp846tLHam2o++qitpguFiYCKbn0w9jyqw= +github.com/gabriel-vasile/mimetype v1.4.12/go.mod h1:d+9Oxyo1wTzWdyVUPMmXFvp4F9tea18J8ufA774AB3s= +github.com/gin-contrib/sse v1.1.0 h1:n0w2GMuUpWDVp7qSpvze6fAu9iRxJY4Hmj6AmBOU05w= +github.com/gin-contrib/sse v1.1.0/go.mod h1:hxRZ5gVpWMT7Z0B0gSNYqqsSCNIJMjzvm6fqCz9vjwM= +github.com/gin-gonic/gin v1.12.0 h1:b3YAbrZtnf8N//yjKeU2+MQsh2mY5htkZidOM7O0wG8= +github.com/gin-gonic/gin v1.12.0/go.mod h1:VxccKfsSllpKshkBWgVgRniFFAzFb9csfngsqANjnLc= +github.com/glebarez/go-sqlite v1.21.2 h1:3a6LFC4sKahUunAmynQKLZceZCOzUthkRkEAl9gAXWo= +github.com/glebarez/go-sqlite v1.21.2/go.mod h1:sfxdZyhQjTM2Wry3gVYWaW072Ri1WMdWJi0k6+3382k= +github.com/glebarez/sqlite v1.11.0 h1:wSG0irqzP6VurnMEpFGer5Li19RpIRi2qvQz++w0GMw= +github.com/glebarez/sqlite v1.11.0/go.mod h1:h8/o8j5wiAsqSPoWELDUdJXhjAhsVliSn7bWZjOhrgQ= +github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/oXslEjJA= +github.com/go-playground/locales v0.14.1/go.mod h1:hxrqLVvrK65+Rwrd5Fc6F2O76J/NuW9t0sjnWqG1slY= +github.com/go-playground/universal-translator v0.18.1 h1:Bcnm0ZwsGyWbCzImXv+pAJnYK9S473LQFuzCbDbfSFY= +github.com/go-playground/universal-translator v0.18.1/go.mod h1:xekY+UJKNuX9WP91TpwSH2VMlDf28Uj24BCp08ZFTUY= +github.com/go-playground/validator/v10 v10.30.1 h1:f3zDSN/zOma+w6+1Wswgd9fLkdwy06ntQJp0BBvFG0w= +github.com/go-playground/validator/v10 v10.30.1/go.mod h1:oSuBIQzuJxL//3MelwSLD5hc2Tu889bF0Idm9Dg26cM= +github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4= +github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= +github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM= +github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= +github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= +github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= +github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0= +github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= +github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E= +github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc= +github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= +github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= +github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= +github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= +github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y= +github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0= +github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ= +github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= +github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg= +github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= +github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M= +github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= +github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4= +github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= +github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= +github.com/quic-go/quic-go v0.59.0 h1:OLJkp1Mlm/aS7dpKgTc6cnpynnD2Xg7C1pwL6vy/SAw= +github.com/quic-go/quic-go v0.59.0/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU= +github.com/remyoudompheng/bigfft v0.0.0-20200410134404-eec4a21b6bb0/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +github.com/shopspring/decimal v1.4.0 h1:bxl37RwXBklmTi0C79JfXCEBD1cqqHt0bbgBAGFp81k= +github.com/shopspring/decimal v1.4.0/go.mod h1:gawqmDU56v4yIKSwfBSFip1HdCCXN8/+DMd9qYNcwME= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= +github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= +github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= +github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= +github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= +github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI= +github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08= +github.com/ugorji/go/codec v1.3.1 h1:waO7eEiFDwidsBN6agj1vJQ4AG7lh2yqXyOXqhgQuyY= +github.com/ugorji/go/codec v1.3.1/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4= +go.mongodb.org/mongo-driver/v2 v2.5.0 h1:yXUhImUjjAInNcpTcAlPHiT7bIXhshCTL3jVBkF3xaE= +go.mongodb.org/mongo-driver/v2 v2.5.0/go.mod h1:yOI9kBsufol30iFsl1slpdq1I0eHPzybRWdyYUs8K/0= +go.uber.org/multierr v1.10.0 h1:S0h4aNzvfcFsC3dRF1jLoaov7oRaKqRGC/pUEJ2yvPQ= +go.uber.org/multierr v1.10.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y= +go.uber.org/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo= +go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q= +golang.org/x/arch v0.22.0 h1:c/Zle32i5ttqRXjdLyyHZESLD/bB90DCU1g9l/0YBDI= +golang.org/x/arch v0.22.0/go.mod h1:dNHoOeKiyja7GTvF9NJS1l3Z2yntpQNzgrjh1cU103A= +golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M= +golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis= +golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= +golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= +golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= +golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +google.golang.org/protobuf v1.36.10 h1:AYd7cD/uASjIL6Q9LiTjz8JLcrh/88q5UObnmY3aOOE= +google.golang.org/protobuf v1.36.10/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gorm.io/driver/postgres v1.6.2 h1:BvXQ/cNUg63q5TFNg672DmDcowZSFrNLkkA3Xe6GXq4= +gorm.io/driver/postgres v1.6.2/go.mod h1:0c4fQA44XhOklXDkgtuKqysHCycTa5i9e3EIpDGCwXk= +gorm.io/gorm v1.31.2 h1:3o8FXNo9v9S858gil+3LlZA1LkCOzgb4g5BL64FgaCo= +gorm.io/gorm v1.31.2/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs= +modernc.org/libc v1.22.5 h1:91BNch/e5B0uPbJFgqbxXuOnxBQjlS//icfQEGmvyjE= +modernc.org/libc v1.22.5/go.mod h1:jj+Z7dTNX8fBScMVNRAYZ/jF91K8fdT2hYMThc3YjBY= +modernc.org/mathutil v1.5.0 h1:rV0Ko/6SfM+8G+yKiyI830l3Wuz1zRutdslNoQ0kfiQ= +modernc.org/mathutil v1.5.0/go.mod h1:mZW8CKdRPY1v87qxC/wUdX5O1qDzXMP5TH3wjfpga6E= +modernc.org/memory v1.5.0 h1:N+/8c5rE6EqugZwHii4IFsaJ7MUhoWX07J5tC/iI5Ds= +modernc.org/memory v1.5.0/go.mod h1:PkUhL0Mugw21sHPeskwZW4D6VscE/GQJOnIpCnW6pSU= +modernc.org/sqlite v1.23.1 h1:nrSBg4aRQQwq59JpvGEQ15tNxoO5pX/kUjcRNwSAGQM= +modernc.org/sqlite v1.23.1/go.mod h1:OrDj17Mggn6MhE+iPbBNf7RGKODDE9NFT0f3EwDzJqk= diff --git a/server/internal/admin/channels.go b/server/internal/admin/channels.go new file mode 100644 index 0000000..c89620c --- /dev/null +++ b/server/internal/admin/channels.go @@ -0,0 +1,329 @@ +package admin + +import ( + "encoding/json" + "fmt" + "net/http" + "strconv" + + "github.com/gin-gonic/gin" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/channel" + "openteam/server/internal/pkg/httpx" + "openteam/server/internal/store" +) + +type channelHandler struct { + db *gorm.DB + ch *channel.Service + log *zap.Logger +} + +type channelInput struct { + Name string `json:"name" binding:"required"` + Provider string `json:"provider" binding:"omitempty,oneof=openai anthropic compatible"` + BaseURL string `json:"baseUrl" binding:"required"` + APIKey string `json:"apiKey"` + Weight int `json:"weight"` + Priority int `json:"priority"` + TimeoutMs int `json:"timeoutMs"` + MaxConcurrency int `json:"maxConcurrency"` + Enabled *bool `json:"enabled"` + Formats []string `json:"formats"` +} + +// ListChannels handles GET /api/v1/admin/channels. +func (h *channelHandler) List(c *gin.Context) { + var channels []store.Channel + if err := h.db.Order("id ASC").Find(&channels).Error; err != nil { + h.log.Warn("list channels failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "list channels failed") + return + } + out := make([]gin.H, 0, len(channels)) + for i := range channels { + out = append(out, channelDTO(&channels[i])) + } + httpx.OK(c, out) +} + +// Create handles POST /api/v1/admin/channels. +func (h *channelHandler) Create(c *gin.Context) { + var in channelInput + if !httpx.Bind(c, &in) { + return + } + provider := in.Provider + if len(in.Formats) > 0 { + if err := validateFormats(in.Formats); err != nil { + httpx.Fail(c, http.StatusBadRequest, err.Error()) + return + } + provider = deriveProvider(in.Formats) + } + if provider == "" { + provider = "openai" + } + enc, err := h.ch.EncryptKey(in.APIKey) + if err != nil { + httpx.Fail(c, http.StatusInternalServerError, "encrypt key failed") + return + } + ch := &store.Channel{ + Name: in.Name, Provider: provider, BaseURL: in.BaseURL, Formats: in.Formats, + APIKeyEnc: enc, Weight: defaultIf(in.Weight, 1), Priority: in.Priority, + TimeoutMs: defaultIf(in.TimeoutMs, 300000), MaxConcurrency: defaultIf(in.MaxConcurrency, 100), + Enabled: true, + } + if in.Enabled != nil { + ch.Enabled = *in.Enabled + } + if err := h.db.Create(ch).Error; err != nil { + h.log.Warn("create channel failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "create channel failed") + return + } + httpx.Created(c, channelDTO(ch)) +} + +// Update handles PUT /api/v1/admin/channels/:id. +func (h *channelHandler) Update(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid channel id") + return + } + var ch store.Channel + if err := h.db.First(&ch, id).Error; err != nil { + httpx.Fail(c, http.StatusNotFound, "channel not found") + return + } + var in channelInput + if !httpx.Bind(c, &in) { + return + } + updates := map[string]any{ + "name": in.Name, "base_url": in.BaseURL, + "weight": defaultIf(in.Weight, ch.Weight), "priority": in.Priority, + "timeout_ms": defaultIf(in.TimeoutMs, ch.TimeoutMs), + "max_concurrency": defaultIf(in.MaxConcurrency, ch.MaxConcurrency), + } + if len(in.Formats) > 0 { + if err := validateFormats(in.Formats); err != nil { + httpx.Fail(c, http.StatusBadRequest, err.Error()) + return + } + updates["provider"] = deriveProvider(in.Formats) + // GORM map updates skip the json serializer, so encode explicitly; + // a raw []string would be emitted as a SQL row-value list. + formatsJSON, err := json.Marshal(in.Formats) + if err != nil { + httpx.Fail(c, http.StatusInternalServerError, "serialize formats failed") + return + } + updates["formats"] = string(formatsJSON) + } else if in.Provider != "" { + updates["provider"] = in.Provider + } + if in.Enabled != nil { + updates["enabled"] = *in.Enabled + } + if in.APIKey != "" { + enc, err := h.ch.EncryptKey(in.APIKey) + if err != nil { + httpx.Fail(c, http.StatusInternalServerError, "encrypt key failed") + return + } + updates["api_key_enc"] = enc + } + if err := h.db.Model(&ch).Updates(updates).Error; err != nil { + h.log.Warn("update channel failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "update channel failed") + return + } + h.db.First(&ch, id) + httpx.OK(c, channelDTO(&ch)) +} + +// Delete handles DELETE /api/v1/admin/channels/:id. +func (h *channelHandler) Delete(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid channel id") + return + } + if err := h.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("channel_id = ?", id).Delete(&store.ChannelModelBinding{}).Error; err != nil { + return err + } + return tx.Delete(&store.Channel{}, id).Error + }); err != nil { + h.log.Warn("delete channel failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "delete channel failed") + return + } + httpx.OK(c, gin.H{"ok": true}) +} + +// Test handles POST /api/v1/admin/channels/:id/test. +func (h *channelHandler) Test(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid channel id") + return + } + var ch store.Channel + if err := h.db.First(&ch, id).Error; err != nil { + httpx.Fail(c, http.StatusNotFound, "channel not found") + return + } + ok, latency, model, err := h.ch.Test(&ch) + if err != nil || !ok { + msg := "test failed" + if err != nil { + msg = err.Error() + } + httpx.Fail(c, http.StatusBadGateway, msg) + return + } + httpx.OK(c, gin.H{"ok": true, "latencyMs": latency, "model": model}) +} + +// ImportModels handles POST /api/v1/admin/channels/:id/import-models. +func (h *channelHandler) ImportModels(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid channel id") + return + } + var ch store.Channel + if err := h.db.First(&ch, id).Error; err != nil { + httpx.Fail(c, http.StatusNotFound, "channel not found") + return + } + names, err := h.ch.ImportModels(&ch) + if err != nil { + httpx.Fail(c, http.StatusBadGateway, "import failed: "+err.Error()) + return + } + httpx.OK(c, gin.H{"imported": names, "count": len(names)}) +} + +// ListBindings handles GET /api/v1/admin/channels/:id/bindings. +func (h *channelHandler) ListBindings(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid channel id") + return + } + var bindings []store.ChannelModelBinding + if err := h.db.Preload("Model").Where("channel_id = ?", id).Find(&bindings).Error; err != nil { + h.log.Warn("list bindings failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "list bindings failed") + return + } + out := make([]gin.H, 0, len(bindings)) + for _, b := range bindings { + out = append(out, gin.H{ + "id": b.ID, "modelId": b.ModelID, "modelName": b.Model.Name, + "upstreamModel": b.UpstreamModel, "weight": b.Weight, + }) + } + httpx.OK(c, out) +} + +// SaveBindings handles PUT /api/v1/admin/channels/:id/bindings. +func (h *channelHandler) SaveBindings(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid channel id") + return + } + var in struct { + Items []struct { + ModelID int64 `json:"modelId"` + UpstreamModel string `json:"upstreamModel"` + Weight int `json:"weight"` + } `json:"items"` + } + if !httpx.Bind(c, &in) { + return + } + if err := h.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("channel_id = ?", id).Delete(&store.ChannelModelBinding{}).Error; err != nil { + return err + } + for _, item := range in.Items { + if item.ModelID == 0 { + continue + } + b := store.ChannelModelBinding{ + ChannelID: id, ModelID: item.ModelID, + UpstreamModel: item.UpstreamModel, Weight: defaultIf(item.Weight, 1), + } + if err := tx.Create(&b).Error; err != nil { + return err + } + } + return nil + }); err != nil { + h.log.Warn("save bindings failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "save bindings failed") + return + } + httpx.OK(c, gin.H{"ok": true}) +} + +func channelDTO(ch *store.Channel) gin.H { + return gin.H{ + "id": ch.ID, "name": ch.Name, "provider": ch.Provider, "baseUrl": ch.BaseURL, + "formats": ch.FormatsResolved(), + "weight": ch.Weight, "priority": ch.Priority, "timeoutMs": ch.TimeoutMs, + "maxConcurrency": ch.MaxConcurrency, "healthStatus": ch.HealthStatus, + "healthFailures": ch.HealthFailures, "enabled": ch.Enabled, + "createdAt": ch.CreatedAt, "updatedAt": ch.UpdatedAt, + } +} + +var validFormats = map[string]bool{ + store.FormatOpenAIChat: true, + store.FormatOpenAIResponses: true, + store.FormatAnthropic: true, +} + +func validateFormats(fs []string) error { + for _, f := range fs { + if !validFormats[f] { + return fmt.Errorf("invalid api format: %s", f) + } + } + return nil +} + +// deriveProvider computes the conversion-target provider from the declared +// formats. Anthropic-only channels convert to Claude; anything else targets the +// OpenAI chat-completions family. +func deriveProvider(formats []string) string { + hasAnthropic, hasOpenAI := false, false + for _, f := range formats { + switch f { + case store.FormatAnthropic: + hasAnthropic = true + case store.FormatOpenAIChat, store.FormatOpenAIResponses: + hasOpenAI = true + } + } + if hasAnthropic && !hasOpenAI { + return "anthropic" + } + return "openai" +} + +func defaultIf(v, def int) int { + if v == 0 { + return def + } + return v +} diff --git a/server/internal/admin/config.go b/server/internal/admin/config.go new file mode 100644 index 0000000..6081eae --- /dev/null +++ b/server/internal/admin/config.go @@ -0,0 +1,57 @@ +package admin + +import ( + "encoding/json" + "net/http" + + "github.com/gin-gonic/gin" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/pkg/httpx" + "openteam/server/internal/store" +) + +type configHandler struct { + db *gorm.DB + log *zap.Logger +} + +// Get handles GET /api/v1/admin/config — returns the full config map. +func (h *configHandler) Get(c *gin.Context) { + var confs []store.SystemConfig + if err := h.db.Find(&confs).Error; err != nil { + h.log.Warn("load config failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "load config failed") + return + } + out := map[string]any{} + for _, conf := range confs { + var v any + if err := json.Unmarshal(conf.Value, &v); err == nil { + out[conf.Key] = v + } + } + httpx.OK(c, out) +} + +// Put handles PUT /api/v1/admin/config — upserts key/value pairs. +func (h *configHandler) Put(c *gin.Context) { + var in map[string]any + if !httpx.Bind(c, &in) { + return + } + for k, v := range in { + raw, err := json.Marshal(v) + if err != nil { + continue + } + conf := store.SystemConfig{Key: k, Value: raw} + if err := h.db.Save(&conf).Error; err != nil { + h.log.Warn("save config failed", zap.Error(err), zap.String("key", k)) + httpx.Fail(c, http.StatusInternalServerError, "save config failed") + return + } + } + httpx.OK(c, gin.H{"ok": true}) +} diff --git a/server/internal/admin/handler.go b/server/internal/admin/handler.go new file mode 100644 index 0000000..1ecd324 --- /dev/null +++ b/server/internal/admin/handler.go @@ -0,0 +1,57 @@ +package admin + +import ( + "github.com/gin-gonic/gin" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/billing" + "openteam/server/internal/channel" + "openteam/server/internal/user" +) + +// Register wires the admin routes onto a gin group. +func Register(g *gin.RouterGroup, db *gorm.DB, ch *channel.Service, bill *billing.Service, log *zap.Logger) { + // Handler instances. + chHandler := &channelHandler{db: db, ch: ch, log: log} + modelHandler := &modelHandler{db: db, log: log} + userHandler := &userHandler{db: db, bill: bill, log: log} + usageHandler := &usageHandler{db: db, bill: bill, log: log} + configHandler := &configHandler{db: db, log: log} + + g = g.Group("", user.RequireAdmin()) + + // Channels. + g.GET("/channels", chHandler.List) + g.POST("/channels", chHandler.Create) + g.PUT("/channels/:id", chHandler.Update) + g.DELETE("/channels/:id", chHandler.Delete) + g.POST("/channels/:id/test", chHandler.Test) + g.POST("/channels/:id/import-models", chHandler.ImportModels) + g.GET("/channels/:id/bindings", chHandler.ListBindings) + g.PUT("/channels/:id/bindings", chHandler.SaveBindings) + + // Models. + g.GET("/models", modelHandler.List) + g.POST("/models", modelHandler.Create) + g.PUT("/models/:id", modelHandler.Update) + g.PUT("/models/:id/price", modelHandler.UpdatePrice) + + // Users. + g.GET("/users", userHandler.List) + g.PATCH("/users/:id", userHandler.Update) + g.POST("/users/:id/balance", userHandler.AdjustBalance) + + // Recharges (reserved; approve/reject wire the paused state machine). + g.GET("/recharges", usageHandler.Recharges) + g.POST("/recharges/:id/approve", usageHandler.ApproveRecharge) + g.POST("/recharges/:id/reject", usageHandler.RejectRecharge) + + // Usage & stats. + g.GET("/usage", usageHandler.Usage) + g.GET("/stats/overview", usageHandler.Overview) + + // System config. + g.GET("/config", configHandler.Get) + g.PUT("/config", configHandler.Put) +} diff --git a/server/internal/admin/models.go b/server/internal/admin/models.go new file mode 100644 index 0000000..e79b019 --- /dev/null +++ b/server/internal/admin/models.go @@ -0,0 +1,171 @@ +package admin + +import ( + "net/http" + "strconv" + + "github.com/gin-gonic/gin" + "github.com/shopspring/decimal" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/pkg/httpx" + "openteam/server/internal/store" +) + +type modelHandler struct { + db *gorm.DB + log *zap.Logger +} + +type modelInput struct { + Name string `json:"name" binding:"required"` + DisplayName string `json:"displayName"` + InputPrice string `json:"inputPrice"` + OutputPrice string `json:"outputPrice"` + CacheReadPrice string `json:"cacheReadPrice"` + Enabled *bool `json:"enabled"` + Sort *int `json:"sort"` +} + +// List handles GET /api/v1/admin/models. +func (h *modelHandler) List(c *gin.Context) { + var models []store.Model + if err := h.db.Order("sort ASC, id ASC").Find(&models).Error; err != nil { + h.log.Warn("list models failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "list models failed") + return + } + out := make([]gin.H, 0, len(models)) + for i := range models { + out = append(out, modelDTO(&models[i])) + } + httpx.OK(c, out) +} + +// Create handles POST /api/v1/admin/models. +func (h *modelHandler) Create(c *gin.Context) { + var in modelInput + if !httpx.Bind(c, &in) { + return + } + m := &store.Model{ + Name: in.Name, + DisplayName: in.DisplayName, + InputPrice: parsePrice(in.InputPrice), + OutputPrice: parsePrice(in.OutputPrice), + CacheReadPrice: parsePrice(in.CacheReadPrice), + Enabled: true, + } + if in.Enabled != nil { + m.Enabled = *in.Enabled + } + if in.Sort != nil { + m.Sort = *in.Sort + } + if err := h.db.Create(m).Error; err != nil { + h.log.Warn("create model failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "create model failed") + return + } + httpx.Created(c, modelDTO(m)) +} + +// Update handles PUT /api/v1/admin/models/:id. +func (h *modelHandler) Update(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid model id") + return + } + var m store.Model + if err := h.db.First(&m, id).Error; err != nil { + httpx.Fail(c, http.StatusNotFound, "model not found") + return + } + var in modelInput + if !httpx.Bind(c, &in) { + return + } + updates := map[string]any{ + "name": in.Name, + } + if in.DisplayName != "" { + updates["display_name"] = in.DisplayName + } + if in.InputPrice != "" { + updates["input_price"] = parsePrice(in.InputPrice) + } + if in.OutputPrice != "" { + updates["output_price"] = parsePrice(in.OutputPrice) + } + if in.CacheReadPrice != "" { + updates["cache_read_price"] = parsePrice(in.CacheReadPrice) + } + if in.Enabled != nil { + updates["enabled"] = *in.Enabled + } + if in.Sort != nil { + updates["sort"] = *in.Sort + } + if err := h.db.Model(&m).Updates(updates).Error; err != nil { + h.log.Warn("update model failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "update model failed") + return + } + h.db.First(&m, id) + httpx.OK(c, modelDTO(&m)) +} + +// UpdatePrice handles PUT /api/v1/admin/models/:id/price. +func (h *modelHandler) UpdatePrice(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid model id") + return + } + var in struct { + InputPrice string `json:"inputPrice"` + OutputPrice string `json:"outputPrice"` + CacheReadPrice string `json:"cacheReadPrice"` + } + if !httpx.Bind(c, &in) { + return + } + updates := map[string]any{} + if in.InputPrice != "" { + updates["input_price"] = parsePrice(in.InputPrice) + } + if in.OutputPrice != "" { + updates["output_price"] = parsePrice(in.OutputPrice) + } + if in.CacheReadPrice != "" { + updates["cache_read_price"] = parsePrice(in.CacheReadPrice) + } + if len(updates) == 0 { + httpx.Fail(c, http.StatusBadRequest, "no price fields provided") + return + } + res := h.db.Model(&store.Model{}).Where("id = ?", id).Updates(updates) + if res.RowsAffected == 0 { + httpx.Fail(c, http.StatusNotFound, "model not found") + return + } + httpx.OK(c, gin.H{"ok": true}) +} + +func parsePrice(s string) decimal.Decimal { + d, err := decimal.NewFromString(s) + if err != nil { + return decimal.Zero + } + return d.Round(8) +} + +func modelDTO(m *store.Model) gin.H { + return gin.H{ + "id": m.ID, "name": m.Name, "displayName": m.DisplayName, + "inputPrice": m.InputPrice.String(), "outputPrice": m.OutputPrice.String(), + "cacheReadPrice": m.CacheReadPrice.String(), "enabled": m.Enabled, "sort": m.Sort, + } +} diff --git a/server/internal/admin/usage.go b/server/internal/admin/usage.go new file mode 100644 index 0000000..d91459e --- /dev/null +++ b/server/internal/admin/usage.go @@ -0,0 +1,222 @@ +package admin + +import ( + "net/http" + "strconv" + "time" + + "github.com/gin-gonic/gin" + "github.com/shopspring/decimal" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/billing" + "openteam/server/internal/pkg/httpx" + "openteam/server/internal/store" + "openteam/server/internal/user" +) + +type usageHandler struct { + db *gorm.DB + bill *billing.Service + log *zap.Logger +} + +// Overview handles GET /api/v1/admin/stats/overview. +func (h *usageHandler) Overview(c *gin.Context) { + now := time.Now() + today := now.Format("2006-01-02") + month := now.Format("2006-01-02") + + var totals struct { + Requests int + Cost decimal.Decimal + } + h.db.Model(&store.UsageDaily{}). + Select("COALESCE(SUM(requests),0) as requests, COALESCE(SUM(cost),0) as cost"). + Scan(&totals) + + var todayAgg struct { + Requests int + Cost decimal.Decimal + InputTokens int64 + OutputTokens int64 + } + h.db.Model(&store.UsageDaily{}).Where("date = ?", today). + Select("COALESCE(SUM(requests),0) as requests, COALESCE(SUM(cost),0) as cost, COALESCE(SUM(input_tokens),0) as input_tokens, COALESCE(SUM(output_tokens),0) as output_tokens"). + Scan(&todayAgg) + + var monthAgg struct { + Requests int + Cost decimal.Decimal + } + h.db.Model(&store.UsageDaily{}).Where("date >= ? AND date <= ?", month[:7]+"-01", today). + Select("COALESCE(SUM(requests),0) as requests, COALESCE(SUM(cost),0) as cost"). + Scan(&monthAgg) + + var userCount, channelCount, modelCount int64 + h.db.Model(&store.User{}).Count(&userCount) + h.db.Model(&store.Channel{}).Count(&channelCount) + h.db.Model(&store.Model{}).Count(&modelCount) + + httpx.OK(c, gin.H{ + "total": gin.H{ + "requests": totals.Requests, + "cost": totals.Cost.String(), + "users": userCount, + "channels": channelCount, + "models": modelCount, + }, + "today": gin.H{ + "requests": todayAgg.Requests, + "cost": todayAgg.Cost.String(), + "inputTokens": todayAgg.InputTokens, + "outputTokens": todayAgg.OutputTokens, + }, + "month": gin.H{ + "requests": monthAgg.Requests, + "cost": monthAgg.Cost.String(), + }, + }) +} + +// Usage handles GET /api/v1/admin/usage?from&to&userId&model&group=day|model|user. +func (h *usageHandler) Usage(c *gin.Context) { + from := c.Query("from") + to := c.Query("to") + group := c.DefaultQuery("group", "day") + + q := h.db.Model(&store.UsageDaily{}) + if from != "" { + q = q.Where("date >= ?", from) + } + if to != "" { + q = q.Where("date <= ?", to) + } + if uid := c.Query("userId"); uid != "" { + q = q.Where("user_id = ?", uid) + } + if m := c.Query("model"); m != "" { + q = q.Where("model_id = ?", m) + } + + var rows []struct { + Key string `gorm:"column:g"` + Requests int + InputTokens int64 + OutputTokens int64 + Cost decimal.Decimal + } + switch group { + case "model": + q = q.Joins("JOIN models ON models.id = usage_dailies.model_id"). + Select("models.name as g, COALESCE(SUM(requests),0) as requests, COALESCE(SUM(input_tokens),0) as input_tokens, COALESCE(SUM(output_tokens),0) as output_tokens, COALESCE(SUM(cost),0) as cost"). + Group("models.name") + case "user": + q = q.Joins("JOIN users ON users.id = usage_dailies.user_id"). + Select("users.username as g, COALESCE(SUM(requests),0) as requests, COALESCE(SUM(input_tokens),0) as input_tokens, COALESCE(SUM(output_tokens),0) as output_tokens, COALESCE(SUM(cost),0) as cost"). + Group("users.username") + default: + q = q.Select("date as g, COALESCE(SUM(requests),0) as requests, COALESCE(SUM(input_tokens),0) as input_tokens, COALESCE(SUM(output_tokens),0) as output_tokens, COALESCE(SUM(cost),0) as cost"). + Group("date").Order("date ASC") + } + if err := q.Scan(&rows).Error; err != nil { + h.log.Warn("admin usage query failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "usage query failed") + return + } + out := make([]gin.H, 0, len(rows)) + for _, r := range rows { + out = append(out, gin.H{ + "key": r.Key, "requests": r.Requests, + "inputTokens": r.InputTokens, "outputTokens": r.OutputTokens, + "cost": r.Cost.String(), + }) + } + httpx.OK(c, out) +} + +// Recharges handles GET /api/v1/admin/recharges. +func (h *usageHandler) Recharges(c *gin.Context) { + status := c.Query("status") + q := h.db.Preload("User").Model(&store.RechargeOrder{}) + if status != "" { + q = q.Where("status = ?", status) + } + var orders []store.RechargeOrder + if err := q.Order("id DESC").Limit(100).Find(&orders).Error; err != nil { + h.log.Warn("list recharges failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "list recharges failed") + return + } + out := make([]gin.H, 0, len(orders)) + for i := range orders { + o := &orders[i] + out = append(out, gin.H{ + "id": o.ID, "userId": o.UserID, "username": o.User.Username, + "amount": o.Amount.String(), "status": o.Status, "method": o.Method, + "remark": o.Remark, "createdAt": o.CreatedAt, + }) + } + httpx.OK(c, out) +} + +// ApproveRecharge handles POST /api/v1/admin/recharges/:id/approve. +func (h *usageHandler) ApproveRecharge(c *gin.Context) { + h.decideRecharge(c, "approve") +} + +// RejectRecharge handles POST /api/v1/admin/recharges/:id/reject. +func (h *usageHandler) RejectRecharge(c *gin.Context) { + h.decideRecharge(c, "reject") +} + +func (h *usageHandler) decideRecharge(c *gin.Context, action string) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid order id") + return + } + admin := user.Current(c) + + var order store.RechargeOrder + if err := h.db.First(&order, id).Error; err != nil { + httpx.Fail(c, http.StatusNotFound, "order not found") + return + } + if order.Status != "pending" { + httpx.Fail(c, http.StatusBadRequest, "order already processed") + return + } + + var in struct { + Remark string `json:"remark"` + } + _ = c.ShouldBindJSON(&in) + + now := time.Now() + if action == "approve" { + err = h.db.Transaction(func(tx *gorm.DB) error { + if _, cerr := h.bill.Credit(order.UserID, order.Amount, "recharge", orderIDStr(order.ID)); cerr != nil { + return cerr + } + return tx.Model(&order).Updates(map[string]any{ + "status": "credited", "reviewed_by": admin.ID, "reviewed_at": now, "remark": in.Remark, + }).Error + }) + } else { + err = h.db.Model(&order).Updates(map[string]any{ + "status": "rejected", "reviewed_by": admin.ID, "reviewed_at": now, "remark": in.Remark, + }).Error + } + if err != nil { + h.log.Warn("recharge decision failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "recharge decision failed") + return + } + httpx.OK(c, gin.H{"ok": true, "status": order.Status}) +} + +func orderIDStr(id int64) string { + return strconv.FormatInt(id, 10) +} diff --git a/server/internal/admin/users.go b/server/internal/admin/users.go new file mode 100644 index 0000000..2a449c7 --- /dev/null +++ b/server/internal/admin/users.go @@ -0,0 +1,140 @@ +package admin + +import ( + "net/http" + "strconv" + + "github.com/gin-gonic/gin" + "github.com/shopspring/decimal" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/billing" + "openteam/server/internal/pkg/httpx" + "openteam/server/internal/store" +) + +type userHandler struct { + db *gorm.DB + bill *billing.Service + log *zap.Logger +} + +// List handles GET /api/v1/admin/users?page&search. +func (h *userHandler) List(c *gin.Context) { + page, _ := strconv.Atoi(c.DefaultQuery("page", "1")) + if page < 1 { + page = 1 + } + pageSize := 20 + search := c.Query("search") + + q := h.db.Model(&store.User{}) + if search != "" { + like := "%" + search + "%" + q = q.Where("username LIKE ? OR email LIKE ?", like, like) + } + var total int64 + q.Count(&total) + + var users []store.User + if err := q.Order("id DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&users).Error; err != nil { + h.log.Warn("list users failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "list users failed") + return + } + out := make([]gin.H, 0, len(users)) + for i := range users { + out = append(out, userDTO(&users[i])) + } + httpx.OK(c, gin.H{"total": total, "page": page, "items": out}) +} + +// Update handles PATCH /api/v1/admin/users/:id. +func (h *userHandler) Update(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid user id") + return + } + var u store.User + if err := h.db.First(&u, id).Error; err != nil { + httpx.Fail(c, http.StatusNotFound, "user not found") + return + } + var in struct { + Role *string `json:"role"` + Status *string `json:"status"` + } + if !httpx.Bind(c, &in) { + return + } + updates := map[string]any{} + if in.Role != nil { + if *in.Role != "admin" && *in.Role != "user" { + httpx.Fail(c, http.StatusBadRequest, "invalid role") + return + } + updates["role"] = *in.Role + } + if in.Status != nil { + if *in.Status != "active" && *in.Status != "disabled" { + httpx.Fail(c, http.StatusBadRequest, "invalid status") + return + } + updates["status"] = *in.Status + } + if len(updates) > 0 { + if err := h.db.Model(&u).Updates(updates).Error; err != nil { + h.log.Warn("update user failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "update user failed") + return + } + } + h.db.First(&u, id) + httpx.OK(c, userDTO(&u)) +} + +// AdjustBalance handles POST /api/v1/admin/users/:id/balance. +func (h *userHandler) AdjustBalance(c *gin.Context) { + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid user id") + return + } + var in struct { + Amount string `json:"amount" binding:"required"` + Remark string `json:"remark"` + } + if !httpx.Bind(c, &in) { + return + } + amount, err := decimal.NewFromString(in.Amount) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid amount") + return + } + if amount.IsZero() { + httpx.Fail(c, http.StatusBadRequest, "amount must be non-zero") + return + } + after, err := h.bill.AdminAdjust(id, amount.Round(8), in.Remark) + if err != nil { + if err == billing.ErrInsufficientBalance { + httpx.Fail(c, http.StatusBadRequest, "amount would make balance negative") + return + } + h.log.Warn("adjust balance failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "adjust balance failed") + return + } + httpx.OK(c, gin.H{"balanceAfter": after.String()}) +} + +func userDTO(u *store.User) gin.H { + return gin.H{ + "id": u.ID, "username": u.Username, "email": u.Email, + "role": u.Role, "balance": u.Balance.String(), "status": u.Status, + "lastLoginAt": u.LastLoginAt, "createdAt": u.CreatedAt, + } +} diff --git a/server/internal/apikey/handler.go b/server/internal/apikey/handler.go new file mode 100644 index 0000000..5f24a87 --- /dev/null +++ b/server/internal/apikey/handler.go @@ -0,0 +1,99 @@ +package apikey + +import ( + "errors" + "net/http" + "strconv" + + "github.com/gin-gonic/gin" + "go.uber.org/zap" + + "openteam/server/internal/pkg/httpx" + "openteam/server/internal/user" +) + +type Handler struct { + svc *Service + log *zap.Logger +} + +func NewHandler(svc *Service, log *zap.Logger) *Handler { + return &Handler{svc: svc, log: log} +} + +// List handles GET /api/v1/keys. +func (h *Handler) List(c *gin.Context) { + u := user.Current(c) + keys, err := h.svc.List(u.ID) + if err != nil { + h.log.Warn("list keys failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "list keys failed") + return + } + out := make([]map[string]any, 0, len(keys)) + for i := range keys { + out = append(out, Public(&keys[i])) + } + httpx.OK(c, out) +} + +// Create handles POST /api/v1/keys. +func (h *Handler) Create(c *gin.Context) { + u := user.Current(c) + var in CreateInput + if !httpx.Bind(c, &in) { + return + } + rec, plain, err := h.svc.Create(u.ID, in) + if err != nil { + h.log.Warn("create key failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "create key failed") + return + } + data := Public(rec) + data["key"] = plain // shown only once + httpx.Created(c, data) +} + +// Update handles PATCH /api/v1/keys/:id. +func (h *Handler) Update(c *gin.Context) { + u := user.Current(c) + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid key id") + return + } + var in UpdateInput + if !httpx.Bind(c, &in) { + return + } + rec, err := h.svc.Update(u.ID, id, in) + if err != nil { + if errors.Is(err, ErrKeyNotFound) { + httpx.Fail(c, http.StatusNotFound, "api key not found") + return + } + httpx.Fail(c, http.StatusInternalServerError, "update key failed") + return + } + httpx.OK(c, Public(rec)) +} + +// Delete handles DELETE /api/v1/keys/:id. +func (h *Handler) Delete(c *gin.Context) { + u := user.Current(c) + id, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + httpx.Fail(c, http.StatusBadRequest, "invalid key id") + return + } + if err := h.svc.Delete(u.ID, id); err != nil { + if errors.Is(err, ErrKeyNotFound) { + httpx.Fail(c, http.StatusNotFound, "api key not found") + return + } + httpx.Fail(c, http.StatusInternalServerError, "revoke key failed") + return + } + httpx.OK(c, gin.H{"ok": true}) +} diff --git a/server/internal/apikey/service.go b/server/internal/apikey/service.go new file mode 100644 index 0000000..3b260c3 --- /dev/null +++ b/server/internal/apikey/service.go @@ -0,0 +1,160 @@ +package apikey + +import ( + "encoding/json" + "errors" + "time" + + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/pkg/crypto" + "openteam/server/internal/pkg/rand" + "openteam/server/internal/store" +) + +var ErrKeyNotFound = errors.New("api key not found") + +type Service struct { + db *gorm.DB + log *zap.Logger +} + +func NewService(db *gorm.DB, log *zap.Logger) *Service { + return &Service{db: db, log: log} +} + +type CreateInput struct { + Name string `json:"name" binding:"required,max=128"` + QuotaTokensPerDay *int64 `json:"quotaTokensPerDay"` + QuotaRequestsPerDay *int `json:"quotaRequestsPerDay"` + AllowedModels []string `json:"allowedModels"` + ExpiresInDays *int `json:"expiresInDays"` +} + +type UpdateInput struct { + Name *string `json:"name"` + QuotaTokensPerDay *int64 `json:"quotaTokensPerDay"` + QuotaRequestsPerDay *int `json:"quotaRequestsPerDay"` + AllowedModels []string `json:"allowedModels"` + ExpiresAt *time.Time `json:"expiresAt"` + Status *string `json:"status"` +} + +// Create generates a key and returns it once along with the stored record. +func (s *Service) Create(userID int64, in CreateInput) (*store.ApiKey, string, error) { + plain, err := rand.Base62(48) + if err != nil { + return nil, "", err + } + full := "sk-" + plain + rec := &store.ApiKey{ + UserID: userID, + Name: in.Name, + KeyHash: crypto.HashSHA256(full), + KeyPrefix: "sk-" + plain[:8], + QuotaTokensPerDay: in.QuotaTokensPerDay, + QuotaRequestsPerDay: in.QuotaRequestsPerDay, + AllowedModels: in.AllowedModels, + Status: "active", + } + if in.ExpiresInDays != nil && *in.ExpiresInDays > 0 { + t := time.Now().AddDate(0, 0, *in.ExpiresInDays) + rec.ExpiresAt = &t + } + if err := s.db.Create(rec).Error; err != nil { + return nil, "", err + } + return rec, full, nil +} + +// List returns the user's keys (never the hash). +func (s *Service) List(userID int64) ([]store.ApiKey, error) { + var keys []store.ApiKey + err := s.db.Where("user_id = ?", userID).Order("id DESC").Find(&keys).Error + return keys, err +} + +// Update patches a key owned by the user. +func (s *Service) Update(userID, keyID int64, in UpdateInput) (*store.ApiKey, error) { + var k store.ApiKey + if err := s.db.First(&k, "id = ? AND user_id = ?", keyID, userID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, ErrKeyNotFound + } + return nil, err + } + updates := map[string]any{} + if in.Name != nil { + updates["name"] = *in.Name + } + if in.QuotaTokensPerDay != nil { + updates["quota_tokens_per_day"] = *in.QuotaTokensPerDay + } + if in.QuotaRequestsPerDay != nil { + updates["quota_requests_per_day"] = *in.QuotaRequestsPerDay + } + if in.AllowedModels != nil { + // GORM map updates skip the json serializer, so encode explicitly. + modelsJSON, err := json.Marshal(in.AllowedModels) + if err != nil { + return nil, err + } + updates["allowed_models"] = string(modelsJSON) + } + if in.ExpiresAt != nil { + updates["expires_at"] = in.ExpiresAt + } + if in.Status != nil { + updates["status"] = *in.Status + } + if len(updates) > 0 { + if err := s.db.Model(&k).Updates(updates).Error; err != nil { + return nil, err + } + } + return s.Get(userID, keyID) +} + +// Delete revokes a key (soft revoke by status). +func (s *Service) Delete(userID, keyID int64) error { + res := s.db.Model(&store.ApiKey{}). + Where("id = ? AND user_id = ?", keyID, userID). + Update("status", "revoked") + if res.Error != nil { + return res.Error + } + if res.RowsAffected == 0 { + return ErrKeyNotFound + } + return nil +} + +// Get loads one key owned by the user. +func (s *Service) Get(userID, keyID int64) (*store.ApiKey, error) { + var k store.ApiKey + if err := s.db.First(&k, "id = ? AND user_id = ?", keyID, userID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, ErrKeyNotFound + } + return nil, err + } + return &k, nil +} + +// Public returns a DTO without the hash. +func Public(k *store.ApiKey) map[string]any { + return map[string]any{ + "id": k.ID, + "userId": k.UserID, + "name": k.Name, + "keyPrefix": k.KeyPrefix, + "quotaTokensPerDay": k.QuotaTokensPerDay, + "quotaRequestsPerDay": k.QuotaRequestsPerDay, + "allowedModels": k.AllowedModels, + "expiresAt": k.ExpiresAt, + "status": k.Status, + "lastUsedAt": k.LastUsedAt, + "createdAt": k.CreatedAt, + } +} diff --git a/server/internal/billing/service.go b/server/internal/billing/service.go new file mode 100644 index 0000000..ba8f7ac --- /dev/null +++ b/server/internal/billing/service.go @@ -0,0 +1,117 @@ +package billing + +import ( + "errors" + "time" + + "github.com/shopspring/decimal" + "go.uber.org/zap" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "openteam/server/internal/store" +) + +var ( + ErrInsufficientBalance = errors.New("insufficient balance") + ErrUserNotFound = errors.New("user not found") +) + +// PriceSnapshot mirrors a model's prices at billing time. +type PriceSnapshot struct { + InputPrice decimal.Decimal + OutputPrice decimal.Decimal + CacheReadPrice decimal.Decimal +} + +// CostFromPrices computes cost for token counts using per-1M-token prices. +func CostFromPrices(in, out, cacheRead int64, prices PriceSnapshot) decimal.Decimal { + perM := decimal.NewFromInt(1_000_000) + cost := prices.InputPrice.Mul(decimal.NewFromInt(in)).Div(perM). + Add(prices.OutputPrice.Mul(decimal.NewFromInt(out)).Div(perM)). + Add(prices.CacheReadPrice.Mul(decimal.NewFromInt(cacheRead)).Div(perM)) + return cost.Round(8) +} + +// EstimateCost approximates cost from token estimates. +func (s *Service) EstimateCost(modelID int64, in, out, cacheRead int64) (decimal.Decimal, error) { + var m store.Model + if err := s.db.First(&m, modelID).Error; err != nil { + return decimal.Zero, err + } + return CostFromPrices(in, out, cacheRead, PriceSnapshot{ + InputPrice: m.InputPrice, OutputPrice: m.OutputPrice, CacheReadPrice: m.CacheReadPrice, + }), nil +} + +type Service struct { + db *gorm.DB + log *zap.Logger +} + +func NewService(db *gorm.DB, log *zap.Logger) *Service { + return &Service{db: db, log: log} +} + +// CheckBalance returns whether the user can afford the estimated cost. +func (s *Service) CheckBalance(userID int64, estimated decimal.Decimal) error { + var u store.User + if err := s.db.First(&u, userID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return ErrUserNotFound + } + return err + } + if u.Balance.LessThan(estimated) { + return ErrInsufficientBalance + } + return nil +} + +// Deduct atomically debits the user balance and appends a ledger entry. +func (s *Service) Deduct(userID int64, change decimal.Decimal, logType, refID string) (decimal.Decimal, error) { + return s.mutateBalance(userID, change.Neg(), logType, refID) +} + +// Credit adds to the user balance (recharge / refund / admin adjust). +func (s *Service) Credit(userID int64, change decimal.Decimal, logType, refID string) (decimal.Decimal, error) { + return s.mutateBalance(userID, change, logType, refID) +} + +// mutateBalance performs the balance update in a transaction so ledger and +// balance always agree. +func (s *Service) mutateBalance(userID int64, delta decimal.Decimal, logType, refID string) (decimal.Decimal, error) { + var after decimal.Decimal + err := s.db.Transaction(func(tx *gorm.DB) error { + var u store.User + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&u, userID).Error; err != nil { + return err + } + newBal := u.Balance.Add(delta) + if newBal.LessThan(decimal.Zero) { + return ErrInsufficientBalance + } + if err := tx.Model(&u).Update("balance", newBal).Error; err != nil { + return err + } + after = newBal + return tx.Create(&store.BalanceLog{ + UserID: userID, + Change: delta, + BalanceAfter: newBal, + Type: logType, + RefID: refID, + CreatedAt: time.Now(), + }).Error + }) + if err != nil { + return decimal.Zero, err + } + return after, nil +} + +// AdminAdjust changes a user's balance with an optional reason. +func (s *Service) AdminAdjust(userID int64, amount decimal.Decimal, reason string) (decimal.Decimal, error) { + return s.mutateBalance(userID, amount, "admin_adjust", reason) +} + diff --git a/server/internal/channel/health.go b/server/internal/channel/health.go new file mode 100644 index 0000000..b19a3ad --- /dev/null +++ b/server/internal/channel/health.go @@ -0,0 +1,53 @@ +package channel + +import ( + "time" + + "go.uber.org/zap" + + "openteam/server/internal/store" +) + +// StartHealthCheck runs the periodic health-check loop in a goroutine. +// Only channels that are enabled and have a bound test model are checked. +func (s *Service) StartHealthCheck() { + go func() { + ticker := time.NewTicker(s.cfg.HealthCheck.Interval) + defer ticker.Stop() + s.runHealthCheck() + for range ticker.C { + s.runHealthCheck() + } + }() +} + +func (s *Service) runHealthCheck() { + var channels []store.Channel + if err := s.db.Where("enabled = ?", true).Find(&channels).Error; err != nil { + s.log.Warn("health check query failed", zap.Error(err)) + return + } + for i := range channels { + ch := &channels[i] + // Skip channels already in cooldown until the cooldown elapses. + if ch.HealthStatus == "cooldown" { + continue + } + ok, _, _, err := s.Test(ch) + if err != nil || !ok { + s.MarkFailure(ch.ID) + continue + } + s.MarkHealthy(ch.ID, "healthy") + } +} + +// RecoverCooldown moves channels back from cooldown to degraded after the +// cooldown window, giving them another chance to pass the health check. +// Called periodically; keeps cooldown bounded. +func (s *Service) RecoverCooldown() { + cutoff := time.Now().Add(-s.cfg.HealthCheck.Cooldown) + s.db.Model(&store.Channel{}). + Where("health_status = ? AND updated_at < ?", "cooldown", cutoff). + Updates(map[string]any{"health_status": "degraded", "health_failures": 0}) +} diff --git a/server/internal/channel/ops.go b/server/internal/channel/ops.go new file mode 100644 index 0000000..d5431ca --- /dev/null +++ b/server/internal/channel/ops.go @@ -0,0 +1,220 @@ +package channel + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/pkg/crypto" + "openteam/server/internal/store" +) + +// EncryptKey encrypts an upstream API key for storage. +func (s *Service) EncryptKey(plain string) (string, error) { + return crypto.Encrypt(plain, s.master) +} + +// ImportModels fetches the channel's model list (GET /v1/models) and creates +// global Model records plus bindings. Returns the imported model names. +func (s *Service) ImportModels(ch *store.Channel) ([]string, error) { + key, err := s.DecryptKey(ch.APIKeyEnc) + if err != nil { + return nil, err + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + base := strings.TrimSuffix(ch.BaseURL, "/") + url := base + "/v1/models" + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+key) + req.Header.Set("Accept", "application/json") + + resp, err := http.DefaultClient.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + raw, err := io.ReadAll(io.LimitReader(resp.Body, 4<<20)) + if err != nil { + return nil, err + } + if resp.StatusCode >= 400 { + return nil, fmt.Errorf("channel returned %d: %s", resp.StatusCode, truncate(string(raw), 200)) + } + + var payload struct { + Data []struct { + ID string `json:"id"` + } `json:"data"` + } + if err := json.Unmarshal(raw, &payload); err != nil { + return nil, err + } + if len(payload.Data) == 0 { + return nil, errors.New("no models returned") + } + + imported := []string{} + err = s.db.Transaction(func(tx *gorm.DB) error { + for _, m := range payload.Data { + if m.ID == "" { + continue + } + var model store.Model + err := tx.Where("name = ?", m.ID).First(&model).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + model = store.Model{Name: m.ID, DisplayName: m.ID} + if err := tx.Create(&model).Error; err != nil { + return err + } + } else if err != nil { + return err + } + var binding store.ChannelModelBinding + err = tx.Where("channel_id = ? AND model_id = ?", ch.ID, model.ID).First(&binding).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + binding = store.ChannelModelBinding{ + ChannelID: ch.ID, ModelID: model.ID, UpstreamModel: m.ID, Weight: 1, + } + if err := tx.Create(&binding).Error; err != nil { + return err + } + } + imported = append(imported, m.ID) + } + return nil + }) + return imported, err +} + +// Test sends a cheap test request to the channel and reports success + latency. +func (s *Service) Test(ch *store.Channel) (bool, int, string, error) { + key, err := s.DecryptKey(ch.APIKeyEnc) + if err != nil { + return false, 0, "", err + } + model := s.cfg.HealthCheck.TestModel + if model == "" { + // Pick the first bound model, if any. + var b store.ChannelModelBinding + if err := s.db.Where("channel_id = ?", ch.ID).First(&b).Error; err == nil { + model = b.UpstreamModel + } + } + if model == "" { + return false, 0, "", errors.New("no test model configured; set HEALTHCHECK_TEST_MODEL or bind a model") + } + + timeout := time.Duration(s.cfg.HealthCheck.TimeoutMs) * time.Millisecond + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + // Test the first endpoint the channel actually serves, so channels that + // omit chat completions (e.g. Anthropic-only) still get a valid probe. + testFmt := pickTestFormat(ch) + var body []byte + switch testFmt { + case store.FormatAnthropic: + body, _ = json.Marshal(map[string]any{ + "model": model, "max_tokens": 8, + "messages": []any{map[string]any{"role": "user", "content": "ping"}}, + }) + case store.FormatOpenAIResponses: + body, _ = json.Marshal(map[string]any{"model": model, "input": "ping"}) + default: + body, _ = json.Marshal(map[string]any{ + "model": model, + "messages": []any{map[string]any{"role": "user", "content": "ping"}}, + "max_tokens": 8, + }) + } + + url := strings.TrimSuffix(ch.BaseURL, "/") + healthCheckPath(testFmt) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, strings.NewReader(string(body))) + if err != nil { + return false, 0, "", err + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+key) + if testFmt == store.FormatAnthropic { + req.Header.Set("anthropic-version", "2023-06-01") + } + + start := time.Now() + resp, err := http.DefaultClient.Do(req) + latency := int(time.Since(start).Milliseconds()) + if err != nil { + return false, latency, "", err + } + defer resp.Body.Close() + if resp.StatusCode >= 400 { + raw, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + return false, latency, "", fmt.Errorf("test request failed: %d %s", resp.StatusCode, truncate(string(raw), 160)) + } + return true, latency, model, nil +} + +// MarkHealthy updates health status and resets the failure counter. +func (s *Service) MarkHealthy(chID int64, status string) { + s.db.Model(&store.Channel{}).Where("id = ?", chID). + Updates(map[string]any{"health_status": status, "health_failures": 0}) +} + +// MarkFailure increments the failure counter and sets cooldown when maxed. +func (s *Service) MarkFailure(chID int64) { + var ch store.Channel + if err := s.db.First(&ch, chID).Error; err != nil { + return + } + failures := ch.HealthFailures + 1 + status := "degraded" + if failures >= s.cfg.HealthCheck.MaxFailures { + status = "cooldown" + } + s.db.Model(&ch).Updates(map[string]any{"health_failures": failures, "health_status": status}) + s.log.Info("channel health failure", zap.Int64("channel_id", chID), + zap.Int("failures", failures), zap.String("status", status)) +} + +// pickTestFormat chooses a supported format to probe, preferring the cheapest +// endpoint. Falls back to chat completions so legacy channels keep working. +func pickTestFormat(ch *store.Channel) string { + for _, want := range []string{store.FormatOpenAIChat, store.FormatAnthropic, store.FormatOpenAIResponses} { + for _, f := range ch.FormatsResolved() { + if f == want { + return want + } + } + } + return store.FormatOpenAIChat +} + +func healthCheckPath(f string) string { + switch f { + case store.FormatAnthropic: + return "/v1/messages" + case store.FormatOpenAIResponses: + return "/v1/responses" + default: + return "/v1/chat/completions" + } +} + +func truncate(s string, n int) string { + if len(s) <= n { + return s + } + return s[:n] + "..." +} diff --git a/server/internal/channel/service.go b/server/internal/channel/service.go new file mode 100644 index 0000000..4bf7c7e --- /dev/null +++ b/server/internal/channel/service.go @@ -0,0 +1,156 @@ +package channel + +import ( + "errors" + "math/rand" + "sort" + "sync" + + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/config" + "openteam/server/internal/pkg/crypto" + "openteam/server/internal/store" +) + +var ErrNoChannel = errors.New("no available channel for model") + +// Service manages channels, model bindings and load-balanced selection. +type Service struct { + db *gorm.DB + cfg *config.Config + log *zap.Logger + + mu sync.RWMutex + sem map[int64]chan struct{} // per-channel concurrency limiter + master string +} + +func NewService(db *gorm.DB, cfg *config.Config, log *zap.Logger) *Service { + return &Service{ + db: db, cfg: cfg, log: log, + sem: map[int64]chan struct{}{}, + master: cfg.MasterKey, + } +} + +// DecryptKey decrypts a channel's stored upstream API key. +func (s *Service) DecryptKey(enc string) (string, error) { + return crypto.Decrypt(enc, s.master) +} + +// Acquire takes a concurrency slot for a channel (blocks if saturated). +func (s *Service) Acquire(channelID int64) (func(), error) { + s.mu.Lock() + lim, ok := s.sem[channelID] + if !ok { + var rec store.Channel + if err := s.db.First(&rec, channelID).Error; err != nil { + s.mu.Unlock() + return nil, err + } + limit := rec.MaxConcurrency + if limit <= 0 { + limit = 100 + } + lim = make(chan struct{}, limit) + s.sem[channelID] = lim + } + s.mu.Unlock() + lim <- struct{}{} + return func() { <-lim }, nil +} + +// ResolveModel returns the model registry record. +func (s *Service) ResolveModel(name string) (*store.Model, error) { + var m store.Model + if err := s.db.Where("name = ? AND enabled = ?", name, true).First(&m).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, errors.New("model not found: " + name) + } + return nil, err + } + return &m, nil +} + +// SelectChannel picks a healthy channel bound to the given model. +// Channels with higher priority (lower number) and higher weight win; +// cooldown channels are skipped, degraded channels are deprioritized. +// Channels in exclude (already tried in this request) are skipped. +func (s *Service) SelectChannel(modelID int64, exclude map[int64]bool) (*store.Channel, *store.ChannelModelBinding, error) { + var bindings []store.ChannelModelBinding + if err := s.db.Preload("Channel").Where("model_id = ?", modelID).Find(&bindings).Error; err != nil { + return nil, nil, err + } + var candidates []struct { + ch *store.Channel + b *store.ChannelModelBinding + score float64 + } + for i := range bindings { + ch := &bindings[i].Channel + if !ch.Enabled { + continue + } + if ch.HealthStatus == "cooldown" { + continue + } + if exclude[ch.ID] { + continue + } + score := float64(bindings[i].Weight) + if ch.HealthStatus == "degraded" { + score *= 0.1 + } + candidates = append(candidates, struct { + ch *store.Channel + b *store.ChannelModelBinding + score float64 + }{ch: ch, b: &bindings[i], score: score}) + } + if len(candidates) == 0 { + return nil, nil, ErrNoChannel + } + // Sort by priority asc, then weight desc. + sort.SliceStable(candidates, func(i, j int) bool { + if candidates[i].ch.Priority != candidates[j].ch.Priority { + return candidates[i].ch.Priority < candidates[j].ch.Priority + } + return candidates[i].score > candidates[j].score + }) + // Weighted random pick among the top priority group. + total := 0.0 + for _, c := range candidates { + total += c.score + } + if total <= 0 { + return candidates[0].ch, candidates[0].b, nil + } + pick := rand.Float64() * total + for _, c := range candidates { + pick -= c.score + if pick <= 0 { + return c.ch, c.b, nil + } + } + return candidates[0].ch, candidates[0].b, nil +} + +// ListModelsForChannel returns the model names a channel serves. +func (s *Service) ListModelsForChannel(channelID int64) ([]string, error) { + var names []string + err := s.db.Table("channel_model_bindings"). + Joins("JOIN models ON models.id = channel_model_bindings.model_id"). + Where("channel_model_bindings.channel_id = ? AND models.enabled = ?", channelID, true). + Order("models.sort ASC, models.id ASC"). + Pluck("models.name", &names).Error + return names, err +} + +// ListEnabledModels returns all enabled models for GET /v1/models. +func (s *Service) ListEnabledModels() ([]store.Model, error) { + var models []store.Model + err := s.db.Where("enabled = ?", true).Order("sort ASC, id ASC").Find(&models).Error + return models, err +} diff --git a/server/internal/config/config.go b/server/internal/config/config.go new file mode 100644 index 0000000..5c5d85b --- /dev/null +++ b/server/internal/config/config.go @@ -0,0 +1,143 @@ +package config + +import ( + "fmt" + "os" + "strings" + "time" +) + +// Config is the runtime configuration, loaded from environment variables. +type Config struct { + Env string + Debug bool + + HTTPPort string + PublicBase string // external base URL, used for cookies + AllowOrigins []string + + DB struct { + Driver string // "sqlite" (dev default) or "postgres" + DSN string + } + + Redis struct { + Addr string + Password string + Enabled bool + } + + Auth struct { + AccessTokenTTL time.Duration + RefreshTokenTTL time.Duration + JWTSecret string + RefreshCookieName string + RefreshCookieSecure bool // false for local http dev; set true behind TLS + RefreshCookieSameSite string + } + + // MasterKey encrypts channel API keys at rest (AES-GCM). + MasterKey string + + Proxy struct { + DefaultMaxTokens int + DefaultTimeoutMs int + BillingExactBalance bool // reject when estimated cost > balance + MaxRetries int // additional channel attempts on transport/5xx failures + } + + Registration struct { + Mode string // "open" | "invite" + } + + HealthCheck struct { + Interval time.Duration + MaxFailures int + Cooldown time.Duration + TimeoutMs int + TestModel string + } + + RateLimit struct { + RequestsPerMin int // per user global limiter + Burst int + } + + MetricsEnabled bool +} + +func Load() *Config { + c := &Config{} + c.Env = get("APP_ENV", "development") + c.Debug = strings.EqualFold(get("DEBUG", "false"), "true") + c.HTTPPort = get("HTTP_PORT", "8080") + c.PublicBase = get("PUBLIC_BASE", "http://localhost:8080") + if o := get("CORS_ORIGINS", "*"); o != "*" { + c.AllowOrigins = strings.Split(o, ",") + } else { + c.AllowOrigins = []string{"*"} + } + + c.DB.Driver = get("DB_DRIVER", "sqlite") + if c.DB.Driver == "postgres" { + c.DB.DSN = get("DATABASE_URL", "host=localhost user=postgres password=postgres dbname=openteam port=5432 sslmode=disable") + } else { + path := get("SQLITE_PATH", "data/openteam.db") + c.DB.DSN = path + } + + c.Redis.Addr = get("REDIS_ADDR", "localhost:6379") + c.Redis.Password = get("REDIS_PASSWORD", "") + c.Redis.Enabled = strings.EqualFold(get("REDIS_ENABLED", "false"), "true") + + c.Auth.AccessTokenTTL = duration(get("ACCESS_TOKEN_TTL", "2h"), 2*time.Hour) + c.Auth.RefreshTokenTTL = duration(get("REFRESH_TOKEN_TTL", "168h"), 7*24*time.Hour) + c.Auth.JWTSecret = get("JWT_SECRET", "dev-only-secret-change-me") + c.Auth.RefreshCookieName = get("REFRESH_COOKIE_NAME", "ot_refresh") + c.Auth.RefreshCookieSecure = strings.EqualFold(get("REFRESH_COOKIE_SECURE", "false"), "true") + c.Auth.RefreshCookieSameSite = get("REFRESH_COOKIE_SAMESITE", "lax") + + c.MasterKey = get("MASTER_KEY", "dev-only-master-key-change-me") + + c.Proxy.DefaultMaxTokens = intVal(get("DEFAULT_MAX_TOKENS", "4096"), 4096) + c.Proxy.DefaultTimeoutMs = intVal(get("PROXY_TIMEOUT_MS", "300000"), 300000) + c.Proxy.BillingExactBalance = strings.EqualFold(get("BILLING_EXACT_BALANCE", "false"), "true") + c.Proxy.MaxRetries = intVal(get("PROXY_MAX_RETRIES", "1"), 1) + + c.Registration.Mode = get("REGISTRATION_MODE", "open") + + c.HealthCheck.Interval = duration(get("HEALTHCHECK_INTERVAL", "60s"), time.Minute) + c.HealthCheck.MaxFailures = intVal(get("HEALTHCHECK_MAX_FAILURES", "3"), 3) + c.HealthCheck.Cooldown = duration(get("HEALTHCHECK_COOLDOWN", "300s"), 5*time.Minute) + c.HealthCheck.TimeoutMs = intVal(get("HEALTHCHECK_TIMEOUT_MS", "15000"), 15000) + c.HealthCheck.TestModel = get("HEALTHCHECK_TEST_MODEL", "") + + c.RateLimit.RequestsPerMin = intVal(get("RATE_LIMIT_PER_MIN", "60"), 60) + c.RateLimit.Burst = intVal(get("RATE_LIMIT_BURST", "120"), 120) + + c.MetricsEnabled = strings.EqualFold(get("METRICS_ENABLED", "false"), "true") + return c +} + +func get(key, def string) string { + if v := os.Getenv(key); v != "" { + return v + } + return def +} + +func intVal(s string, def int) int { + n := 0 + if _, err := fmt.Sscanf(s, "%d", &n); err != nil || n <= 0 { + return def + } + return n +} + +func duration(s string, def time.Duration) time.Duration { + d, err := time.ParseDuration(s) + if err != nil || d <= 0 { + return def + } + return d +} diff --git a/server/internal/pkg/crypto/crypto.go b/server/internal/pkg/crypto/crypto.go new file mode 100644 index 0000000..5e42d07 --- /dev/null +++ b/server/internal/pkg/crypto/crypto.go @@ -0,0 +1,64 @@ +package crypto + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "errors" + "fmt" + "io" +) + +// Encrypt seals a secret with AES-GCM. The master key can be any string; +// it is hashed to a fixed-size AES key. +func Encrypt(plaintext, masterKey string) (string, error) { + key := sha256.Sum256([]byte(masterKey)) + block, err := aes.NewCipher(key[:]) + if err != nil { + return "", err + } + gcm, err := cipher.NewGCM(block) + if err != nil { + return "", err + } + nonce := make([]byte, gcm.NonceSize()) + if _, err := io.ReadFull(rand.Reader, nonce); err != nil { + return "", err + } + sealed := gcm.Seal(nonce, nonce, []byte(plaintext), nil) + return base64.StdEncoding.EncodeToString(sealed), nil +} + +// Decrypt opens a ciphertext produced by Encrypt. +func Decrypt(ciphertext, masterKey string) (string, error) { + key := sha256.Sum256([]byte(masterKey)) + data, err := base64.StdEncoding.DecodeString(ciphertext) + if err != nil { + return "", err + } + block, err := aes.NewCipher(key[:]) + if err != nil { + return "", err + } + gcm, err := cipher.NewGCM(block) + if err != nil { + return "", err + } + if len(data) < gcm.NonceSize() { + return "", errors.New("ciphertext too short") + } + nonce, sealed := data[:gcm.NonceSize()], data[gcm.NonceSize():] + plain, err := gcm.Open(nil, nonce, sealed, nil) + if err != nil { + return "", fmt.Errorf("decrypt failed (bad master key?): %w", err) + } + return string(plain), nil +} + +// HashSHA256 returns the hex SHA-256 of a string (used for API key lookup). +func HashSHA256(s string) string { + sum := sha256.Sum256([]byte(s)) + return fmt.Sprintf("%x", sum) +} diff --git a/server/internal/pkg/httpx/httpx.go b/server/internal/pkg/httpx/httpx.go new file mode 100644 index 0000000..779037f --- /dev/null +++ b/server/internal/pkg/httpx/httpx.go @@ -0,0 +1,42 @@ +package httpx + +import ( + "net/http" + + "github.com/gin-gonic/gin" +) + +// APIError is the standard error body for the management API. +type APIError struct { + Code int `json:"code"` + Message string `json:"message"` +} + +// OK writes a JSON success response. +func OK(c *gin.Context, data any) { + c.JSON(http.StatusOK, gin.H{"code": 0, "data": data}) +} + +// Created writes a 201 response. +func Created(c *gin.Context, data any) { + c.JSON(http.StatusCreated, gin.H{"code": 0, "data": data}) +} + +// Fail writes an error response with the given status. +func Fail(c *gin.Context, status int, message string) { + c.AbortWithStatusJSON(status, gin.H{"code": status, "message": message}) +} + +// FailWithCode writes an error with a custom business code. +func FailWithCode(c *gin.Context, status, code int, message string) { + c.AbortWithStatusJSON(status, gin.H{"code": code, "message": message}) +} + +// Bind parses the JSON body and aborts with 400 on failure. +func Bind(c *gin.Context, dst any) bool { + if err := c.ShouldBindJSON(dst); err != nil { + Fail(c, http.StatusBadRequest, "invalid request body: "+err.Error()) + return false + } + return true +} diff --git a/server/internal/pkg/jwt/jwt.go b/server/internal/pkg/jwt/jwt.go new file mode 100644 index 0000000..4a74557 --- /dev/null +++ b/server/internal/pkg/jwt/jwt.go @@ -0,0 +1,66 @@ +package jwt + +import ( + "errors" + "strconv" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +// Claims is the payload of a signed token. +type Claims struct { + UserID int64 `json:"uid"` + Username string `json:"username"` + Role string `json:"role"` + Type string `json:"typ"` // access | refresh + jwt.RegisteredClaims +} + +func sign(secret string, c Claims) (string, error) { + t := jwt.NewWithClaims(jwt.SigningMethodHS256, c) + return t.SignedString([]byte(secret)) +} + +// SignAccess issues a short-lived access token. +func SignAccess(secret string, userID int64, username, role string, ttl time.Duration) (string, error) { + return sign(secret, Claims{ + UserID: userID, Username: username, Role: role, Type: "access", + RegisteredClaims: jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(time.Now().Add(ttl)), + IssuedAt: jwt.NewNumericDate(time.Now()), + Subject: fmtID(userID), + }, + }) +} + +// SignRefresh issues a long-lived refresh token. +func SignRefresh(secret string, userID int64, ttl time.Duration) (string, error) { + return sign(secret, Claims{ + UserID: userID, Type: "refresh", + RegisteredClaims: jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(time.Now().Add(ttl)), + IssuedAt: jwt.NewNumericDate(time.Now()), + Subject: fmtID(userID), + }, + }) +} + +// Parse validates a token and returns its claims. +func Parse(secret, token string) (*Claims, error) { + var c Claims + parsed, err := jwt.ParseWithClaims(token, &c, func(t *jwt.Token) (interface{}, error) { + if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok { + return nil, errors.New("unexpected signing method") + } + return []byte(secret), nil + }) + if err != nil || !parsed.Valid { + return nil, errors.New("invalid token") + } + return &c, nil +} + +func fmtID(id int64) string { + return strconv.FormatInt(id, 10) +} diff --git a/server/internal/pkg/password/password.go b/server/internal/pkg/password/password.go new file mode 100644 index 0000000..d186fd1 --- /dev/null +++ b/server/internal/pkg/password/password.go @@ -0,0 +1,62 @@ +package password + +import ( + "crypto/rand" + "crypto/subtle" + "encoding/base64" + "errors" + "fmt" + "strings" + + "golang.org/x/crypto/argon2" +) + +const ( + argonTime = 3 + argonMemory = 64 * 1024 + argonThreads = 2 + argonKeyLen = 32 + argonSaltLen = 16 +) + +// Hash hashes a plaintext password with argon2id. +func Hash(plain string) (string, error) { + salt := make([]byte, argonSaltLen) + if _, err := rand.Read(salt); err != nil { + return "", err + } + key := argon2.IDKey([]byte(plain), salt, argonTime, argonMemory, argonThreads, argonKeyLen) + enc := base64.RawStdEncoding + return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s", + argon2.Version, argonMemory, argonTime, argonThreads, + enc.EncodeToString(salt), enc.EncodeToString(key)), nil +} + +// Verify checks a plaintext password against an argon2id hash string. +func Verify(plain, encoded string) (bool, error) { + parts := strings.Split(encoded, "$") + if len(parts) != 6 || parts[1] != "argon2id" { + return false, errors.New("malformed password hash") + } + var version int + var memory uint32 + var time_ uint32 + var threads uint8 + if _, err := fmt.Sscanf(parts[2], "v=%d", &version); err != nil { + return false, err + } + if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time_, &threads); err != nil { + return false, err + } + enc := base64.RawStdEncoding + salt, err := enc.DecodeString(parts[4]) + if err != nil { + return false, err + } + want, err := enc.DecodeString(parts[5]) + if err != nil { + return false, err + } + got := argon2.IDKey([]byte(plain), salt, time_, memory, threads, uint32(len(want))) + return subtle.ConstantTimeCompare(got, want) == 1, nil +} diff --git a/server/internal/pkg/rand/rand.go b/server/internal/pkg/rand/rand.go new file mode 100644 index 0000000..0f76dfd --- /dev/null +++ b/server/internal/pkg/rand/rand.go @@ -0,0 +1,21 @@ +package rand + +import ( + crand "crypto/rand" + "math/big" +) + +const base62 = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" + +// Base62 returns a cryptographically random base62 string of length n. +func Base62(n int) (string, error) { + out := make([]byte, n) + for i := range out { + idx, err := crand.Int(crand.Reader, big.NewInt(int64(len(base62)))) + if err != nil { + return "", err + } + out[i] = base62[idx.Int64()] + } + return string(out), nil +} diff --git a/server/internal/pkg/ratelimit/ratelimit.go b/server/internal/pkg/ratelimit/ratelimit.go new file mode 100644 index 0000000..23d18f4 --- /dev/null +++ b/server/internal/pkg/ratelimit/ratelimit.go @@ -0,0 +1,73 @@ +package ratelimit + +import ( + "sync" + "time" +) + +// Limiter is a token-bucket rate limiter keyed by string. +type Limiter struct { + mu sync.Mutex + rate float64 // tokens per second + burst float64 + tokens map[string]*bucket +} + +type bucket struct { + tokens float64 + lastFill time.Time +} + +// New creates a limiter refilling `rate` tokens/sec with `burst` capacity. +func New(rate float64, burst int) *Limiter { + return &Limiter{ + rate: rate, + burst: float64(burst), + tokens: map[string]*bucket{}, + } +} + +// Allow checks whether `key` may take one token now. +func (l *Limiter) Allow(key string) bool { + return l.Take(key, 1) +} + +// Take checks whether `key` may take n tokens now. +func (l *Limiter) Take(key string, n float64) bool { + l.mu.Lock() + defer l.mu.Unlock() + now := time.Now() + b, ok := l.tokens[key] + if !ok { + b = &bucket{tokens: l.burst, lastFill: now} + l.tokens[key] = b + } + // Refill based on elapsed time. + elapsed := now.Sub(b.lastFill).Seconds() + b.tokens = minF(l.burst, b.tokens+elapsed*l.rate) + b.lastFill = now + if b.tokens >= n { + b.tokens -= n + return true + } + return false +} + +// Sweep removes idle buckets to bound memory. Call periodically. +func (l *Limiter) Sweep(olderThan time.Duration) { + l.mu.Lock() + defer l.mu.Unlock() + cutoff := time.Now().Add(-olderThan) + for k, b := range l.tokens { + if b.lastFill.Before(cutoff) { + delete(l.tokens, k) + } + } +} + +func minF(a, b float64) float64 { + if a < b { + return a + } + return b +} diff --git a/server/internal/proxy/claude/types.go b/server/internal/proxy/claude/types.go new file mode 100644 index 0000000..794b3e8 --- /dev/null +++ b/server/internal/proxy/claude/types.go @@ -0,0 +1,50 @@ +package claude + +import "encoding/json" + +// Request is an Anthropic Messages request. +type Request struct { + Model string `json:"model"` + MaxTokens int `json:"max_tokens"` + Stream bool `json:"stream"` + System json.RawMessage `json:"system"` + Messages json.RawMessage `json:"messages"` + Temperature *float64 `json:"temperature"` + TopP *float64 `json:"top_p"` + TopK *int `json:"top_k"` + StopSequences json.RawMessage `json:"stop_sequences"` + Tools json.RawMessage `json:"tools"` + ToolChoice json.RawMessage `json:"tool_choice"` + Metadata json.RawMessage `json:"metadata"` +} + +// Usage is the Claude usage block. +type Usage struct { + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + CacheReadInputTokens int64 `json:"cache_read_input_tokens,omitempty"` + CacheCreationInputTokens int64 `json:"cache_creation_input_tokens,omitempty"` +} + +// MessageResponse is the non-stream response. +type MessageResponse struct { + ID string `json:"id"` + Type string `json:"type"` + Role string `json:"role"` + Model string `json:"model"` + Content json.RawMessage `json:"content"` + StopReason string `json:"stop_reason"` + StopSequence string `json:"stop_sequence"` + Usage *Usage `json:"usage"` +} + +// StreamEvent is one event in the Claude streaming event sequence. +type StreamEvent struct { + Type string `json:"type"` + Message json.RawMessage `json:"message,omitempty"` + Index *int `json:"index,omitempty"` + Delta json.RawMessage `json:"delta,omitempty"` + Usage json.RawMessage `json:"usage,omitempty"` + ContentBlock json.RawMessage `json:"content_block,omitempty"` + Error json.RawMessage `json:"error,omitempty"` +} diff --git a/server/internal/proxy/convert/canonical.go b/server/internal/proxy/convert/canonical.go new file mode 100644 index 0000000..ad05d6f --- /dev/null +++ b/server/internal/proxy/convert/canonical.go @@ -0,0 +1,196 @@ +package convert + +import ( + "encoding/json" + "strings" + + "openteam/server/internal/proxy/claude" + "openteam/server/internal/proxy/openai" + "openteam/server/internal/proxy/responses" +) + +// CanonicalRequest is the gateway-internal standard form (OpenAI chat shape). +// Every protocol is converted into this before being emitted to a channel. +type CanonicalRequest struct { + Model string + System string + Messages []openai.ChatMessage + Stream bool + Temperature *float64 + TopP *float64 + MaxTokens *int + Stop []string + Tools []openai.Tool + ToolChoice json.RawMessage + ResponseFormat json.RawMessage + // Extra passthrough-only fields for OpenAI channels. + RawOpenAIExtras map[string]json.RawMessage +} + +// contentString returns a plain-text representation of a message content, +// handling both string and structured (Claude-style blocks) content. +func contentString(raw json.RawMessage) string { + if len(raw) == 0 { + return "" + } + var s string + if err := json.Unmarshal(raw, &s); err == nil { + return s + } + var blocks []struct { + Type string `json:"type"` + Text string `json:"text"` + } + if err := json.Unmarshal(raw, &blocks); err == nil { + out := "" + for _, b := range blocks { + if b.Text != "" { + out += b.Text + } + } + return out + } + return "" +} + +// requestFromOpenAIChat parses an OpenAI chat body into the canonical form. +func requestFromOpenAIChat(body []byte) (*CanonicalRequest, error) { + var r openai.ChatRequest + if err := json.Unmarshal(body, &r); err != nil { + return nil, err + } + var msgs []openai.ChatMessage + if err := json.Unmarshal(r.Messages, &msgs); err != nil { + return nil, err + } + req := &CanonicalRequest{ + Model: r.Model, + Messages: msgs, + Stream: r.Stream, + Temperature: r.Temperature, + TopP: r.TopP, + Stop: parseStop(r.Stop), + ToolChoice: r.ToolChoice, + } + if r.MaxTokens != nil { + req.MaxTokens = r.MaxTokens + } else if r.MaxCompl != nil { + req.MaxTokens = r.MaxCompl + } + if len(r.Tools) > 0 { + _ = json.Unmarshal(r.Tools, &req.Tools) + } + req.ResponseFormat = r.ResponseFmt + req.System = extractSystem(msgs) + req.RawOpenAIExtras = rawExtras(body, "model", "messages", "stream", "temperature", "top_p", "max_tokens", "max_completion_tokens", "stop", "tools", "tool_choice", "response_format", "user") + return req, nil +} + +func requestFromClaude(body []byte) (*CanonicalRequest, error) { + var r claude.Request + if err := json.Unmarshal(body, &r); err != nil { + return nil, err + } + var msgs []openai.ChatMessage + // Claude messages: content can be a string or blocks; tool_use/tool_result + // blocks map to assistant.tool_calls and role=tool messages. + if err := claudeMessagesToChat(r.Messages, &msgs); err != nil { + return nil, err + } + req := &CanonicalRequest{ + Model: r.Model, + Messages: msgs, + Stream: r.Stream, + Temperature: r.Temperature, + TopP: r.TopP, + Stop: parseStop(r.StopSequences), + } + if r.MaxTokens > 0 { + mt := r.MaxTokens + req.MaxTokens = &mt + } + if len(r.Tools) > 0 { + _ = json.Unmarshal(r.Tools, &req.Tools) + } + req.ToolChoice = r.ToolChoice + // Claude system can be a string or array of blocks. + req.System = contentString(r.System) + return req, nil +} + +func requestFromResponses(body []byte) (*CanonicalRequest, error) { + var r responses.Request + if err := json.Unmarshal(body, &r); err != nil { + return nil, err + } + req := &CanonicalRequest{ + Model: r.Model, + Stream: r.Stream, + } + // instructions → system. + req.System = contentString(r.Instructions) + // Parse input items into chat messages. + msgs, err := responsesInputToChat(r.Input) + if err != nil { + return nil, err + } + req.Messages = msgs + if r.MaxOutputTokens != nil { + req.MaxTokens = r.MaxOutputTokens + } + if len(r.Tools) > 0 { + _ = json.Unmarshal(r.Tools, &req.Tools) + } + // output_format / text.format → response_format. + if len(r.OutputFormat) > 0 { + req.ResponseFormat = r.OutputFormat + } else if len(r.Text) > 0 { + var t struct { + Format json.RawMessage `json:"format"` + } + if json.Unmarshal(r.Text, &t) == nil && len(t.Format) > 0 { + req.ResponseFormat = t.Format + } + } + return req, nil +} + +func extractSystem(msgs []openai.ChatMessage) string { + var parts []string + for _, m := range msgs { + if m.Role == "system" { + if s := contentString(m.Content); s != "" { + parts = append(parts, s) + } + } + } + return strings.Join(parts, "\n") +} + +func parseStop(raw json.RawMessage) []string { + if len(raw) == 0 { + return nil + } + var one string + if err := json.Unmarshal(raw, &one); err == nil { + return []string{one} + } + var many []string + if err := json.Unmarshal(raw, &many); err == nil { + return many + } + return nil +} + +// rawExtras captures fields not otherwise modeled so they can be re-emitted +// on OpenAI passthrough-style conversions. +func rawExtras(body []byte, skip ...string) map[string]json.RawMessage { + var obj map[string]json.RawMessage + if json.Unmarshal(body, &obj) != nil { + return nil + } + for _, k := range skip { + delete(obj, k) + } + return obj +} diff --git a/server/internal/proxy/convert/chat_claude.go b/server/internal/proxy/convert/chat_claude.go new file mode 100644 index 0000000..188bd25 --- /dev/null +++ b/server/internal/proxy/convert/chat_claude.go @@ -0,0 +1,269 @@ +package convert + +import ( + "encoding/json" + + "openteam/server/internal/proxy/openai" +) + +// claudeMessagesToChat converts Claude messages JSON into OpenAI chat messages. +func claudeMessagesToChat(raw json.RawMessage, out *[]openai.ChatMessage) error { + var msgs []struct { + Role string `json:"role"` + Content json.RawMessage `json:"content"` + } + if err := json.Unmarshal(raw, &msgs); err != nil { + return err + } + for _, m := range msgs { + var text string + if err := json.Unmarshal(m.Content, &text); err == nil { + content, _ := json.Marshal(text) + *out = append(*out, openai.ChatMessage{Role: m.Role, Content: content}) + continue + } + var blocks []map[string]json.RawMessage + if err := json.Unmarshal(m.Content, &blocks); err != nil { + return err + } + // Split blocks into: plain content parts, tool_use (→ tool_calls), + // and tool_result (→ separate role=tool messages). + var parts []json.RawMessage + var toolCalls []json.RawMessage + for _, b := range blocks { + var typ string + _ = json.Unmarshal(b["type"], &typ) + switch typ { + case "tool_use": + tc := map[string]any{ + "id": rawString(b["id"]), + "type": "function", + "function": map[string]any{ + "name": rawString(b["name"]), + "arguments": string(b["input"]), + }, + } + encoded, _ := json.Marshal(tc) + toolCalls = append(toolCalls, encoded) + case "tool_result": + msg := openai.ChatMessage{ + Role: "tool", + ToolCallID: rawString(b["tool_use_id"]), + } + // content may be a string or array of text blocks. + if b["content"] != nil { + var s string + if json.Unmarshal(b["content"], &s) == nil { + content, _ := json.Marshal(s) + msg.Content = content + } else { + var texts []struct { + Type string `json:"type"` + Text string `json:"text"` + } + if json.Unmarshal(b["content"], &texts) == nil { + var buf string + for _, t := range texts { + buf += t.Text + } + content, _ := json.Marshal(buf) + msg.Content = content + } + } + } + encoded, _ := json.Marshal(msg) + *out = append(*out, msg) + _ = encoded + default: + // text / image blocks → OpenAI content part. + part, err := claudeBlockToOpenAIContent(b) + if err != nil { + return err + } + if part != nil { + parts = append(parts, part) + } + } + } + if len(toolCalls) > 0 { + msg := openai.ChatMessage{Role: "assistant"} + if len(parts) > 0 { + content, _ := json.Marshal(partsToText(parts)) + msg.Content = content + } + tcArr, _ := json.Marshal(toolCalls) + msg.ToolCalls = tcArr + *out = append(*out, msg) + } else if len(parts) > 0 { + if len(parts) == 1 { + // Collapse a single text part back to a plain string. + var s string + if json.Unmarshal(parts[0], &s) == nil { + *out = append(*out, openai.ChatMessage{Role: m.Role, Content: parts[0]}) + } else { + arr, _ := json.Marshal(parts) + *out = append(*out, openai.ChatMessage{Role: m.Role, Content: arr}) + } + } else { + arr, _ := json.Marshal(parts) + *out = append(*out, openai.ChatMessage{Role: m.Role, Content: arr}) + } + } + } + return nil +} + +func claudeBlockToOpenAIContent(b map[string]json.RawMessage) (json.RawMessage, error) { + var typ string + _ = json.Unmarshal(b["type"], &typ) + switch typ { + case "text": + part := map[string]any{"type": "text", "text": rawString(b["text"])} + return json.Marshal(part) + case "image": + // b["source"] may not be present; be defensive. + if b["source"] != nil { + var src struct { + Type string `json:"type"` + URL string `json:"url"` + Data string `json:"data"` + MediaType string `json:"media_type"` + } + _ = json.Unmarshal(b["source"], &src) + if src.URL != "" { + return json.Marshal(map[string]any{"type": "image_url", "image_url": map[string]any{"url": src.URL}}) + } + if src.Data != "" { + return json.Marshal(map[string]any{"type": "image_url", "image_url": map[string]any{"url": "data:" + src.MediaType + ";base64," + src.Data}}) + } + } + return nil, nil + default: + return nil, nil + } +} + +func partsToText(parts []json.RawMessage) string { + var out string + for _, p := range parts { + var t struct { + Type string `json:"type"` + Text string `json:"text"` + } + if json.Unmarshal(p, &t) == nil { + out += t.Text + } + } + return out +} + +func rawString(raw json.RawMessage) string { + if len(raw) == 0 { + return "" + } + var s string + if err := json.Unmarshal(raw, &s); err != nil { + return string(raw) + } + return s +} + +// chatToClaudeMessages converts chat messages (with system already removed) +// into Claude Messages body. +func chatToClaudeMessages(msgs []openai.ChatMessage, system string) (map[string]any, error) { + claudeMsgs := []map[string]any{} + for _, m := range msgs { + if m.Role == "system" { + continue + } + // Tool calls on assistant messages → tool_use blocks. + if m.Role == "assistant" && len(m.ToolCalls) > 0 { + content := []map[string]any{} + // Preserve any text content. + if text := contentString(m.Content); text != "" { + content = append(content, map[string]any{"type": "text", "text": text}) + } + var calls []struct { + ID string `json:"id"` + Function struct { + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` + } `json:"function"` + } + _ = json.Unmarshal(m.ToolCalls, &calls) + for _, c := range calls { + var input map[string]any + if err := json.Unmarshal(c.Function.Arguments, &input); err != nil { + input = map[string]any{"raw": string(c.Function.Arguments)} + } + content = append(content, map[string]any{ + "type": "tool_use", + "id": c.ID, + "name": c.Function.Name, + "input": input, + }) + } + claudeMsgs = append(claudeMsgs, map[string]any{"role": "assistant", "content": content}) + continue + } + // Tool role messages → tool_result blocks. + if m.Role == "tool" { + content := []map[string]any{{ + "type": "tool_result", + "tool_use_id": m.ToolCallID, + "content": contentString(m.Content), + }} + claudeMsgs = append(claudeMsgs, map[string]any{"role": "user", "content": content}) + continue + } + claudeMsgs = append(claudeMsgs, map[string]any{"role": m.Role, "content": contentString(m.Content)}) + } + if len(claudeMsgs) == 0 { + claudeMsgs = append(claudeMsgs, map[string]any{"role": "user", "content": "Hi"}) + } + body := map[string]any{"messages": claudeMsgs} + if system != "" { + body["system"] = system + } + return body, nil +} + +// toolsToClaude converts OpenAI function tools to Claude tools. +func toolsToClaude(tools []openai.Tool) []map[string]any { + out := []map[string]any{} + for _, t := range tools { + if t.Type != "" && t.Type != "function" { + continue + } + out = append(out, map[string]any{ + "name": t.Function.Name, + "description": t.Function.Description, + "input_schema": json.RawMessage(t.Function.Parameters), + }) + } + return out +} + +// claudeToolsToOpenAI converts Claude tools to OpenAI function tools. +func claudeToolsToOpenAI(tools json.RawMessage) []openai.Tool { + var list []struct { + Name string `json:"name"` + Description string `json:"description"` + InputSchema json.RawMessage `json:"input_schema"` + } + if err := json.Unmarshal(tools, &list); err != nil { + return nil + } + out := []openai.Tool{} + for _, t := range list { + out = append(out, openai.Tool{ + Type: "function", + Function: openai.FunctionTool{ + Name: t.Name, + Description: t.Description, + Parameters: t.InputSchema, + }, + }) + } + return out +} diff --git a/server/internal/proxy/convert/convert.go b/server/internal/proxy/convert/convert.go new file mode 100644 index 0000000..1a323a6 --- /dev/null +++ b/server/internal/proxy/convert/convert.go @@ -0,0 +1,357 @@ +package convert + +import ( + "encoding/json" + + "openteam/server/internal/proxy/claude" + "openteam/server/internal/proxy/openai" +) + +// ClientProtocol mirrors the gateway's Protocol but kept here to avoid an +// import cycle with the proxy package. +type ClientProtocol string + +const ( + ClientOpenAIChat ClientProtocol = "openai-chat" + ClientOpenAIResponses ClientProtocol = "openai-responses" + ClientAnthropic ClientProtocol = "anthropic" +) + +// Request converts a request body from the client protocol to the channel +// provider's native format. +func Request(client ClientProtocol, provider string, body []byte, upstreamModel string) ([]byte, error) { + var canon *CanonicalRequest + var err error + switch client { + case ClientOpenAIChat: + canon, err = requestFromOpenAIChat(body) + case ClientOpenAIResponses: + canon, err = requestFromResponses(body) + case ClientAnthropic: + canon, err = requestFromClaude(body) + default: + return nil, jsonError("unsupported client protocol") + } + if err != nil { + return nil, err + } + canon.Model = upstreamModel + if canon.Model == "" { + canon.Model = "" + } + + switch provider { + case "anthropic": + return canonicalToClaude(canon) + case "openai", "compatible": + return canonicalToOpenAIChat(canon) + default: + return nil, jsonError("unsupported provider") + } +} + +// Response converts a non-stream upstream response into the client protocol. +func Response(client ClientProtocol, provider string, body []byte) ([]byte, error) { + if client == ClientOpenAIChat && provider == "openai" { + return body, nil // passthrough + } + if client == ClientOpenAIResponses && provider == "openai" { + return body, nil + } + if client == ClientAnthropic && provider == "anthropic" { + return body, nil + } + if provider == "anthropic" { + var mr claude.MessageResponse + if err := json.Unmarshal(body, &mr); err != nil { + return nil, err + } + switch client { + case ClientOpenAIChat: + return claudeResponseToChat(&mr) + case ClientOpenAIResponses: + return claudeResponseToResponses(&mr) + } + } + if provider == "openai" || provider == "compatible" { + var cc openai.ChatCompletion + if err := json.Unmarshal(body, &cc); err != nil { + return nil, err + } + switch client { + case ClientAnthropic: + return chatResponseToClaude(&cc) + } + } + return nil, jsonError("no conversion path") +} + +// canonicalToOpenAIChat re-emits the canonical request as OpenAI chat JSON. +func canonicalToOpenAIChat(c *CanonicalRequest) ([]byte, error) { + obj := map[string]any{ + "model": c.Model, + "messages": c.Messages, + } + if c.Stream { + obj["stream"] = true + } + if c.Temperature != nil { + obj["temperature"] = *c.Temperature + } + if c.TopP != nil { + obj["top_p"] = *c.TopP + } + if c.MaxTokens != nil { + obj["max_tokens"] = *c.MaxTokens + } + if len(c.Stop) > 0 { + obj["stop"] = c.Stop + } + if len(c.Tools) > 0 { + obj["tools"] = c.Tools + } + if len(c.ToolChoice) > 0 { + obj["tool_choice"] = json.RawMessage(c.ToolChoice) + } + if len(c.ResponseFormat) > 0 { + obj["response_format"] = json.RawMessage(c.ResponseFormat) + } + for k, v := range c.RawOpenAIExtras { + obj[k] = json.RawMessage(v) + } + return json.Marshal(obj) +} + +// canonicalToClaude emits the canonical request as an Anthropic Messages body. +func canonicalToClaude(c *CanonicalRequest) ([]byte, error) { + body, err := chatToClaudeMessages(c.Messages, c.System) + if err != nil { + return nil, err + } + maxTokens := 4096 + if c.MaxTokens != nil { + maxTokens = *c.MaxTokens + } + body["model"] = c.Model + body["max_tokens"] = maxTokens + if c.Stream { + body["stream"] = true + } + if c.Temperature != nil { + // Claude clamps temperature to [0,1]. + t := *c.Temperature + if t > 1 { + t = 1 + } + if t < 0 { + t = 0 + } + body["temperature"] = t + } + if c.TopP != nil { + body["top_p"] = *c.TopP + } + if len(c.Stop) > 0 { + body["stop_sequences"] = c.Stop + } + if tools := toolsToClaude(c.Tools); len(tools) > 0 { + body["tools"] = tools + } + if len(c.ToolChoice) > 0 { + var tc json.RawMessage + if err := json.Unmarshal(c.ToolChoice, &tc); err == nil { + body["tool_choice"] = tc + } + } + return json.Marshal(body) +} + +// claudeResponseToChat converts a Claude non-stream response to OpenAI chat. +func claudeResponseToChat(mr *claude.MessageResponse) ([]byte, error) { + var blocks []struct { + Type string `json:"type"` + Text string `json:"text"` + ID string `json:"id"` + Name string `json:"name"` + Input json.RawMessage `json:"input"` + } + _ = json.Unmarshal(mr.Content, &blocks) + + content := "" + var toolCalls []map[string]any + for _, b := range blocks { + switch b.Type { + case "text": + if content == "" { + content = b.Text + } else { + content += b.Text + } + case "tool_use": + input, _ := json.Marshal(b.Input) + toolCalls = append(toolCalls, map[string]any{ + "id": b.ID, + "type": "function", + "function": map[string]any{ + "name": b.Name, + "arguments": string(input), + }, + }) + } + } + msg := map[string]any{"role": "assistant", "content": content} + if len(toolCalls) > 0 { + msg["tool_calls"] = toolCalls + } + finish := mapClaudeStopReason(mr.StopReason) + choices := []any{map[string]any{"index": 0, "message": msg, "finish_reason": finish}} + resp := map[string]any{ + "id": "chatcmpl-" + mr.ID, + "object": "chat.completion", + "created": json.Number("0"), + "model": mr.Model, + "choices": choices, + } + if mr.Usage != nil { + resp["usage"] = map[string]any{ + "prompt_tokens": mr.Usage.InputTokens, + "completion_tokens": mr.Usage.OutputTokens, + "total_tokens": mr.Usage.InputTokens + mr.Usage.OutputTokens, + } + } + return json.Marshal(resp) +} + +// chatResponseToClaude converts an OpenAI non-stream response to Claude. +func chatResponseToClaude(cc *openai.ChatCompletion) ([]byte, error) { + content := []map[string]any{} + if len(cc.Choices) > 0 { + ch := cc.Choices[0] + if text := contentString(ch.Message.Content); text != "" { + content = append(content, map[string]any{"type": "text", "text": text}) + } + if len(ch.Message.ToolCalls) > 0 { + var calls []struct { + ID string `json:"id"` + Function struct { + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` + } `json:"function"` + } + _ = json.Unmarshal(ch.Message.ToolCalls, &calls) + for _, call := range calls { + var input map[string]any + _ = json.Unmarshal(call.Function.Arguments, &input) + content = append(content, map[string]any{ + "type": "tool_use", + "id": call.ID, + "name": call.Function.Name, + "input": input, + }) + } + } + } + resp := map[string]any{ + "id": mrID(cc.ID), + "type": "message", + "role": "assistant", + "model": cc.Model, + "content": content, + "stop_reason": mapChatStopReason(cc), + "usage": map[string]any{ + "input_tokens": usageInt64(cc, true), + "output_tokens": usageInt64(cc, false), + }, + } + return json.Marshal(resp) +} + +// claudeResponseToResponses converts a Claude non-stream response to a +// Responses API response. +func claudeResponseToResponses(mr *claude.MessageResponse) ([]byte, error) { + chatCC := &openai.ChatCompletion{ + ID: mr.ID, + Model: mr.Model, + Choices: []openai.ChatChoice{{FinishReason: mapClaudeStopReason(mr.StopReason)}}, + } + // Reuse the chat conversion then re-shape into responses items. + chatBody, err := claudeResponseToChat(mr) + if err == nil { + var cc openai.ChatCompletion + if json.Unmarshal(chatBody, &cc) == nil { + chatCC = &cc + } + } + items, status := chatToResponsesOutput(chatCC) + resp := map[string]any{ + "id": "resp_" + mr.ID, + "object": "response", + "created": json.Number("0"), + "model": mr.Model, + "status": status, + "output": items, + } + if mr.Usage != nil { + resp["usage"] = map[string]any{ + "input_tokens": mr.Usage.InputTokens, + "output_tokens": mr.Usage.OutputTokens, + "total_tokens": mr.Usage.InputTokens + mr.Usage.OutputTokens, + "input_tokens_details": map[string]any{ + "cached_tokens": mr.Usage.CacheReadInputTokens, + }, + } + } + return json.Marshal(resp) +} + +func mapClaudeStopReason(reason string) string { + switch reason { + case "end_turn": + return "stop" + case "max_tokens": + return "length" + case "stop_sequence": + return "stop" + case "tool_use": + return "tool_calls" + case "refusal": + return "content_filter" + default: + return "stop" + } +} + +func mapChatStopReason(cc *openai.ChatCompletion) string { + if len(cc.Choices) == 0 { + return "end_turn" + } + switch cc.Choices[0].FinishReason { + case "stop": + return "end_turn" + case "length": + return "max_tokens" + case "tool_calls": + return "tool_use" + case "content_filter": + return "refusal" + default: + return "end_turn" + } +} + +func usageInt64(cc *openai.ChatCompletion, input bool) int64 { + if cc.Usage == nil { + return 0 + } + if input { + return cc.Usage.PromptTokens + } + return cc.Usage.CompletionTokens +} + +func mrID(id string) string { + if id == "" { + return "msg_unknown" + } + return id +} diff --git a/server/internal/proxy/convert/responses.go b/server/internal/proxy/convert/responses.go new file mode 100644 index 0000000..c6af030 --- /dev/null +++ b/server/internal/proxy/convert/responses.go @@ -0,0 +1,161 @@ +package convert + +import ( + "encoding/json" + "errors" + + "openteam/server/internal/proxy/openai" +) + +// responsesInputToChat converts Responses API input items into chat messages. +func responsesInputToChat(raw json.RawMessage) ([]openai.ChatMessage, error) { + var msgs []openai.ChatMessage + + // `input` may be a plain string. + var s string + if json.Unmarshal(raw, &s) == nil { + content, _ := json.Marshal(s) + msgs = append(msgs, openai.ChatMessage{Role: "user", Content: content}) + return msgs, nil + } + + // Or an array of content parts (text/image). + var parts []map[string]json.RawMessage + if json.Unmarshal(raw, &parts) == nil { + chat, err := contentPartsToChat(parts) + if err == nil { + return chat, nil + } + } + + var items []map[string]json.RawMessage + if err := json.Unmarshal(raw, &items); err != nil { + return nil, err + } + for _, item := range items { + var typ string + _ = json.Unmarshal(item["type"], &typ) + switch typ { + case "message": + var role string + _ = json.Unmarshal(item["role"], &role) + content, err := contentPartsToText(item["content"]) + if err != nil { + continue + } + msgs = append(msgs, openai.ChatMessage{Role: role, Content: content}) + case "function_call": + tc := map[string]any{ + "id": rawString(item["call_id"]), + "type": "function", + "function": map[string]any{ + "name": rawString(item["name"]), + "arguments": rawString(item["arguments"]), + }, + } + tcArr, _ := json.Marshal([]any{tc}) + msgs = append(msgs, openai.ChatMessage{Role: "assistant", ToolCalls: tcArr}) + case "function_call_output": + content, _ := json.Marshal(rawString(item["output"])) + msgs = append(msgs, openai.ChatMessage{Role: "tool", ToolCallID: rawString(item["call_id"]), Content: content}) + case "reasoning", "computer_call", "web_search_call": + // Not representable in chat; drop. + } + } + if len(msgs) == 0 { + content, _ := json.Marshal("") + msgs = append(msgs, openai.ChatMessage{Role: "user", Content: content}) + } + return msgs, nil +} + +// contentPartsToText flattens an array of content parts into a string. +func contentPartsToText(raw json.RawMessage) (json.RawMessage, error) { + var parts []map[string]json.RawMessage + if err := json.Unmarshal(raw, &parts); err != nil { + var s string + if err := json.Unmarshal(raw, &s); err == nil { + c, _ := json.Marshal(s) + return c, nil + } + return nil, err + } + var buf string + for _, p := range parts { + var typ string + _ = json.Unmarshal(p["type"], &typ) + switch typ { + case "input_text", "output_text", "text": + buf += rawString(p["text"]) + } + } + c, _ := json.Marshal(buf) + return c, nil +} + +func contentPartsToChat(parts []map[string]json.RawMessage) ([]openai.ChatMessage, error) { + out := []map[string]any{} + for _, p := range parts { + var typ string + _ = json.Unmarshal(p["type"], &typ) + switch typ { + case "input_text", "output_text", "text": + out = append(out, map[string]any{"type": "text", "text": rawString(p["text"])}) + case "input_image": + out = append(out, map[string]any{ + "type": "image_url", + "image_url": map[string]any{ + "url": "data:" + rawString(p["media_type"]) + ";base64," + rawString(p["image_url"]), + }, + }) + } + } + if len(out) == 0 { + return nil, jsonError("empty content") + } + arr, _ := json.Marshal(out) + msg := openai.ChatMessage{Role: "user", Content: arr} + return []openai.ChatMessage{msg}, nil +} + +// chatToResponsesOutput builds Responses output items from a chat completion. +func chatToResponsesOutput(cc *openai.ChatCompletion) ([]map[string]any, string) { + var items []map[string]any + var status = "completed" + if len(cc.Choices) == 0 { + return items, status + } + ch := cc.Choices[0] + if ch.FinishReason == "length" { + status = "incomplete" + } + content := []map[string]any{} + if text := contentString(ch.Message.Content); text != "" { + content = append(content, map[string]any{"type": "output_text", "text": text}) + } + items = append(items, map[string]any{"type": "message", "role": "assistant", "content": content}) + if len(ch.Message.ToolCalls) > 0 { + var calls []struct { + ID string `json:"id"` + Function struct { + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` + } `json:"function"` + } + if json.Unmarshal(ch.Message.ToolCalls, &calls) == nil { + for _, call := range calls { + items = append(items, map[string]any{ + "type": "function_call", + "call_id": call.ID, + "name": call.Function.Name, + "arguments": rawString(call.Function.Arguments), + }) + } + } + } + return items, status +} + +func jsonError(msg string) error { + return errors.New(msg) +} diff --git a/server/internal/proxy/convert/stream.go b/server/internal/proxy/convert/stream.go new file mode 100644 index 0000000..09dcef1 --- /dev/null +++ b/server/internal/proxy/convert/stream.go @@ -0,0 +1,416 @@ +package convert + +import ( + "encoding/json" + + "openteam/server/internal/proxy/claude" + "openteam/server/internal/proxy/openai" + "openteam/server/internal/proxy/stream" +) + +// Usage is the streaming token usage snapshot. +type Usage struct { + Input int64 + Output int64 + CacheRead int64 + CacheCreation int64 +} + +// Translator converts SSE events from an upstream stream into client frames. +type Translator interface { + // Feed handles one upstream SSE event, returning client frames to write. + Feed(ev stream.SSEEvent) ([]stream.SSEEvent, error) + // Finish is called at end-of-stream, returning final frames. + Finish() ([]stream.SSEEvent, error) + // Usage returns the latest known usage. + Usage() *Usage +} + +// claudeToChatTranslator converts a Claude stream to OpenAI chat chunks. +type claudeToChatTranslator struct { + model string + usage *Usage + started bool + finishSent bool + toolCallIndex int + toolCallID string + toolCallName string +} + +func (t *claudeToChatTranslator) Feed(ev stream.SSEEvent) ([]stream.SSEEvent, error) { + if ev.Done { + return nil, nil + } + var e claude.StreamEvent + if err := json.Unmarshal([]byte(ev.Data), &e); err != nil { + return nil, nil + } + var out []stream.SSEEvent + + switch e.Type { + case "message_start": + if e.Message != nil { + var msg struct { + Model string `json:"model"` + } + _ = json.Unmarshal(e.Message, &msg) + t.model = msg.Model + } + chunk, _ := json.Marshal(openai.ChatChunk{ + ID: "chatcmpl-stream", Object: "chat.completion.chunk", Model: t.model, + Choices: []openai.ChatChunkChoice{{Index: 0, Delta: openai.ChatDelta{Role: "assistant"}}}, + }) + out = append(out, stream.SSEEvent{Data: string(chunk)}) + t.started = true + + case "content_block_start": + var cb struct { + Index int `json:"index"` + Block json.RawMessage `json:"content_block"` + } + _ = json.Unmarshal([]byte(ev.Data), &cb) + var block struct { + Type string `json:"type"` + ID string `json:"id"` + Name string `json:"name"` + } + _ = json.Unmarshal(cb.Block, &block) + if block.Type == "tool_use" { + t.toolCallIndex = cb.Index + t.toolCallID = block.ID + t.toolCallName = block.Name + tc, _ := json.Marshal([]map[string]any{{ + "index": cb.Index, "id": block.ID, "type": "function", + "function": map[string]any{"name": block.Name, "arguments": ""}, + }}) + chunk, _ := json.Marshal(openai.ChatChunk{ + ID: "chatcmpl-stream", Object: "chat.completion.chunk", Model: t.model, + Choices: []openai.ChatChunkChoice{{Index: 0, Delta: openai.ChatDelta{ToolCalls: tc}}}, + }) + out = append(out, stream.SSEEvent{Data: string(chunk)}) + } + + case "content_block_delta": + var d struct { + Index int `json:"index"` + Delta json.RawMessage `json:"delta"` + } + _ = json.Unmarshal([]byte(ev.Data), &d) + var delta struct { + Type string `json:"type"` + Text string `json:"text"` + PartialJSON string `json:"partial_json"` + } + _ = json.Unmarshal(d.Delta, &delta) + if delta.Type == "text_delta" && delta.Text != "" { + chunk, _ := json.Marshal(openai.ChatChunk{ + ID: "chatcmpl-stream", Object: "chat.completion.chunk", Model: t.model, + Choices: []openai.ChatChunkChoice{{Index: 0, Delta: openai.ChatDelta{Content: delta.Text}}}, + }) + out = append(out, stream.SSEEvent{Data: string(chunk)}) + } else if delta.Type == "input_json_delta" && delta.PartialJSON != "" { + tc, _ := json.Marshal([]map[string]any{{ + "index": d.Index, "function": map[string]any{"arguments": delta.PartialJSON}, + }}) + chunk, _ := json.Marshal(openai.ChatChunk{ + ID: "chatcmpl-stream", Object: "chat.completion.chunk", Model: t.model, + Choices: []openai.ChatChunkChoice{{Index: 0, Delta: openai.ChatDelta{ToolCalls: tc}}}, + }) + out = append(out, stream.SSEEvent{Data: string(chunk)}) + } + + case "message_delta": + var d struct { + Delta json.RawMessage `json:"delta"` + Usage json.RawMessage `json:"usage"` + } + _ = json.Unmarshal([]byte(ev.Data), &d) + if len(d.Usage) > 0 { + var u claude.Usage + if json.Unmarshal(d.Usage, &u) == nil { + t.usage = &Usage{ + Input: u.InputTokens, Output: u.OutputTokens, + CacheRead: u.CacheReadInputTokens, CacheCreation: u.CacheCreationInputTokens, + } + } + } + if len(d.Delta) > 0 && !t.finishSent { + var delta struct { + StopReason string `json:"stop_reason"` + } + _ = json.Unmarshal(d.Delta, &delta) + if delta.StopReason != "" { + reason := mapClaudeStopReason(delta.StopReason) + chunk, _ := json.Marshal(openai.ChatChunk{ + ID: "chatcmpl-stream", Object: "chat.completion.chunk", Model: t.model, + Choices: []openai.ChatChunkChoice{{Index: 0, Delta: openai.ChatDelta{}, FinishReason: &reason}}, + }) + out = append(out, stream.SSEEvent{Data: string(chunk)}) + t.finishSent = true + } + } + } + return out, nil +} + +func (t *claudeToChatTranslator) Finish() ([]stream.SSEEvent, error) { + if !t.finishSent { + reason := "stop" + chunk, _ := json.Marshal(openai.ChatChunk{ + ID: "chatcmpl-stream", Object: "chat.completion.chunk", Model: t.model, + Choices: []openai.ChatChunkChoice{{Index: 0, Delta: openai.ChatDelta{}, FinishReason: &reason}}, + }) + t.finishSent = true + return []stream.SSEEvent{{Data: string(chunk)}, {Data: "[DONE]"}}, nil + } + return []stream.SSEEvent{{Data: "[DONE]"}}, nil +} + +func (t *claudeToChatTranslator) Usage() *Usage { return t.usage } + +// chatToClaudeTranslator converts an OpenAI chat stream to Claude events. +type chatToClaudeTranslator struct { + usage *Usage + started bool + openBlock bool + blockType string + toolIndex int + finishSent bool +} + +func (t *chatToClaudeTranslator) Feed(ev stream.SSEEvent) ([]stream.SSEEvent, error) { + if ev.Done { + return nil, nil + } + var chunk openai.ChatChunk + if err := json.Unmarshal([]byte(ev.Data), &chunk); err != nil { + return nil, nil + } + if chunk.Usage != nil { + t.usage = &Usage{Input: chunk.Usage.PromptTokens, Output: chunk.Usage.CompletionTokens} + } + var out []stream.SSEEvent + if len(chunk.Choices) == 0 { + return out, nil + } + ch := chunk.Choices[0] + + if !t.started { + msg, _ := json.Marshal(map[string]any{ + "id": "msg_stream", "type": "message", "role": "assistant", + "model": chunk.Model, "content": []any{}, + }) + start, _ := json.Marshal(map[string]any{"type": "message_start", "message": json.RawMessage(msg)}) + out = append(out, stream.SSEEvent{Data: string(start)}) + t.started = true + } + + if ch.Delta.Content != "" { + if !t.openBlock || t.blockType != "text" { + start, _ := json.Marshal(map[string]any{ + "type": "content_block_start", "index": 0, + "content_block": map[string]any{"type": "text", "text": ""}, + }) + out = append(out, stream.SSEEvent{Data: string(start)}) + t.openBlock = true + t.blockType = "text" + t.toolIndex = 0 + } + delta, _ := json.Marshal(map[string]any{ + "type": "content_block_delta", "index": 0, + "delta": map[string]any{"type": "text_delta", "text": ch.Delta.Content}, + }) + out = append(out, stream.SSEEvent{Data: string(delta)}) + } + + if len(ch.Delta.ToolCalls) > 0 { + var calls []struct { + Index *int `json:"index"` + ID string `json:"id"` + Function struct { + Name string `json:"name"` + Arguments string `json:"arguments"` + } `json:"function"` + } + _ = json.Unmarshal(ch.Delta.ToolCalls, &calls) + for _, call := range calls { + idx := 0 + if call.Index != nil { + idx = *call.Index + } + if !t.openBlock || t.blockType != "tool_use" || idx != t.toolIndex { + start, _ := json.Marshal(map[string]any{ + "type": "content_block_start", "index": idx, + "content_block": map[string]any{ + "type": "tool_use", "id": call.ID, "name": call.Function.Name, "input": map[string]any{}, + }, + }) + out = append(out, stream.SSEEvent{Data: string(start)}) + t.openBlock = true + t.blockType = "tool_use" + t.toolIndex = idx + } + if call.Function.Arguments != "" { + delta, _ := json.Marshal(map[string]any{ + "type": "content_block_delta", "index": idx, + "delta": map[string]any{"type": "input_json_delta", "partial_json": call.Function.Arguments}, + }) + out = append(out, stream.SSEEvent{Data: string(delta)}) + } + } + } + + if ch.FinishReason != nil && !t.finishSent { + reason := mapChatStopReasonToClaude(*ch.FinishReason) + md, _ := json.Marshal(map[string]any{ + "type": "message_delta", + "delta": map[string]any{"stop_reason": reason, "stop_sequence": nil}, + }) + out = append(out, stream.SSEEvent{Data: string(md)}) + t.finishSent = true + } + return out, nil +} + +func (t *chatToClaudeTranslator) Finish() ([]stream.SSEEvent, error) { + if !t.finishSent { + md, _ := json.Marshal(map[string]any{ + "type": "message_delta", + "delta": map[string]any{"stop_reason": "end_turn", "stop_sequence": nil}, + }) + t.finishSent = true + return []stream.SSEEvent{{Data: string(md)}, {Data: "{\"type\":\"message_stop\"}"}}, nil + } + return []stream.SSEEvent{{Data: "{\"type\":\"message_stop\"}"}}, nil +} + +func (t *chatToClaudeTranslator) Usage() *Usage { return t.usage } + +// claudeToResponsesTranslator converts a Claude stream to Responses events. +type claudeToResponsesTranslator struct { + usage *Usage + model string + completed bool +} + +func (t *claudeToResponsesTranslator) Feed(ev stream.SSEEvent) ([]stream.SSEEvent, error) { + if ev.Done { + return nil, nil + } + var e claude.StreamEvent + if err := json.Unmarshal([]byte(ev.Data), &e); err != nil { + return nil, nil + } + var out []stream.SSEEvent + switch e.Type { + case "message_start": + if e.Message != nil { + var msg struct { + Model string `json:"model"` + } + _ = json.Unmarshal(e.Message, &msg) + t.model = msg.Model + } + created, _ := json.Marshal(map[string]any{ + "type": "response.created", + "response": map[string]any{"id": "resp_stream", "object": "response", "status": "in_progress", "model": t.model, "output": []any{}}, + }) + out = append(out, stream.SSEEvent{Data: string(created)}) + case "content_block_delta": + var d struct { + Index int `json:"index"` + Delta json.RawMessage `json:"delta"` + } + _ = json.Unmarshal([]byte(ev.Data), &d) + var delta struct { + Type string `json:"type"` + Text string `json:"text"` + } + _ = json.Unmarshal(d.Delta, &delta) + if delta.Type == "text_delta" && delta.Text != "" { + item, _ := json.Marshal(map[string]any{ + "type": "response.output_text.delta", "item_id": "msg_stream", "output_index": 0, + "delta": delta.Text, + }) + out = append(out, stream.SSEEvent{Data: string(item)}) + } + case "message_delta": + var d struct { + Usage json.RawMessage `json:"usage"` + } + _ = json.Unmarshal([]byte(ev.Data), &d) + if len(d.Usage) > 0 { + var u claude.Usage + if json.Unmarshal(d.Usage, &u) == nil { + t.usage = &Usage{ + Input: u.InputTokens, Output: u.OutputTokens, + CacheRead: u.CacheReadInputTokens, CacheCreation: u.CacheCreationInputTokens, + } + } + } + if !t.completed { + item, _ := json.Marshal(map[string]any{ + "type": "response.completed", + "response": map[string]any{ + "id": "resp_stream", "object": "response", "status": "completed", "model": t.model, + "output": []map[string]any{{"type": "message", "role": "assistant", "content": []any{}}}, + }, + }) + out = append(out, stream.SSEEvent{Data: string(item)}) + t.completed = true + } + } + return out, nil +} + +func (t *claudeToResponsesTranslator) Finish() ([]stream.SSEEvent, error) { + if !t.completed { + item, _ := json.Marshal(map[string]any{ + "type": "response.completed", + "response": map[string]any{ + "id": "resp_stream", "object": "response", "status": "completed", "model": t.model, + "output": []map[string]any{{"type": "message", "role": "assistant", "content": []any{}}}, + }, + }) + t.completed = true + return []stream.SSEEvent{{Data: string(item)}, {Data: "[DONE]"}}, nil + } + return []stream.SSEEvent{{Data: "[DONE]"}}, nil +} + +func (t *claudeToResponsesTranslator) Usage() *Usage { return t.usage } + +func mapChatStopReasonToClaude(reason string) string { + switch reason { + case "stop": + return "end_turn" + case "length": + return "max_tokens" + case "tool_calls": + return "tool_use" + case "content_filter": + return "refusal" + default: + return "end_turn" + } +} + +// NewTranslator returns the stream translator for a client protocol + +// upstream provider pair, or nil for passthrough (no translation needed). +func NewTranslator(client ClientProtocol, provider string) Translator { + switch client { + case ClientOpenAIChat: + if provider == "anthropic" { + return &claudeToChatTranslator{} + } + case ClientOpenAIResponses: + if provider == "anthropic" { + return &claudeToResponsesTranslator{} + } + case ClientAnthropic: + if provider == "openai" || provider == "compatible" { + return &chatToClaudeTranslator{} + } + } + return nil +} diff --git a/server/internal/proxy/convert_wrappers.go b/server/internal/proxy/convert_wrappers.go new file mode 100644 index 0000000..794d129 --- /dev/null +++ b/server/internal/proxy/convert_wrappers.go @@ -0,0 +1,17 @@ +package proxy + +import ( + "openteam/server/internal/proxy/convert" +) + +// convertRequest translates a request body from the client protocol to the +// channel provider's native format. +func convertRequest(route Route, provider string, body []byte, upstreamModel string) ([]byte, error) { + return convert.Request(clientProto(route), provider, body, upstreamModel) +} + +// convertResponse translates a non-stream upstream response back to the +// client protocol. +func convertResponse(route Route, provider string, body []byte) ([]byte, error) { + return convert.Response(clientProto(route), provider, body) +} diff --git a/server/internal/proxy/gateway.go b/server/internal/proxy/gateway.go new file mode 100644 index 0000000..5be34a2 --- /dev/null +++ b/server/internal/proxy/gateway.go @@ -0,0 +1,414 @@ +package proxy + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strconv" + "strings" + "time" + + "github.com/gin-gonic/gin" + "github.com/google/uuid" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/apikey" + "openteam/server/internal/billing" + "openteam/server/internal/channel" + "openteam/server/internal/config" + "openteam/server/internal/pkg/ratelimit" + "openteam/server/internal/usage" + "openteam/server/internal/store" +) + +// Protocol identifies the client-facing API protocol. +type Protocol string + +const ( + ProtocolOpenAIChat Protocol = "openai-chat" + ProtocolOpenAIResponses Protocol = "openai-responses" + ProtocolAnthropic Protocol = "anthropic" +) + +// Route describes one proxied endpoint. +type Route struct { + Protocol Protocol + UpstreamPath string // suffix after base URL, e.g. /v1/chat/completions + NativeProvider string // provider type that matches this protocol ("openai" | "anthropic") +} + +var Routes = []Route{ + {Protocol: ProtocolOpenAIChat, UpstreamPath: "/v1/chat/completions", NativeProvider: "openai"}, + {Protocol: ProtocolOpenAIResponses, UpstreamPath: "/v1/responses", NativeProvider: "openai"}, + {Protocol: ProtocolAnthropic, UpstreamPath: "/v1/messages", NativeProvider: "anthropic"}, +} + +const maxBodyBytes = 16 << 20 // 16 MiB + +type Gateway struct { + db *gorm.DB + cfg *config.Config + log *zap.Logger + channel *channel.Service + billing *billing.Service + usage *usage.Service + apiKeys *apikey.Service + limiter *ratelimit.Limiter + client *http.Client +} + +func NewGateway(db *gorm.DB, cfg *config.Config, log *zap.Logger, + ch *channel.Service, bill *billing.Service, use *usage.Service, ak *apikey.Service) *Gateway { + return &Gateway{ + db: db, cfg: cfg, log: log, + channel: ch, billing: bill, usage: use, apiKeys: ak, + limiter: ratelimit.New(float64(cfg.RateLimit.RequestsPerMin)/60.0, cfg.RateLimit.Burst), + client: &http.Client{ + // Transport-level timeout; stream reads rely on the context so a + // connected-but-silent upstream is still bounded. + Timeout: time.Duration(cfg.Proxy.DefaultTimeoutMs) * time.Millisecond, + }, + } +} + +// authContext carries the authenticated key + user through a proxy request. +type authContext struct { + key *store.ApiKey + user *store.User +} + +// Handle builds a gin handler for a route. +func (g *Gateway) Handle(route Route) gin.HandlerFunc { + return func(c *gin.Context) { + g.proxy(c, route) + } +} + +func (g *Gateway) proxy(c *gin.Context, route Route) { + start := time.Now() + reqID := uuid.NewString() + c.Header("X-Request-Id", reqID) + + auth, err := g.authenticate(c) + if err != nil { + g.writeProxyError(c, route.Protocol, http.StatusUnauthorized, "invalid API key") + return + } + + body, err := io.ReadAll(io.LimitReader(c.Request.Body, maxBodyBytes)) + if err != nil { + g.writeProxyError(c, route.Protocol, http.StatusBadRequest, "failed to read request body") + return + } + if len(body) == 0 { + g.writeProxyError(c, route.Protocol, http.StatusBadRequest, "empty request body") + return + } + + var meta struct { + Model string `json:"model"` + Stream bool `json:"stream"` + } + if err := json.Unmarshal(body, &meta); err != nil { + g.writeProxyError(c, route.Protocol, http.StatusBadRequest, "invalid JSON body: "+err.Error()) + return + } + if meta.Model == "" { + g.writeProxyError(c, route.Protocol, http.StatusBadRequest, "missing model field") + return + } + + // Per-key model whitelist. + if len(auth.key.AllowedModels) > 0 && !contains(auth.key.AllowedModels, meta.Model) { + g.writeProxyError(c, route.Protocol, http.StatusForbidden, + "model not allowed for this API key: "+meta.Model) + return + } + + // Rate limit (user + key). + if !g.limiter.Allow(fmt.Sprintf("u:%d", auth.user.ID)) { + g.writeProxyError(c, route.Protocol, http.StatusTooManyRequests, "rate limit exceeded") + return + } + if !g.limiter.Allow(fmt.Sprintf("k:%d", auth.key.ID)) { + g.writeProxyError(c, route.Protocol, http.StatusTooManyRequests, "key rate limit exceeded") + return + } + + // Resolve model. + model, err := g.channel.ResolveModel(meta.Model) + if err != nil { + g.writeProxyError(c, route.Protocol, http.StatusNotFound, err.Error()) + return + } + + // Channel attempts with failover: on transport/5xx failures the request is + // retried against another channel bound to the same model. + exclude := map[int64]bool{} + attempts := 1 + g.cfg.Proxy.MaxRetries + if attempts < 1 { + attempts = 1 + } + var lastCh *store.Channel + formatBlocked := false + attempted := false + for attempt := 0; attempt < attempts; attempt++ { + ch, binding, err := g.channel.SelectChannel(model.ID, exclude) + if err != nil { + // Every bound channel was skipped for format reasons: say so + // clearly instead of reporting a generic upstream failure. + if formatBlocked && !attempted { + g.writeProxyError(c, route.Protocol, http.StatusBadRequest, + "no channel supports the "+string(route.Protocol)+" API format") + return + } + if attempt == 0 { + g.writeProxyError(c, route.Protocol, http.StatusServiceUnavailable, "no available channel for model") + } else { + g.recordError(reqID, auth, model, lastCh, start, "upstream_error") + g.writeProxyError(c, route.Protocol, http.StatusBadGateway, "all upstream channels failed") + } + return + } + exclude[ch.ID] = true + lastCh = ch + + // Skip channels that cannot serve this protocol at all (no native + // support and no conversion path). + if !channelServesFormat(ch, route) { + formatBlocked = true + g.log.Debug("channel cannot serve protocol", + zap.String("protocol", string(route.Protocol)), zap.Int64("channel_id", ch.ID)) + continue + } + + // Optional balance pre-check with an estimate. + if g.cfg.Proxy.BillingExactBalance { + estIn := int64(len(body) / 4) + estOut := int64(512) + if meta.Stream { + estOut = int64(g.cfg.Proxy.DefaultMaxTokens) + } + est, cerr := g.billing.EstimateCost(model.ID, estIn, estOut, 0) + if cerr == nil { + if berr := g.billing.CheckBalance(auth.user.ID, est); berr != nil { + g.recordError(reqID, auth, model, ch, start, "insufficient_balance") + g.writeProxyError(c, route.Protocol, http.StatusPaymentRequired, "insufficient balance") + return + } + } + } + + upstreamKey, err := g.channel.DecryptKey(ch.APIKeyEnc) + if err != nil { + g.log.Error("decrypt channel key", zap.Error(err), zap.Int64("channel_id", ch.ID)) + g.writeProxyError(c, route.Protocol, http.StatusBadGateway, "channel key unavailable") + return + } + + release, err := g.channel.Acquire(ch.ID) + if err != nil { + g.writeProxyError(c, route.Protocol, http.StatusServiceUnavailable, "channel unavailable") + return + } + attempted = true + retryable := g.forward(c, route, meta.Model, ch, binding.UpstreamModel, upstreamKey, body, start, reqID, auth, model) + release() + if !retryable { + return + } + g.log.Warn("upstream failed, retrying on another channel", + zap.Int64("model_id", model.ID), zap.Int64("channel_id", ch.ID), zap.Int("attempt", attempt+1)) + } + if formatBlocked && !attempted { + g.writeProxyError(c, route.Protocol, http.StatusBadRequest, + "no channel supports the "+string(route.Protocol)+" API format") + return + } + g.recordError(reqID, auth, model, lastCh, start, "upstream_error") + g.writeProxyError(c, route.Protocol, http.StatusBadGateway, "all upstream channels failed") +} + +// forward decides passthrough vs conversion and calls the upstream. +// It returns true when the failure is retryable on another channel. +func (g *Gateway) forward(c *gin.Context, route Route, clientModel string, + ch *store.Channel, upstreamModel, upstreamKey string, body []byte, + start time.Time, reqID string, auth *authContext, model *store.Model) bool { + + upstreamBody, converted := g.prepareUpstreamBody(route, ch, body, upstreamModel) + if converted { + c.Header("x-converted", "true") + } + + // Build upstream request bound to the client context so disconnects cancel it. + ctx := c.Request.Context() + upstreamURL := strings.TrimSuffix(ch.BaseURL, "/") + upstreamPath(route, ch) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamURL, bytes.NewReader(upstreamBody)) + if err != nil { + g.writeProxyError(c, route.Protocol, http.StatusBadGateway, "failed to build upstream request") + return false + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+upstreamKey) + req.Header.Set("Accept", "application/json") + req.Header.Set("X-Request-Id", reqID) + // Explicitly drop hop-by-hop / auth-ish headers we don't want forwarded. + copyProxyHeaders(c, req) + + resp, err := g.client.Do(req) + if err != nil { + // Client disconnect vs upstream failure. + if errors.Is(ctx.Err(), context.Canceled) { + g.recordCanceled(reqID, auth, model, ch, start) + return false + } + g.log.Warn("upstream request failed", zap.Error(err), zap.Int64("channel_id", ch.ID)) + return true + } + defer resp.Body.Close() + + if resp.StatusCode >= 500 { + errBody, _ := io.ReadAll(resp.Body) + g.log.Warn("upstream returned 5xx", zap.Int("status", resp.StatusCode), + zap.Int64("channel_id", ch.ID), zap.String("body", truncateText(string(errBody), 512))) + return true + } + if resp.StatusCode >= 400 { + errBody, _ := io.ReadAll(resp.Body) + g.recordError(reqID, auth, model, ch, start, "upstream_"+strconv.Itoa(resp.StatusCode)) + g.writeUpstreamError(c, route.Protocol, resp.StatusCode, errBody) + return false + } + + streaming := bodyStreamFlag(body, route) + if streaming { + g.streamResponse(c, route, resp, start, reqID, auth, model, ch) + } else { + g.plainResponse(c, route, resp, start, reqID, auth, model, ch) + } + return false +} + +// prepareUpstreamBody rewrites the model name, or converts the body when the +// channel does not serve the client protocol natively. Returns the payload and +// whether any conversion happened. +func (g *Gateway) prepareUpstreamBody(route Route, ch *store.Channel, body []byte, upstreamModel string) ([]byte, bool) { + if ch.SupportsFormat(string(route.Protocol)) { + if upstreamModel == "" || sameModel(body, upstreamModel) { + return body, false + } + rewritten, err := setModelField(body, upstreamModel) + if err != nil { + return body, false + } + return rewritten, true + } + converted, err := convertRequest(route, ch.Provider, body, upstreamModel) + if err != nil { + g.log.Warn("request conversion failed, falling back to passthrough", + zap.Error(err), zap.String("route", string(route.Protocol)), zap.String("provider", ch.Provider)) + return body, false + } + return converted, true +} + +// streamResponse forwards an SSE stream to the client while extracting usage. +func (g *Gateway) streamResponse(c *gin.Context, route Route, resp *http.Response, + start time.Time, reqID string, auth *authContext, model *store.Model, ch *store.Channel) { + + c.Header("Content-Type", "text/event-stream") + c.Header("Cache-Control", "no-cache") + c.Header("Connection", "keep-alive") + + if upstreamProvider(route, ch) == "anthropic" { + g.streamAnthropic(c, route, resp, start, reqID, auth, model, ch) + return + } + g.streamOpenAI(c, route, resp, start, reqID, auth, model, ch) +} + +// plainResponse buffers a non-stream upstream response and returns it. +func (g *Gateway) plainResponse(c *gin.Context, route Route, resp *http.Response, + start time.Time, reqID string, auth *authContext, model *store.Model, ch *store.Channel) { + + raw, err := io.ReadAll(resp.Body) + if err != nil { + g.log.Warn("read upstream body", zap.Error(err)) + g.writeProxyError(c, route.Protocol, http.StatusBadGateway, "failed to read upstream response") + return + } + + out := raw + usageInfo := parseUsageForProtocol(route.Protocol, raw) + if !ch.SupportsFormat(string(route.Protocol)) { + converted, cerr := convertResponse(route, ch.Provider, raw) + if cerr == nil { + out = converted + usageInfo = parseUsageForProtocol(route.Protocol, out) + c.Header("x-converted", "true") + } else { + g.log.Warn("response conversion failed, forwarding raw", + zap.Error(cerr), zap.String("protocol", string(route.Protocol))) + } + } + c.Data(http.StatusOK, "application/json", out) + g.afterComplete(start, reqID, auth, model, ch, usageInfo, "success", "") +} + +// afterComplete performs billing and usage accounting for a finished request. +func (g *Gateway) afterComplete(start time.Time, reqID string, auth *authContext, + model *store.Model, ch *store.Channel, usageInfo *tokenUsage, status, errCode string) { + + latency := int(time.Since(start).Milliseconds()) + if usageInfo == nil { + usageInfo = &tokenUsage{} + } + cost := billing.CostFromPrices(usageInfo.input, usageInfo.output, usageInfo.cacheRead, + billing.PriceSnapshot{ + InputPrice: model.InputPrice, OutputPrice: model.OutputPrice, CacheReadPrice: model.CacheReadPrice, + }) + + go func() { + // Asynchronous: deduct balance first, then record usage. + if cost.IsPositive() { + if _, err := g.billing.Deduct(auth.user.ID, cost, "usage", reqID); err != nil { + g.log.Warn("deduct balance failed", zap.Error(err), + zap.Int64("user_id", auth.user.ID), zap.String("request_id", reqID)) + } + } + g.usage.Record(usage.Record{ + RequestID: reqID, + UserID: auth.user.ID, + KeyID: auth.key.ID, + ChannelID: ch.ID, + ModelID: model.ID, + ModelName: model.Name, + InputTokens: usageInfo.input, + OutputTokens: usageInfo.output, + CacheReadTokens: usageInfo.cacheRead, + CacheCreationTokens: usageInfo.cacheCreation, + InputPrice: model.InputPrice, + OutputPrice: model.OutputPrice, + CacheReadPrice: model.CacheReadPrice, + Cost: cost, + LatencyMs: latency, + Status: status, + ErrorCode: errCode, + }) + g.db.Model(&store.ApiKey{}).Where("id = ?", auth.key.ID). + Update("last_used_at", time.Now()) + }() +} + +func (g *Gateway) recordError(reqID string, auth *authContext, model *store.Model, ch *store.Channel, start time.Time, errCode string) { + g.afterComplete(start, reqID, auth, model, ch, &tokenUsage{}, "error", errCode) +} + +func (g *Gateway) recordCanceled(reqID string, auth *authContext, model *store.Model, ch *store.Channel, start time.Time) { + g.afterComplete(start, reqID, auth, model, ch, &tokenUsage{}, "canceled", "client_disconnect") +} diff --git a/server/internal/proxy/helpers.go b/server/internal/proxy/helpers.go new file mode 100644 index 0000000..cb4280b --- /dev/null +++ b/server/internal/proxy/helpers.go @@ -0,0 +1,248 @@ +package proxy + +import ( + "encoding/json" + "net/http" + "strings" + "time" + + "github.com/gin-gonic/gin" + + "openteam/server/internal/pkg/crypto" + "openteam/server/internal/store" +) + +// tokenUsage is the normalized usage extracted from any protocol. +type tokenUsage struct { + input int64 + output int64 + cacheRead int64 + cacheCreation int64 +} + +// respStreamEvent is the streaming shape of the OpenAI Responses API used for +// usage sniffing on passthrough responses streams. +type respStreamEvent struct { + Type string `json:"type"` + Response *struct { + Usage *struct { + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + } `json:"usage"` + } `json:"response"` +} + +// authenticate resolves the Bearer API key to a key + user. +func (g *Gateway) authenticate(c *gin.Context) (*authContext, error) { + auth := c.GetHeader("Authorization") + token := "" + if strings.HasPrefix(auth, "Bearer ") { + token = strings.TrimPrefix(auth, "Bearer ") + } else if strings.HasPrefix(auth, "sk-") { + // Some clients send the raw key without the Bearer scheme. + token = auth + } else { + return nil, errUnauthorized + } + + var key store.ApiKey + if err := g.db.Where("key_hash = ?", crypto.HashSHA256(token)).First(&key).Error; err != nil { + return nil, errUnauthorized + } + if key.Status != "active" { + return nil, errUnauthorized + } + if key.ExpiresAt != nil && key.ExpiresAt.Before(now()) { + return nil, errUnauthorized + } + var user store.User + if err := g.db.First(&user, key.UserID).Error; err != nil { + return nil, errUnauthorized + } + if user.Status != "active" { + return nil, errUnauthorized + } + return &authContext{key: &key, user: &user}, nil +} + +var errUnauthorized = &unauthorizedError{} + +type unauthorizedError struct{} + +func (*unauthorizedError) Error() string { return "invalid API key" } + +// writeProxyError writes a gateway-generated error in the client's protocol. +func (g *Gateway) writeProxyError(c *gin.Context, proto Protocol, status int, message string) { + c.Header("Content-Type", "application/json") + switch proto { + case ProtocolAnthropic: + c.AbortWithStatusJSON(status, map[string]any{ + "type": "error", + "error": map[string]any{"type": statusType(status), "message": message}, + }) + default: + c.AbortWithStatusJSON(status, map[string]any{ + "error": map[string]any{"message": message, "type": "gateway_error", "code": "gateway_error"}, + }) + } +} + +// writeUpstreamError maps an upstream error body to the client protocol. +func (g *Gateway) writeUpstreamError(c *gin.Context, proto Protocol, status int, body []byte) { + switch proto { + case ProtocolAnthropic: + // Extract the upstream Claude error if present. + var up struct { + Error struct { + Type string `json:"type"` + Message string `json:"message"` + } `json:"error"` + } + if json.Unmarshal(body, &up) == nil && up.Error.Message != "" { + c.AbortWithStatusJSON(status, map[string]any{ + "type": "error", + "error": map[string]any{"type": up.Error.Type, "message": up.Error.Message}, + }) + return + } + c.AbortWithStatusJSON(status, map[string]any{ + "type": "error", + "error": map[string]any{"type": statusType(status), "message": upstreamMessage(status, body)}, + }) + default: + var up struct { + Error struct { + Message string `json:"message"` + Type string `json:"type"` + Code string `json:"code"` + Param string `json:"param"` + } `json:"error"` + } + if json.Unmarshal(body, &up) == nil && up.Error.Message != "" { + c.AbortWithStatusJSON(status, map[string]any{ + "error": map[string]any{ + "message": up.Error.Message, "type": up.Error.Type, "code": up.Error.Code, "param": up.Error.Param, + }, + }) + return + } + c.AbortWithStatusJSON(status, map[string]any{ + "error": map[string]any{ + "message": upstreamMessage(status, body), "type": statusType(status), "code": statusType(status), + }, + }) + } +} + +func statusType(status int) string { + switch { + case status == 429: + return "rate_limit_error" + case status >= 500: + return "api_error" + case status >= 400: + return "invalid_request_error" + default: + return "api_error" + } +} + +func upstreamMessage(status int, body []byte) string { + msg := strings.TrimSpace(string(body)) + if msg == "" { + msg = http.StatusText(status) + } + return truncateText(msg, 512) +} + +func truncateText(s string, n int) string { + if len(s) <= n { + return s + } + return s[:n] + "..." +} + +// copyProxyHeaders forwards selected request headers upstream. +func copyProxyHeaders(c *gin.Context, req *http.Request) { + for _, h := range []string{"OpenAI-Organization", "OpenAI-Beta", "anthropic-version", "anthropic-beta", "X-Stainless-Lang", "X-Stainless-Package-Version"} { + if v := c.GetHeader(h); v != "" { + req.Header.Set(h, v) + } + } +} + +// bodyStreamFlag determines streaming intent from the raw body + route. +func bodyStreamFlag(body []byte, route Route) bool { + var m struct { + Stream bool `json:"stream"` + } + _ = json.Unmarshal(body, &m) + return m.Stream +} + +// setModelField rewrites the "model" key in a JSON object. +func setModelField(body []byte, model string) ([]byte, error) { + var obj map[string]json.RawMessage + if err := json.Unmarshal(body, &obj); err != nil { + return nil, err + } + m, _ := json.Marshal(model) + obj["model"] = m + return json.Marshal(obj) +} + +// sameModel reports whether the body's model already equals upstreamModel. +func sameModel(body []byte, upstreamModel string) bool { + var m struct { + Model string `json:"model"` + } + _ = json.Unmarshal(body, &m) + return m.Model == upstreamModel +} + +func contains(list []string, s string) bool { + for _, v := range list { + if v == s { + return true + } + } + return false +} + +// upstreamPath returns the endpoint to POST to. When the channel serves the +// client protocol natively the route's own path is used; otherwise the request +// is converted and must hit the channel's conversion-target path. +func upstreamPath(route Route, ch *store.Channel) string { + if ch.SupportsFormat(string(route.Protocol)) { + return route.UpstreamPath + } + if ch.Provider == "anthropic" { + return "/v1/messages" + } + return "/v1/chat/completions" +} + +// upstreamProvider returns the format family ("openai" | "anthropic") the +// channel will actually speak for this request: the client's own family on +// passthrough, otherwise the channel's conversion target. +func upstreamProvider(route Route, ch *store.Channel) string { + if ch.SupportsFormat(string(route.Protocol)) { + return route.NativeProvider + } + return ch.Provider +} + +// channelServesFormat reports whether the channel can handle the route's +// protocol: natively, or via a conversion path that exists. Every protocol can +// convert to either Claude or chat completions except one case: Responses-API +// requests have no conversion into an openai-family upstream (the +// responses->chat response side is unimplemented), so an openai-family channel +// that does not declare responses support cannot serve them at all. +func channelServesFormat(ch *store.Channel, route Route) bool { + if ch.SupportsFormat(string(route.Protocol)) { + return true + } + return !(route.Protocol == ProtocolOpenAIResponses && ch.Provider != "anthropic") +} + +func now() time.Time { return time.Now() } diff --git a/server/internal/proxy/models.go b/server/internal/proxy/models.go new file mode 100644 index 0000000..01cf33c --- /dev/null +++ b/server/internal/proxy/models.go @@ -0,0 +1,34 @@ +package proxy + +import ( + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +// ModelsHandler serves GET /v1/models (OpenAI-style list). +type ModelsHandler struct { + db *gorm.DB + gw *Gateway +} + +func NewModelsHandler(db *gorm.DB, gw *Gateway) *ModelsHandler { + return &ModelsHandler{db: db, gw: gw} +} + +func (h *ModelsHandler) List(c *gin.Context) { + models, err := h.gw.channel.ListEnabledModels() + if err != nil { + c.JSON(502, gin.H{"error": gin.H{"message": "list models failed", "type": "api_error"}}) + return + } + data := make([]gin.H, 0, len(models)) + for _, m := range models { + data = append(data, gin.H{ + "id": m.Name, + "object": "model", + "created": m.CreatedAt.Unix(), + "owned_by": "openteam", + }) + } + c.JSON(200, gin.H{"object": "list", "data": data}) +} diff --git a/server/internal/proxy/openai/types.go b/server/internal/proxy/openai/types.go new file mode 100644 index 0000000..9dcb725 --- /dev/null +++ b/server/internal/proxy/openai/types.go @@ -0,0 +1,87 @@ +package openai + +import "encoding/json" + +// ChatRequest is a chat completions request. Only the fields the gateway +// needs are typed; the rest is preserved via Raw for passthrough. +type ChatRequest struct { + Model string `json:"model"` + Messages json.RawMessage `json:"messages"` + Stream bool `json:"stream"` + Temperature *float64 `json:"temperature"` + TopP *float64 `json:"top_p"` + MaxTokens *int `json:"max_tokens"` + MaxCompl *int `json:"max_completion_tokens"` + Stop json.RawMessage `json:"stop"` + Tools json.RawMessage `json:"tools"` + ToolChoice json.RawMessage `json:"tool_choice"` + ResponseFmt json.RawMessage `json:"response_format"` + StreamOpts json.RawMessage `json:"stream_options"` + User string `json:"user"` +} + +// ChatMessage is one message in the canonical/chat shape. +type ChatMessage struct { + Role string `json:"role"` + Content json.RawMessage `json:"content,omitempty"` + Name string `json:"name,omitempty"` + ToolCalls json.RawMessage `json:"tool_calls,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` +} + +// Tool is a function tool definition. +type Tool struct { + Type string `json:"type"` + Function FunctionTool `json:"function"` +} + +type FunctionTool struct { + Name string `json:"name"` + Description string `json:"description"` + Parameters json.RawMessage `json:"parameters"` +} + +// ChatCompletion is the non-stream response. +type ChatCompletion struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []ChatChoice `json:"choices"` + Usage *Usage `json:"usage,omitempty"` +} + +type ChatChoice struct { + Index int `json:"index"` + Message ChatMessage `json:"message"` + FinishReason string `json:"finish_reason"` +} + +// Usage is the token usage block. +type Usage struct { + PromptTokens int64 `json:"prompt_tokens"` + CompletionTokens int64 `json:"completion_tokens"` + TotalTokens int64 `json:"total_tokens"` +} + +// ChatChunk is one streaming chunk. +type ChatChunk struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []ChatChunkChoice `json:"choices"` + Usage *Usage `json:"usage,omitempty"` +} + +type ChatChunkChoice struct { + Index int `json:"index"` + Delta ChatDelta `json:"delta"` + FinishReason *string `json:"finish_reason"` +} + +type ChatDelta struct { + Role string `json:"role,omitempty"` + Content string `json:"content,omitempty"` + ToolCalls json.RawMessage `json:"tool_calls,omitempty"` +} diff --git a/server/internal/proxy/responses/types.go b/server/internal/proxy/responses/types.go new file mode 100644 index 0000000..8bc9235 --- /dev/null +++ b/server/internal/proxy/responses/types.go @@ -0,0 +1,42 @@ +package responses + +import "encoding/json" + +// Request is an OpenAI Responses API request. The gateway only needs the +// model + stream fields for routing; the full body passes through raw. +type Request struct { + Model string `json:"model"` + Stream bool `json:"stream"` + Instructions json.RawMessage `json:"instructions"` + Input json.RawMessage `json:"input"` + MaxOutputTokens *int `json:"max_output_tokens"` + PreviousResponseID string `json:"previous_response_id"` + Tools json.RawMessage `json:"tools"` + Reasoning json.RawMessage `json:"reasoning"` + Text json.RawMessage `json:"text"` + OutputFormat json.RawMessage `json:"output_format"` +} + +// Response is the non-stream Responses response. +type Response struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Status string `json:"status"` + Output json.RawMessage `json:"output"` + Usage *Usage `json:"usage,omitempty"` +} + +type Usage struct { + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + TotalTokens int64 `json:"total_tokens"` + InputTokensDetails UsageDetails `json:"input_tokens_details,omitempty"` + OutputTokensDetails UsageDetails `json:"output_tokens_details,omitempty"` +} + +type UsageDetails struct { + CachedTokens int64 `json:"cached_tokens,omitempty"` + ReasoningTokens int64 `json:"reasoning_tokens,omitempty"` +} diff --git a/server/internal/proxy/stream.go b/server/internal/proxy/stream.go new file mode 100644 index 0000000..3954cef --- /dev/null +++ b/server/internal/proxy/stream.go @@ -0,0 +1,201 @@ +package proxy + +import ( + "bufio" + "bytes" + "encoding/json" + "io" + "net/http" + "time" + + "github.com/gin-gonic/gin" + "go.uber.org/zap" + + "openteam/server/internal/proxy/claude" + "openteam/server/internal/proxy/convert" + "openteam/server/internal/proxy/openai" + "openteam/server/internal/proxy/stream" + "openteam/server/internal/store" +) + +// streamOpenAI forwards an SSE stream to an OpenAI-protocol client. +func (g *Gateway) streamOpenAI(c *gin.Context, route Route, resp *http.Response, + start time.Time, reqID string, auth *authContext, model *store.Model, ch *store.Channel) { + + w := c.Writer + tr := convert.NewTranslator(clientProto(route), upstreamProvider(route, ch)) + done := make(chan struct{}) + + if tr == nil { + // Passthrough: copy raw frames, sniffing usage from OpenAI chunks + // and Responses-API stream events. + var usage *tokenUsage + onEvent := func(ev stream.SSEEvent) { + if ev.Done || ev.Data == "" { + return + } + var chunk openai.ChatChunk + if json.Unmarshal([]byte(ev.Data), &chunk) == nil && chunk.Usage != nil { + usage = &tokenUsage{ + input: chunk.Usage.PromptTokens, + output: chunk.Usage.CompletionTokens, + } + return + } + var resp respStreamEvent + if json.Unmarshal([]byte(ev.Data), &resp) == nil && resp.Response != nil && resp.Response.Usage != nil { + usage = &tokenUsage{ + input: resp.Response.Usage.InputTokens, + output: resp.Response.Usage.OutputTokens, + } + } + } + go func() { + defer close(done) + copyRawSSE(w, resp.Body, onEvent) + }() + <-done + g.afterComplete(start, reqID, auth, model, ch, usage, "success", "") + return + } + + // Converted: read upstream events, emit translated frames. + var finalUsage *tokenUsage + go func() { + defer close(done) + err := stream.ReadSSE(resp.Body, func(ev stream.SSEEvent) error { + frames, ferr := tr.Feed(ev) + if ferr != nil { + return ferr + } + for _, f := range frames { + if err := stream.Write(w, f); err != nil { + return err + } + } + return nil + }) + if err != nil { + g.log.Debug("upstream stream read ended", zap.Error(err)) + } + fin, _ := tr.Finish() + for _, f := range fin { + _ = stream.Write(w, f) + } + if u := tr.Usage(); u != nil { + finalUsage = &tokenUsage{input: u.Input, output: u.Output, cacheRead: u.CacheRead, cacheCreation: u.CacheCreation} + } + }() + <-done + g.afterComplete(start, reqID, auth, model, ch, finalUsage, "success", "") +} + +// streamAnthropic forwards an SSE stream to a Claude-protocol client. +func (g *Gateway) streamAnthropic(c *gin.Context, route Route, resp *http.Response, + start time.Time, reqID string, auth *authContext, model *store.Model, ch *store.Channel) { + + w := c.Writer + tr := convert.NewTranslator(clientProto(route), upstreamProvider(route, ch)) + done := make(chan struct{}) + + if tr == nil { + // Passthrough: copy raw frames, sniffing usage from message_delta. + var usage *tokenUsage + onEvent := func(ev stream.SSEEvent) { + if ev.Done || ev.Data == "" { + return + } + var e claude.StreamEvent + if json.Unmarshal([]byte(ev.Data), &e) != nil || e.Type != "message_delta" { + return + } + var u claude.Usage + if json.Unmarshal(e.Usage, &u) == nil { + usage = &tokenUsage{ + input: u.InputTokens, output: u.OutputTokens, + cacheRead: u.CacheReadInputTokens, cacheCreation: u.CacheCreationInputTokens, + } + } + } + go func() { + defer close(done) + copyRawSSE(w, resp.Body, onEvent) + }() + <-done + g.afterComplete(start, reqID, auth, model, ch, usage, "success", "") + return + } + + // Converted: read upstream events, emit translated frames. + var finalUsage *tokenUsage + go func() { + defer close(done) + err := stream.ReadSSE(resp.Body, func(ev stream.SSEEvent) error { + frames, ferr := tr.Feed(ev) + if ferr != nil { + return ferr + } + for _, f := range frames { + if err := stream.Write(w, f); err != nil { + return err + } + } + return nil + }) + if err != nil { + g.log.Debug("upstream stream read ended", zap.Error(err)) + } + fin, _ := tr.Finish() + for _, f := range fin { + _ = stream.Write(w, f) + } + if u := tr.Usage(); u != nil { + finalUsage = &tokenUsage{input: u.Input, output: u.Output, cacheRead: u.CacheRead, cacheCreation: u.CacheCreation} + } + }() + <-done + g.afterComplete(start, reqID, auth, model, ch, finalUsage, "success", "") +} + +// copyRawSSE copies an upstream SSE stream verbatim, invoking onEvent for +// each data frame (used for passthrough + usage sniffing). +func copyRawSSE(w http.ResponseWriter, r io.Reader, onEvent func(stream.SSEEvent)) { + br := bufio.NewReader(r) + for { + line, err := br.ReadBytes('\n') + if len(line) > 0 { + if _, werr := w.Write(line); werr != nil { + return + } + trimmed := bytes.TrimSpace(line) + if bytes.HasPrefix(trimmed, []byte("data:")) { + data := bytes.TrimSpace(trimmed[len("data:"):]) + ev := stream.SSEEvent{ + Data: string(data), + Done: bytes.Equal(data, []byte("[DONE]")), + } + if ev.Data != "" { + onEvent(ev) + } + } + } + if err != nil { + break + } + } + if f, ok := w.(http.Flusher); ok { + f.Flush() + } +} + +// clientProto maps a proxy Protocol to the convert package's protocol type. +func clientProto(route Route) convert.ClientProtocol { + switch route.Protocol { + case ProtocolOpenAIResponses: + return convert.ClientOpenAIResponses + case ProtocolAnthropic: + return convert.ClientAnthropic + default: + return convert.ClientOpenAIChat + } +} diff --git a/server/internal/proxy/stream/stream.go b/server/internal/proxy/stream/stream.go new file mode 100644 index 0000000..db5d305 --- /dev/null +++ b/server/internal/proxy/stream/stream.go @@ -0,0 +1,79 @@ +package stream + +import ( + "bufio" + "bytes" + "io" + "net/http" + "strings" +) + +// SSEEvent is a single SSE data frame. +type SSEEvent struct { + Data string // the JSON payload of the `data:` line + Done bool // true when the payload is [DONE] +} + +// ReadSSE reads SSE frames from r, calling fn for each `data:` line. +// It is used both for reading upstream streams and, via a pipe, for writing +// converted streams to the client. +func ReadSSE(r io.Reader, fn func(SSEEvent) error) error { + br := bufio.NewReader(r) + for { + line, err := br.ReadString('\n') + if err != nil { + if err == io.EOF { + return nil + } + return err + } + line = trimCRLF(line) + if !bytes.HasPrefix([]byte(line), []byte("data:")) { + continue + } + data := line[len("data:"):] + data = strings.TrimPrefix(data, " ") + if len(data) == 0 { + continue + } + ev := SSEEvent{Data: data, Done: bytes.Equal([]byte(data), []byte("[DONE]"))} + if err := fn(ev); err != nil { + return err + } + } +} + +// Write writes an SSE event to w and flushes it. +func Write(w http.ResponseWriter, ev SSEEvent) error { + if ev.Done { + if _, err := io.WriteString(w, "data: [DONE]\n\n"); err != nil { + return err + } + } else { + if _, err := io.WriteString(w, "data: "+ev.Data+"\n\n"); err != nil { + return err + } + } + if f, ok := w.(http.Flusher); ok { + f.Flush() + } + return nil +} + +// WriteRaw writes a raw SSE frame string (with trailing newlines) and flushes. +func WriteRaw(w http.ResponseWriter, frame []byte) error { + if _, err := w.Write(frame); err != nil { + return err + } + if f, ok := w.(http.Flusher); ok { + f.Flush() + } + return nil +} + +func trimCRLF(s string) string { + for len(s) > 0 && (s[len(s)-1] == '\n' || s[len(s)-1] == '\r') { + s = s[:len(s)-1] + } + return s +} diff --git a/server/internal/proxy/usage.go b/server/internal/proxy/usage.go new file mode 100644 index 0000000..a01f8db --- /dev/null +++ b/server/internal/proxy/usage.go @@ -0,0 +1,48 @@ +package proxy + +import ( + "encoding/json" + + "openteam/server/internal/proxy/claude" + "openteam/server/internal/proxy/openai" + "openteam/server/internal/proxy/responses" +) + +// parseUsageForProtocol extracts token usage from a non-stream response body. +func parseUsageForProtocol(proto Protocol, body []byte) *tokenUsage { + switch proto { + case ProtocolAnthropic: + var mr claude.MessageResponse + if err := json.Unmarshal(body, &mr); err != nil || mr.Usage == nil { + return nil + } + return &tokenUsage{ + input: mr.Usage.InputTokens, + output: mr.Usage.OutputTokens, + cacheRead: mr.Usage.CacheReadInputTokens, + cacheCreation: mr.Usage.CacheCreationInputTokens, + } + case ProtocolOpenAIResponses: + var r responses.Response + if err := json.Unmarshal(body, &r); err != nil || r.Usage == nil { + return nil + } + return &tokenUsage{ + input: r.Usage.InputTokens, + output: r.Usage.OutputTokens, + cacheRead: r.Usage.InputTokensDetails.CachedTokens, + } + default: + var cc openai.ChatCompletion + if err := json.Unmarshal(body, &cc); err != nil || cc.Usage == nil { + return nil + } + return &tokenUsage{input: cc.Usage.PromptTokens, output: cc.Usage.CompletionTokens} + } +} + +// estimateTokensFromText is a coarse fallback used when the upstream omits usage. +func estimateTokensFromText(s string) int64 { + // ~4 chars per token, per OpenAI's common heuristic. + return int64(len(s)/4 + 1) +} diff --git a/server/internal/recharge/handler.go b/server/internal/recharge/handler.go new file mode 100644 index 0000000..167eb41 --- /dev/null +++ b/server/internal/recharge/handler.go @@ -0,0 +1,84 @@ +package recharge + +import ( + "net/http" + + "github.com/gin-gonic/gin" + "github.com/shopspring/decimal" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/pkg/httpx" + "openteam/server/internal/store" + "openteam/server/internal/user" +) + +// Handler exposes the reserved recharge order endpoints. +type Handler struct { + db *gorm.DB + log *zap.Logger +} + +func NewHandler(db *gorm.DB, log *zap.Logger) *Handler { + return &Handler{db: db, log: log} +} + +type CreateInput struct { + Amount decimal.Decimal `json:"amount" binding:"required"` + Remark string `json:"remark"` +} + +// Create handles POST /api/v1/recharges — creates a pending manual order. +func (h *Handler) Create(c *gin.Context) { + u := user.Current(c) + var in CreateInput + if !httpx.Bind(c, &in) { + return + } + if in.Amount.LessThanOrEqual(decimal.Zero) { + httpx.Fail(c, http.StatusBadRequest, "amount must be positive") + return + } + order := &store.RechargeOrder{ + UserID: u.ID, + Amount: in.Amount, + Status: "pending", + Method: "manual", + Remark: in.Remark, + } + if err := h.db.Create(order).Error; err != nil { + h.log.Warn("create recharge order failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "create order failed") + return + } + httpx.Created(c, orderDTO(order)) +} + +// List handles GET /api/v1/recharges. +func (h *Handler) List(c *gin.Context) { + u := user.Current(c) + var orders []store.RechargeOrder + if err := h.db.Where("user_id = ?", u.ID).Order("id DESC").Find(&orders).Error; err != nil { + h.log.Warn("list recharge orders failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "list orders failed") + return + } + out := make([]gin.H, 0, len(orders)) + for i := range orders { + out = append(out, orderDTO(&orders[i])) + } + httpx.OK(c, out) +} + +func orderDTO(o *store.RechargeOrder) gin.H { + return gin.H{ + "id": o.ID, + "amount": o.Amount.String(), + "status": o.Status, + "method": o.Method, + "remark": o.Remark, + "reviewedBy": o.ReviewedBy, + "reviewedAt": o.ReviewedAt, + "createdAt": o.CreatedAt, + } +} diff --git a/server/internal/store/model.go b/server/internal/store/model.go new file mode 100644 index 0000000..d0f48d2 --- /dev/null +++ b/server/internal/store/model.go @@ -0,0 +1,196 @@ +package store + +import ( + "time" + + "github.com/shopspring/decimal" + "gorm.io/gorm" +) + +// User is a registered account. +type User struct { + ID int64 `gorm:"primaryKey"` + Username string `gorm:"uniqueIndex;size:64"` + Email string `gorm:"uniqueIndex;size:255"` + PasswordHash string `gorm:"size:255"` + Role string `gorm:"size:16;default:user"` // admin | user + Balance decimal.Decimal `gorm:"type:numeric(20,8);default:0"` + Status string `gorm:"size:16;default:active"` // active | disabled + InviteCode string `gorm:"size:64"` + LastLoginAt *time.Time + CreatedAt time.Time + UpdatedAt time.Time +} + +// ApiKey is an API key issued to a user. +type ApiKey struct { + ID int64 `gorm:"primaryKey"` + UserID int64 `gorm:"index"` + Name string `gorm:"size:128"` + KeyHash string `gorm:"uniqueIndex;size:128"` + KeyPrefix string `gorm:"size:16"` + QuotaTokensPerDay *int64 + QuotaRequestsPerDay *int + AllowedModels []string `gorm:"serializer:json"` + ExpiresAt *time.Time + Status string `gorm:"size:16;default:active"` // active | revoked + LastUsedAt *time.Time + CreatedAt time.Time + UpdatedAt time.Time +} + +// Channel is an upstream provider. +type Channel struct { + ID int64 `gorm:"primaryKey"` + Name string `gorm:"size:128"` + Provider string `gorm:"size:32"` // openai | anthropic | compatible + BaseURL string `gorm:"size:512"` + APIKeyEnc string `gorm:"size:2048"` // AES-GCM ciphertext + Weight int `gorm:"default:1"` + Priority int `gorm:"default:0"` // lower = preferred + TimeoutMs int `gorm:"default:300000"` + MaxConcurrency int `gorm:"default:100"` + HealthStatus string `gorm:"size:16;default:healthy"` // healthy | degraded | cooldown + HealthFailures int + Enabled bool `gorm:"default:true"` + Formats []string `gorm:"serializer:json"` // client API formats served natively; empty = derive from Provider + CreatedAt time.Time + UpdatedAt time.Time +} + +// Client API format keys a channel can serve natively. They match the proxy +// route protocol values so SupportsFormat can compare them directly. +const ( + FormatOpenAIChat = "openai-chat" + FormatOpenAIResponses = "openai-responses" + FormatAnthropic = "anthropic" +) + +// FormatsResolved returns the API formats the channel serves directly. When +// Formats is empty it falls back to the provider's native formats so legacy +// rows behave exactly as before. +func (c *Channel) FormatsResolved() []string { + if len(c.Formats) > 0 { + out := make([]string, len(c.Formats)) + copy(out, c.Formats) + return out + } + if c.Provider == "anthropic" { + return []string{FormatAnthropic} + } + return []string{FormatOpenAIChat, FormatOpenAIResponses} +} + +// SupportsFormat reports whether the channel serves a client protocol +// natively (passthrough) without conversion. +func (c *Channel) SupportsFormat(f string) bool { + for _, x := range c.FormatsResolved() { + if x == f { + return true + } + } + return false +} + +// Model is the global model registry with pricing. +type Model struct { + ID int64 `gorm:"primaryKey"` + Name string `gorm:"uniqueIndex;size:128"` + DisplayName string `gorm:"size:255"` + InputPrice decimal.Decimal `gorm:"type:numeric(20,8);default:0"` // per 1M tokens + OutputPrice decimal.Decimal `gorm:"type:numeric(20,8);default:0"` + CacheReadPrice decimal.Decimal `gorm:"type:numeric(20,8);default:0"` + Enabled bool `gorm:"default:true"` + Sort int `gorm:"default:0"` + CreatedAt time.Time + UpdatedAt time.Time +} + +// ChannelModelBinding binds a global model to a channel with an upstream name. +type ChannelModelBinding struct { + ID int64 `gorm:"primaryKey"` + ChannelID int64 `gorm:"index;uniqueIndex:idx_channel_model"` + ModelID int64 `gorm:"index;uniqueIndex:idx_channel_model"` + UpstreamModel string `gorm:"size:255"` + Weight int `gorm:"default:1"` + Channel Channel `gorm:"foreignKey:ChannelID"` + Model Model `gorm:"foreignKey:ModelID"` +} + +// UsageLog is one proxied request's billing record. +type UsageLog struct { + ID int64 `gorm:"primaryKey"` + RequestID string `gorm:"size:128"` + UserID int64 `gorm:"index:idx_user_created,priority:1"` + KeyID int64 + ChannelID int64 + ModelID int64 + ModelName string `gorm:"size:128"` + InputTokens int64 + OutputTokens int64 + CacheReadTokens int64 + CacheCreationTokens int64 + InputPrice decimal.Decimal `gorm:"type:numeric(20,8)"` + OutputPrice decimal.Decimal `gorm:"type:numeric(20,8)"` + CacheReadPrice decimal.Decimal `gorm:"type:numeric(20,8)"` + Cost decimal.Decimal `gorm:"type:numeric(20,8)"` + LatencyMs int + Status string `gorm:"size:16"` // success | error | canceled + ErrorCode string `gorm:"size:64"` + CreatedAt time.Time `gorm:"index:idx_user_created,priority:2"` +} + +// UsageDaily is the pre-aggregated per-user per-model daily rollup. +type UsageDaily struct { + ID int64 `gorm:"primaryKey"` + UserID int64 `gorm:"index:idx_user_date,priority:1"` + ModelID int64 + Date string `gorm:"size:10;index:idx_user_date,priority:2"` // YYYY-MM-DD + Requests int + InputTokens int64 + OutputTokens int64 + CacheReadTokens int64 + CacheCreationTokens int64 + Cost decimal.Decimal `gorm:"type:numeric(20,8)"` +} + +// RechargeOrder is reserved for the (paused) recharge feature. +type RechargeOrder struct { + ID int64 `gorm:"primaryKey"` + UserID int64 `gorm:"index"` + Amount decimal.Decimal `gorm:"type:numeric(20,8)"` + Status string `gorm:"size:16;default:pending"` // pending | credited | rejected + Method string `gorm:"size:16;default:manual"` // manual | online + TransactionID string `gorm:"size:128"` + ReviewedBy *int64 + ReviewedAt *time.Time + Remark string `gorm:"size:512"` + CreatedAt time.Time + UpdatedAt time.Time + User User `gorm:"foreignKey:UserID"` +} + +// BalanceLog is a user balance ledger entry. +type BalanceLog struct { + ID int64 `gorm:"primaryKey"` + UserID int64 `gorm:"index:idx_balance_user_created,priority:1"` + Change decimal.Decimal `gorm:"type:numeric(20,8)"` + BalanceAfter decimal.Decimal `gorm:"type:numeric(20,8)"` + Type string `gorm:"size:32"` // recharge | usage | refund | admin_adjust + RefID string `gorm:"size:128"` + CreatedAt time.Time `gorm:"index:idx_balance_user_created,priority:2"` +} + +// SystemConfig is a key/value store for runtime settings. +type SystemConfig struct { + Key string `gorm:"primaryKey;size:128"` + Value []byte `gorm:"serializer:json"` +} + +// Migrate creates/upgrades the schema. +func Migrate(db *gorm.DB) error { + return db.AutoMigrate( + &User{}, &ApiKey{}, &Channel{}, &Model{}, &ChannelModelBinding{}, + &UsageLog{}, &UsageDaily{}, &RechargeOrder{}, &BalanceLog{}, &SystemConfig{}, + ) +} diff --git a/server/internal/store/store.go b/server/internal/store/store.go new file mode 100644 index 0000000..f3ef9f3 --- /dev/null +++ b/server/internal/store/store.go @@ -0,0 +1,44 @@ +package store + +import ( + "fmt" + "os" + "path/filepath" + + "github.com/glebarez/sqlite" + "gorm.io/driver/postgres" + "gorm.io/gorm" + "gorm.io/gorm/logger" +) + +func Open(driver, dsn string, debug bool) (*gorm.DB, error) { + cfg := &gorm.Config{} + if debug { + cfg.Logger = logger.Default.LogMode(logger.Info) + } else { + cfg.Logger = logger.Default.LogMode(logger.Warn) + } + + var dialector gorm.Dialector + switch driver { + case "postgres": + dialector = postgres.Open(dsn) + default: + if err := os.MkdirAll(filepath.Dir(dsn), 0o755); err != nil { + return nil, fmt.Errorf("create data dir: %w", err) + } + dialector = sqlite.Open(dsn) + } + + db, err := gorm.Open(dialector, cfg) + if err != nil { + return nil, err + } + sqlDB, err := db.DB() + if err != nil { + return nil, err + } + sqlDB.SetMaxOpenConns(20) + sqlDB.SetMaxIdleConns(5) + return db, nil +} diff --git a/server/internal/usage/handler.go b/server/internal/usage/handler.go new file mode 100644 index 0000000..c4ccde5 --- /dev/null +++ b/server/internal/usage/handler.go @@ -0,0 +1,193 @@ +package usage + +import ( + "net/http" + "strconv" + "time" + + "github.com/gin-gonic/gin" + "github.com/shopspring/decimal" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/pkg/httpx" + "openteam/server/internal/store" + "openteam/server/internal/user" +) + +type Handler struct { + db *gorm.DB + log *zap.Logger +} + +func NewHandler(db *gorm.DB, log *zap.Logger) *Handler { + return &Handler{db: db, log: log} +} + +// Summary handles GET /api/v1/usage/summary. +func (h *Handler) Summary(c *gin.Context) { + u := user.Current(c) + now := time.Now() + today := now.Format("2006-01-02") + month := now.Format("2006-01") + + type agg struct { + Requests int + InputTokens int64 + OutputTokens int64 + Cost decimal.Decimal + } + + todayAgg := h.aggregate(u.ID, today, today) + monthAgg := h.aggregate(u.ID, month+"-01", now.Format("2006-01-02")) + total := h.aggregate(u.ID, "", "") + + httpx.OK(c, gin.H{ + "today": gin.H{ + "requests": todayAgg.Requests, + "inputTokens": todayAgg.InputTokens, + "outputTokens": todayAgg.OutputTokens, + "cost": todayAgg.Cost.String(), + }, + "month": gin.H{ + "requests": monthAgg.Requests, + "inputTokens": monthAgg.InputTokens, + "outputTokens": monthAgg.OutputTokens, + "cost": monthAgg.Cost.String(), + }, + "total": gin.H{ + "requests": total.Requests, + "inputTokens": total.InputTokens, + "outputTokens": total.OutputTokens, + "cost": total.Cost.String(), + }, + }) +} + +func (h *Handler) aggregate(userID int64, from, to string) struct { + Requests int + InputTokens int64 + OutputTokens int64 + Cost decimal.Decimal +} { + var out struct { + Requests int + InputTokens int64 + OutputTokens int64 + Cost decimal.Decimal + } + q := h.db.Model(&store.UsageDaily{}).Where("user_id = ?", userID) + if from != "" { + q = q.Where("date >= ?", from) + } + if to != "" { + q = q.Where("date <= ?", to) + } + q.Select("COALESCE(SUM(requests),0) as requests, COALESCE(SUM(input_tokens),0) as input_tokens, COALESCE(SUM(output_tokens),0) as output_tokens, COALESCE(SUM(cost),0) as cost"). + Scan(&out) + return out +} + +// Stats handles GET /api/v1/usage/stats?from&to&group=day|model. +func (h *Handler) Stats(c *gin.Context) { + u := user.Current(c) + from := c.Query("from") + to := c.Query("to") + group := c.DefaultQuery("group", "day") + + var rows []struct { + Key string `gorm:"column:g"` + Requests int + InputTokens int64 + OutputTokens int64 + Cost decimal.Decimal + } + + q := h.db.Model(&store.UsageDaily{}).Where("user_id = ?", u.ID) + if from != "" { + q = q.Where("date >= ?", from) + } + if to != "" { + q = q.Where("date <= ?", to) + } + + switch group { + case "model": + q = q.Joins("JOIN models ON models.id = usage_dailies.model_id"). + Select("models.name as g, COALESCE(SUM(requests),0) as requests, COALESCE(SUM(input_tokens),0) as input_tokens, COALESCE(SUM(output_tokens),0) as output_tokens, COALESCE(SUM(cost),0) as cost"). + Group("models.name") + default: + q = q.Select("date as g, COALESCE(SUM(requests),0) as requests, COALESCE(SUM(input_tokens),0) as input_tokens, COALESCE(SUM(output_tokens),0) as output_tokens, COALESCE(SUM(cost),0) as cost"). + Group("date").Order("date ASC") + } + if err := q.Scan(&rows).Error; err != nil { + h.log.Warn("usage stats failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "usage stats failed") + return + } + out := make([]gin.H, 0, len(rows)) + for _, r := range rows { + out = append(out, gin.H{ + "key": r.Key, + "requests": r.Requests, + "inputTokens": r.InputTokens, + "outputTokens": r.OutputTokens, + "cost": r.Cost.String(), + }) + } + httpx.OK(c, out) +} + +// Logs handles GET /api/v1/usage/logs?from&to&page&model&keyId. +func (h *Handler) Logs(c *gin.Context) { + u := user.Current(c) + page, _ := strconv.Atoi(c.DefaultQuery("page", "1")) + pageSize := 20 + if page < 1 { + page = 1 + } + from := c.Query("from") + to := c.Query("to") + + q := h.db.Model(&store.UsageLog{}).Where("user_id = ?", u.ID) + if from != "" { + q = q.Where("created_at >= ?", from) + } + if to != "" { + q = q.Where("created_at <= ?", to+" 23:59:59") + } + if m := c.Query("model"); m != "" { + q = q.Where("model_name = ?", m) + } + if kid := c.Query("keyId"); kid != "" { + q = q.Where("key_id = ?", kid) + } + + var total int64 + q.Count(&total) + + var logs []store.UsageLog + if err := q.Order("id DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&logs).Error; err != nil { + h.log.Warn("usage logs failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "usage logs failed") + return + } + out := make([]gin.H, 0, len(logs)) + for _, l := range logs { + out = append(out, gin.H{ + "id": l.ID, + "requestId": l.RequestID, + "model": l.ModelName, + "channelId": l.ChannelID, + "inputTokens": l.InputTokens, + "outputTokens": l.OutputTokens, + "cacheReadTokens": l.CacheReadTokens, + "cost": l.Cost.String(), + "latencyMs": l.LatencyMs, + "status": l.Status, + "errorCode": l.ErrorCode, + "createdAt": l.CreatedAt, + }) + } + httpx.OK(c, gin.H{"total": total, "page": page, "pageSize": pageSize, "items": out}) +} diff --git a/server/internal/usage/service.go b/server/internal/usage/service.go new file mode 100644 index 0000000..9cdb7fa --- /dev/null +++ b/server/internal/usage/service.go @@ -0,0 +1,91 @@ +package usage + +import ( + "time" + + "github.com/shopspring/decimal" + "go.uber.org/zap" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "openteam/server/internal/store" +) + +type Record struct { + RequestID string + UserID int64 + KeyID int64 + ChannelID int64 + ModelID int64 + ModelName string + InputTokens int64 + OutputTokens int64 + CacheReadTokens int64 + CacheCreationTokens int64 + InputPrice decimal.Decimal + OutputPrice decimal.Decimal + CacheReadPrice decimal.Decimal + Cost decimal.Decimal + LatencyMs int + Status string + ErrorCode string +} + +type Service struct { + db *gorm.DB + log *zap.Logger +} + +func NewService(db *gorm.DB, log *zap.Logger) *Service { + return &Service{db: db, log: log} +} + +// Record inserts a request-level usage log and upserts the daily aggregate. +func (s *Service) Record(r Record) error { + date := time.Now().Format("2006-01-02") + err := s.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Create(&store.UsageLog{ + RequestID: r.RequestID, + UserID: r.UserID, + KeyID: r.KeyID, + ChannelID: r.ChannelID, + ModelID: r.ModelID, + ModelName: r.ModelName, + InputTokens: r.InputTokens, + OutputTokens: r.OutputTokens, + CacheReadTokens: r.CacheReadTokens, + CacheCreationTokens: r.CacheCreationTokens, + InputPrice: r.InputPrice, + OutputPrice: r.OutputPrice, + CacheReadPrice: r.CacheReadPrice, + Cost: r.Cost, + LatencyMs: r.LatencyMs, + Status: r.Status, + ErrorCode: r.ErrorCode, + CreatedAt: time.Now(), + }).Error; err != nil { + return err + } + + var daily store.UsageDaily + err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). + Where("user_id = ? AND model_id = ? AND date = ?", r.UserID, r.ModelID, date). + FirstOrCreate(&daily, store.UsageDaily{ + UserID: r.UserID, ModelID: r.ModelID, Date: date, + }).Error + if err != nil { + return err + } + daily.Requests++ + daily.InputTokens += r.InputTokens + daily.OutputTokens += r.OutputTokens + daily.CacheReadTokens += r.CacheReadTokens + daily.CacheCreationTokens += r.CacheCreationTokens + daily.Cost = daily.Cost.Add(r.Cost) + return tx.Save(&daily).Error + }) + if err != nil { + s.log.Warn("record usage failed", zap.Error(err)) + } + return err +} diff --git a/server/internal/user/handler.go b/server/internal/user/handler.go new file mode 100644 index 0000000..9927548 --- /dev/null +++ b/server/internal/user/handler.go @@ -0,0 +1,107 @@ +package user + +import ( + "errors" + "net/http" + + "github.com/gin-gonic/gin" + "go.uber.org/zap" + + "openteam/server/internal/config" + "openteam/server/internal/pkg/httpx" +) + +type Handler struct { + svc *Service + cfg *config.Config + log *zap.Logger +} + +func NewHandler(svc *Service, cfg *config.Config, log *zap.Logger) *Handler { + return &Handler{svc: svc, cfg: cfg, log: log} +} + +// Register handles POST /api/v1/auth/register. +func (h *Handler) Register(c *gin.Context) { + var in RegisterInput + if !httpx.Bind(c, &in) { + return + } + u, pair, err := h.svc.Register(in) + if err != nil { + if errors.Is(err, ErrUserExists) { + httpx.Fail(c, http.StatusConflict, err.Error()) + return + } + h.log.Warn("register failed", zap.Error(err)) + httpx.Fail(c, http.StatusBadRequest, err.Error()) + return + } + h.setRefreshCookie(c, pair.RefreshToken) + httpx.Created(c, gin.H{"user": GetPublic(u), "token": pair}) +} + +// Login handles POST /api/v1/auth/login. +func (h *Handler) Login(c *gin.Context) { + var in LoginInput + if !httpx.Bind(c, &in) { + return + } + u, pair, err := h.svc.Login(in) + if err != nil { + if errors.Is(err, ErrBadCredentials) || errors.Is(err, ErrUserDisabled) { + httpx.Fail(c, http.StatusUnauthorized, err.Error()) + return + } + h.log.Warn("login failed", zap.Error(err)) + httpx.Fail(c, http.StatusInternalServerError, "login failed") + return + } + h.setRefreshCookie(c, pair.RefreshToken) + httpx.OK(c, gin.H{"user": GetPublic(u), "token": pair}) +} + +// Refresh handles POST /api/v1/auth/refresh. +func (h *Handler) Refresh(c *gin.Context) { + token, err := c.Cookie(h.cfg.Auth.RefreshCookieName) + if err != nil || token == "" { + httpx.Fail(c, http.StatusUnauthorized, "missing refresh token") + return + } + pair, err := h.svc.Refresh(token) + if err != nil { + httpx.Fail(c, http.StatusUnauthorized, "invalid refresh token") + return + } + h.setRefreshCookie(c, pair.RefreshToken) + httpx.OK(c, gin.H{"token": pair}) +} + +// Logout handles POST /api/v1/auth/logout. +func (h *Handler) Logout(c *gin.Context) { + c.SetCookie(h.cfg.Auth.RefreshCookieName, "", -1, "/", "", h.cfg.Auth.RefreshCookieSecure, true) + httpx.OK(c, gin.H{"ok": true}) +} + +// Me handles GET /api/v1/auth/me. +func (h *Handler) Me(c *gin.Context) { + u := Current(c) + httpx.OK(c, GetPublic(u)) +} + +// Profile handles GET /api/v1/user/profile. +func (h *Handler) Profile(c *gin.Context) { + u := Current(c) + httpx.OK(c, GetPublic(u)) +} + +// Balance handles GET /api/v1/user/balance. +func (h *Handler) Balance(c *gin.Context) { + u := Current(c) + httpx.OK(c, gin.H{"balance": u.Balance.String()}) +} + +func (h *Handler) setRefreshCookie(c *gin.Context, token string) { + c.SetCookie(h.cfg.Auth.RefreshCookieName, token, + int(h.cfg.Auth.RefreshTokenTTL.Seconds()), "/", "", h.cfg.Auth.RefreshCookieSecure, true) +} diff --git a/server/internal/user/middleware.go b/server/internal/user/middleware.go new file mode 100644 index 0000000..6377385 --- /dev/null +++ b/server/internal/user/middleware.go @@ -0,0 +1,74 @@ +package user + +import ( + "net/http" + "strings" + + "github.com/gin-gonic/gin" + + "openteam/server/internal/pkg/httpx" + "openteam/server/internal/pkg/jwt" + "openteam/server/internal/store" +) + +const ( + ctxUserKey = "current_user" + ctxClaimsKey = "current_claims" +) + +// Current returns the authenticated user (set by Middleware). +func Current(c *gin.Context) *store.User { + if v, ok := c.Get(ctxUserKey); ok { + if u, ok := v.(*store.User); ok { + return u + } + } + return nil +} + +// Claims returns the JWT claims of the current request. +func Claims(c *gin.Context) *jwt.Claims { + if v, ok := c.Get(ctxClaimsKey); ok { + if cl, ok := v.(*jwt.Claims); ok { + return cl + } + } + return nil +} + +// Middleware authenticates the management API via an access token. +func (s *Service) Middleware(secret string) gin.HandlerFunc { + return func(c *gin.Context) { + auth := c.GetHeader("Authorization") + if !strings.HasPrefix(auth, "Bearer ") { + httpx.Fail(c, http.StatusUnauthorized, "missing bearer token") + return + } + token := strings.TrimPrefix(auth, "Bearer ") + claims, err := jwt.Parse(secret, token) + if err != nil || claims.Type != "access" { + httpx.Fail(c, http.StatusUnauthorized, "invalid or expired token") + return + } + var u store.User + if err := s.db.First(&u, claims.UserID).Error; err != nil || u.Status != "active" { + httpx.Fail(c, http.StatusUnauthorized, "user not found or disabled") + return + } + c.Set(ctxUserKey, &u) + c.Set(ctxClaimsKey, claims) + c.Next() + } +} + +// RequireAdmin guards admin-only routes. +func RequireAdmin() gin.HandlerFunc { + return func(c *gin.Context) { + u := Current(c) + if u == nil || u.Role != "admin" { + httpx.Fail(c, http.StatusForbidden, "admin only") + return + } + c.Next() + } +} diff --git a/server/internal/user/service.go b/server/internal/user/service.go new file mode 100644 index 0000000..40e3a3a --- /dev/null +++ b/server/internal/user/service.go @@ -0,0 +1,192 @@ +package user + +import ( + "encoding/json" + "errors" + "time" + + "github.com/shopspring/decimal" + "go.uber.org/zap" + "gorm.io/gorm" + + "openteam/server/internal/config" + "openteam/server/internal/pkg/jwt" + "openteam/server/internal/pkg/password" + "openteam/server/internal/store" +) + +var ( + ErrUserExists = errors.New("username or email already exists") + ErrBadCredentials = errors.New("invalid username/email or password") + ErrUserDisabled = errors.New("account disabled") + ErrInvalidRefresh = errors.New("invalid refresh token") +) + +type Service struct { + db *gorm.DB + cfg *config.Config + log *zap.Logger +} + +func NewService(db *gorm.DB, cfg *config.Config, log *zap.Logger) *Service { + return &Service{db: db, cfg: cfg, log: log} +} + +type RegisterInput struct { + Username string `json:"username" binding:"required,min=3,max=64"` + Email string `json:"email" binding:"required,email"` + Password string `json:"password" binding:"required,min=8,max=128"` + InviteCode string `json:"inviteCode"` +} + +type TokenPair struct { + AccessToken string `json:"accessToken"` + TokenType string `json:"tokenType"` + ExpiresIn int64 `json:"expiresIn"` + RefreshToken string `json:"refreshToken"` +} + +// Register creates a user. Returns a token pair plus the created user. +func (s *Service) Register(in RegisterInput) (*store.User, *TokenPair, error) { + if s.cfg.Registration.Mode == "invite" { + // Invite mode: validate the invite code before allowing signup. + ok, err := s.validateInviteCode(in.InviteCode) + if err != nil { + return nil, nil, err + } + if !ok { + return nil, nil, errors.New("invalid invite code") + } + } + + hash, err := password.Hash(in.Password) + if err != nil { + return nil, nil, err + } + u := &store.User{ + Username: in.Username, + Email: in.Email, + PasswordHash: hash, + Role: "user", + Balance: decimalZero(), + Status: "active", + InviteCode: in.InviteCode, + } + if err := s.db.Create(u).Error; err != nil { + if errors.Is(err, gorm.ErrDuplicatedKey) { + return nil, nil, ErrUserExists + } + return nil, nil, err + } + pair, err := s.issuePair(u) + if err != nil { + return nil, nil, err + } + return u, pair, nil +} + +type LoginInput struct { + Account string `json:"account" binding:"required"` // username or email + Password string `json:"password" binding:"required"` +} + +// Login authenticates by username or email and returns a token pair. +func (s *Service) Login(in LoginInput) (*store.User, *TokenPair, error) { + var u store.User + err := s.db.Where("username = ? OR email = ?", in.Account, in.Account).First(&u).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil, ErrBadCredentials + } + if err != nil { + return nil, nil, err + } + ok, err := password.Verify(in.Password, u.PasswordHash) + if err != nil || !ok { + return nil, nil, ErrBadCredentials + } + if u.Status != "active" { + return nil, nil, ErrUserDisabled + } + now := time.Now() + u.LastLoginAt = &now + s.db.Model(&u).Update("last_login_at", now) + pair, err := s.issuePair(&u) + if err != nil { + return nil, nil, err + } + return &u, pair, nil +} + +// Refresh issues a new access token from a valid refresh token. +func (s *Service) Refresh(refreshToken string) (*TokenPair, error) { + claims, err := jwt.Parse(s.cfg.Auth.JWTSecret, refreshToken) + if err != nil || claims.Type != "refresh" { + return nil, ErrInvalidRefresh + } + var u store.User + if err := s.db.First(&u, claims.UserID).Error; err != nil { + return nil, ErrInvalidRefresh + } + if u.Status != "active" { + return nil, ErrUserDisabled + } + return s.issuePair(&u) +} + +// GetByID loads a user. +func (s *Service) GetByID(id int64) (*store.User, error) { + var u store.User + if err := s.db.First(&u, id).Error; err != nil { + return nil, err + } + return &u, nil +} + +// GetPublic returns a user DTO without secrets. +func GetPublic(u *store.User) map[string]any { + return map[string]any{ + "id": u.ID, + "username": u.Username, + "email": u.Email, + "role": u.Role, + "balance": u.Balance.String(), + "status": u.Status, + "createdAt": u.CreatedAt, + "lastLoginAt": u.LastLoginAt, + } +} + +func (s *Service) validateInviteCode(code string) (bool, error) { + var conf store.SystemConfig + if err := s.db.First(&conf, "key = ?", "invite_codes").Error; err != nil { + return false, nil + } + var codes map[string]bool + if len(conf.Value) > 0 { + if err := json.Unmarshal(conf.Value, &codes); err != nil { + return false, err + } + } + return codes[code], nil +} + +func (s *Service) issuePair(u *store.User) (*TokenPair, error) { + access, err := jwt.SignAccess(s.cfg.Auth.JWTSecret, u.ID, u.Username, u.Role, s.cfg.Auth.AccessTokenTTL) + if err != nil { + return nil, err + } + refresh, err := jwt.SignRefresh(s.cfg.Auth.JWTSecret, u.ID, s.cfg.Auth.RefreshTokenTTL) + if err != nil { + return nil, err + } + return &TokenPair{ + AccessToken: access, + TokenType: "Bearer", + ExpiresIn: int64(s.cfg.Auth.AccessTokenTTL.Seconds()), + RefreshToken: refresh, + }, nil +} + +func decimalZero() decimal.Decimal { + return decimal.NewFromInt(0) +}