feat(keys): role-based gateway keys with admin management UI and per-user model scope

This commit is contained in:
root
2026-08-09 10:01:40 +08:00
parent dec03238dd
commit 3408c9cb1f
10 changed files with 901 additions and 82 deletions

View File

@ -79,12 +79,17 @@ func isAuto(m string) bool {
// resolveCands picks the ordered candidate providers for a requested model.
// toolCalling requests are anchored: they resolve to exactly one provider
// (highest-priority available) so a tool-call round never switches models.
func (g *Gateway) resolveCands(req *chatRequest) ([]*provider.Provider, string) {
func (g *Gateway) resolveCands(ctx context.Context, req *chatRequest) ([]*provider.Provider, string) {
model := req.Model
if model == "" {
model = g.core.DefaultModel()
}
cands, effective := g.resolveByModel(model)
cands = chatOnly(cands)
allow := g.allowedModels(ctx)
if allow != nil {
cands = filterCandsByModels(cands, allow)
}
if !toolRequest(req) {
return cands, effective
}
@ -100,6 +105,42 @@ func (g *Gateway) resolveCands(req *chatRequest) ([]*provider.Provider, string)
return []*provider.Provider{first}, eff
}
// filterCandsByModels keeps only providers exposing at least one model of the
// whitelist (used for user keys with a restricted model scope).
func filterCandsByModels(cands []*provider.Provider, allow []string) []*provider.Provider {
allowed := make(map[string]bool, len(allow))
for _, m := range allow {
allowed[m] = true
}
out := make([]*provider.Provider, 0, len(cands))
for _, p := range cands {
for _, id := range p.Models() {
if allowed[id] {
out = append(out, p)
break
}
}
}
return out
}
// intersectModels restricts a model list to the whitelist (preserving order).
func intersectModels(models, allow []string) []string {
allowed := make(map[string]bool, len(allow))
for _, m := range allow {
allowed[m] = true
}
out := make([]string, 0, len(models))
seen := map[string]bool{}
for _, m := range models {
if allowed[m] && !seen[m] {
seen[m] = true
out = append(out, m)
}
}
return out
}
func (g *Gateway) resolveByModel(model string) ([]*provider.Provider, string) {
if isAuto(model) {
return g.core.Registry().Resolve("AUTO"), ""
@ -107,6 +148,22 @@ func (g *Gateway) resolveByModel(model string) ([]*provider.Provider, string) {
return g.core.Registry().Resolve(model), model
}
// chatOnly keeps providers that expose at least one chat-capable model, so a
// chat/AUTO request never lands on an image-only source (or borrows its image
// model id). Explicit image-kind requests stay on the imageOnly path.
func chatOnly(cands []*provider.Provider) []*provider.Provider {
out := make([]*provider.Provider, 0, len(cands))
for _, p := range cands {
for _, id := range p.Models() {
if m := p.ModelByID(id); m == nil || m.Kind != "image" {
out = append(out, p)
break
}
}
}
return out
}
// toolRequest reports whether the request participates in a tool-call round.
func toolRequest(req *chatRequest) bool {
if len(req.Tools) > 0 || req.ToolChoice != nil {
@ -138,7 +195,13 @@ func (g *Gateway) handleChat(w http.ResponseWriter, r *http.Request) {
if model == "" {
model = g.core.DefaultModel()
}
cands, effective := g.resolveCands(&req)
if !isAuto(model) {
if allow := g.allowedModels(r.Context()); allow != nil && !containsStr(allow, model) {
writeError(w, http.StatusForbidden, "model_not_allowed", fmt.Sprintf("model %q is not allowed for this key", model))
return
}
}
cands, effective := g.resolveCands(r.Context(), &req)
if len(cands) == 0 {
writeError(w, http.StatusServiceUnavailable, "no_provider", "no LLM source configured")
return
@ -237,10 +300,19 @@ func firstSource(cands []*provider.Provider) string {
return ""
}
func containsStr(list []string, s string) bool {
for _, x := range list {
if x == s {
return true
}
}
return false
}
func (g *Gateway) singleChat(w http.ResponseWriter, ctx context.Context, cands []*provider.Provider, req *types.ChatRequest, effective string, rec *Req) {
rec.LatMs = 0
t0 := time.Now()
resp, err := g.core.Scheduler().Chat(ctx, scheduler.FromRegistry(cands), req)
resp, usedSrc, usedModel, err := g.core.Scheduler().Chat(ctx, scheduler.FromRegistry(cands), req)
rec.LatMs = time.Since(t0).Milliseconds()
if err != nil {
rec.OK = false
@ -254,7 +326,8 @@ func (g *Gateway) singleChat(w http.ResponseWriter, ctx context.Context, cands [
rec.Status = http.StatusOK
rec.Prompt = int64(resp.TokenUsage.Prompt)
rec.Compl = int64(resp.TokenUsage.Completion)
rec.Source = firstSource(cands)
rec.Source = usedSrc
rec.Model = usedModel
g.writeRec(rec)
msg := RespMessage{Role: "assistant", Content: resp.Content}
if resp.ReasoningContent != "" {
@ -296,7 +369,7 @@ func (g *Gateway) streamChat(w http.ResponseWriter, ctx context.Context, cands [
rec.LatMs = time.Since(t0).Milliseconds()
g.writeRec(rec)
}()
chunks, err := g.core.Scheduler().ChatStream(ctx, scheduler.FromRegistry(cands), req)
chunks, _, usedModel, err := g.core.Scheduler().ChatStream(ctx, scheduler.FromRegistry(cands), req)
if err != nil {
rec.OK = false
rec.Status = http.StatusBadGateway
@ -304,6 +377,9 @@ func (g *Gateway) streamChat(w http.ResponseWriter, ctx context.Context, cands [
writeError(w, http.StatusBadGateway, "upstream_error", err.Error())
return
}
if usedModel != "" {
rec.Model = usedModel
}
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
@ -383,8 +459,17 @@ func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) {
if model == "" {
model = g.core.DefaultModel()
}
if !isAuto(model) {
if allow := g.allowedModels(r.Context()); allow != nil && !containsStr(allow, model) {
writeError(w, http.StatusForbidden, "model_not_allowed", fmt.Sprintf("model %q is not allowed for this key", model))
return
}
}
cands, _ := g.resolveByModel(model)
cands = imageOnly(cands)
if allow := g.allowedModels(r.Context()); allow != nil {
cands = filterCandsByModels(cands, allow)
}
if len(cands) == 0 {
writeError(w, http.StatusServiceUnavailable, "no_provider", "no image source configured")
return
@ -393,7 +478,7 @@ func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) {
defer done()
rec := &Req{Key: keyID(reqKey(r.Context())), Type: "image", Model: model, Source: firstSource(cands), OK: false}
t0 := time.Now()
resp, err := g.core.Scheduler().Image(r.Context(), scheduler.FromRegistry(cands), &req)
resp, usedSrc, err := g.core.Scheduler().Image(r.Context(), scheduler.FromRegistry(cands), &req)
rec.LatMs = time.Since(t0).Milliseconds()
if err != nil {
rec.Status = http.StatusBadGateway
@ -402,6 +487,9 @@ func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) {
writeError(w, http.StatusBadGateway, "upstream_error", err.Error())
return
}
if usedSrc != "" {
rec.Source = usedSrc
}
rec.OK = true
rec.Status = http.StatusOK
rec.Compl = int64(len(resp.ImageData))