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 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 5
parent
b0c7439c01
commit
d0e31b198f
@@ -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()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
)
|
||||||
+124
@@ -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=
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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})
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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})
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
|
||||||
@@ -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})
|
||||||
|
}
|
||||||
@@ -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] + "..."
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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"`
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
@@ -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() }
|
||||||
@@ -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})
|
||||||
|
}
|
||||||
@@ -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"`
|
||||||
|
}
|
||||||
@@ -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"`
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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{},
|
||||||
|
)
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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})
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user