mirror of
https://gitcode.com/JianFeeeee/ModelRouter.git
synced 2026-10-03 15:44:05 +00:00
feat(keys): role-based gateway keys with admin management UI and per-user model scope
This commit is contained in:
@ -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))
|
||||
|
||||
Reference in New Issue
Block a user