From 1efa87e50a60d290b92e6f269459babc96e8c8fa Mon Sep 17 00:00:00 2001 From: Max Date: Wed, 29 Apr 2026 10:28:07 +0800 Subject: [PATCH] feat(llm): add Yao Agents preset and LLM management endpoints - Introduced a new preset for Yao Agents in presets.yml, including configuration details such as API URL and capabilities. - Added new LLM management endpoints in the OpenAPI settings, allowing for CRUD operations on LLM providers and roles. - Enhanced data structures to support the new LLM functionality, including LLMPageData for aggregated responses. --- llmprovider/presets.yml | 24 +- openapi/setting/llm.go | 668 ++++++++++++++++++++++++++++++ openapi/setting/setting.go | 9 + openapi/setting/types.go | 11 + openapi/tests/setting/llm_test.go | 453 ++++++++++++++++++++ 5 files changed, 1153 insertions(+), 12 deletions(-) create mode 100644 openapi/setting/llm.go create mode 100644 openapi/tests/setting/llm_test.go diff --git a/llmprovider/presets.yml b/llmprovider/presets.yml index 65552448..8340968e 100644 --- a/llmprovider/presets.yml +++ b/llmprovider/presets.yml @@ -1,3 +1,15 @@ +- key: yaoagents + name: Yao Agents + type: openai + api_url: https://api.yaoagents.com + require_key: false + is_cloud: true + default_models: + - id: default + name: Default + capabilities: [vision, tool_calls, streaming] + enabled: true + - key: openai name: OpenAI type: openai @@ -47,15 +59,3 @@ require_key: true url_editable: true default_models: [] - -- key: yaoagents - name: Yao Agents - type: openai - api_url: https://api.yaoagents.com - require_key: false - is_cloud: true - default_models: - - id: default - name: Default - capabilities: [vision, tool_calls, streaming] - enabled: true diff --git a/openapi/setting/llm.go b/openapi/setting/llm.go new file mode 100644 index 00000000..566028e2 --- /dev/null +++ b/openapi/setting/llm.go @@ -0,0 +1,668 @@ +package setting + +import ( + "encoding/json" + "fmt" + "net/http" + "strings" + "time" + + "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/config" + "github.com/yaoapp/yao/llmprovider" + "github.com/yaoapp/yao/openapi/oauth/authorized" + oauthTypes "github.com/yaoapp/yao/openapi/oauth/types" + "github.com/yaoapp/yao/openapi/response" + "github.com/yaoapp/yao/setting" +) + +const llmRolesNS = "llm.roles" + +func llmEnsureEncKey() { + if llmprovider.Global != nil && config.Conf.DB.AESKey != "" { + llmprovider.Global.SetEncryptionKey(config.Conf.DB.AESKey) + } +} + +func llmOwner(info *oauthTypes.AuthorizedInfo) *llmprovider.ProviderOwner { + if info.TeamID != "" { + return &llmprovider.ProviderOwner{Type: "team", TeamID: info.TeamID} + } + return &llmprovider.ProviderOwner{Type: "user", UserID: info.UserID} +} + +func llmScope(info *oauthTypes.AuthorizedInfo) setting.ScopeID { + if info.TeamID != "" { + return setting.ScopeID{Scope: setting.ScopeTeam, TeamID: info.TeamID} + } + return setting.ScopeID{Scope: setting.ScopeUser, UserID: info.UserID} +} + +func llmCheckOwnership(p *llmprovider.Provider, info *oauthTypes.AuthorizedInfo) error { + owner := llmOwner(info) + if p.Owner.Type != owner.Type { + return fmt.Errorf("provider not found") + } + if owner.Type == "team" && p.Owner.TeamID != owner.TeamID { + return fmt.Errorf("provider not found") + } + if owner.Type == "user" && p.Owner.UserID != owner.UserID { + return fmt.Errorf("provider not found") + } + return nil +} + +func enrichProvider(p *llmprovider.Provider) map[string]interface{} { + raw, _ := json.Marshal(p) + var m map[string]interface{} + json.Unmarshal(raw, &m) + + if p.PresetKey != "" { + if preset := llmprovider.GetPreset(p.PresetKey); preset != nil { + m["is_cloud"] = preset.IsCloud + m["url_editable"] = preset.URLEditable + } + } + + delete(m, "connector_id") + delete(m, "source") + delete(m, "owner") + + return m +} + +// llmModelsURL builds the models endpoint URL. +// Trailing slash means the user already specified the path prefix → append "models". +// No trailing slash → append "/v1/models" (standard OpenAI convention). +func llmModelsURL(apiURL string) string { + if strings.HasSuffix(apiURL, "/") { + return apiURL + "models" + } + return apiURL + "/v1/models" +} + +// llmValidateKey tests connectivity by calling GET {apiURL}/models. +// providerType controls the auth header format (anthropic uses x-api-key). +func llmValidateKey(providerType, apiURL, apiKey string) error { + url := llmModelsURL(apiURL) + client := &http.Client{Timeout: 10 * time.Second} + req, err := http.NewRequest("GET", url, nil) + if err != nil { + return fmt.Errorf("failed to build request: %w", err) + } + if apiKey != "" { + if providerType == "anthropic" { + req.Header.Set("x-api-key", apiKey) + req.Header.Set("anthropic-version", "2023-06-01") + } else { + req.Header.Set("Authorization", "Bearer "+apiKey) + } + } + resp, err := client.Do(req) + if err != nil { + return fmt.Errorf("connection failed: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden { + return fmt.Errorf("invalid API key (HTTP %d)", resp.StatusCode) + } + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("server returned HTTP %d", resp.StatusCode) + } + return nil +} + +// --------------------------------------------------------------------------- +// Handlers +// --------------------------------------------------------------------------- + +// handleLLMTest validates an API URL + Key without saving. +// POST /setting/llm/test +func handleLLMTest(c *gin.Context) { + if !guardOwner(c) { + return + } + + var input struct { + APIURL string `json:"api_url"` + APIKey string `json:"api_key"` + Type string `json:"type"` + } + if err := c.ShouldBindJSON(&input); err != nil { + respondError(c, http.StatusBadRequest, "invalid request body") + return + } + if input.APIURL == "" { + respondError(c, http.StatusBadRequest, "api_url is required") + return + } + + url := llmModelsURL(input.APIURL) + start := time.Now() + client := &http.Client{Timeout: 10 * time.Second} + req, err := http.NewRequest("GET", url, nil) + if err != nil { + respondError(c, http.StatusInternalServerError, err.Error()) + return + } + if input.APIKey != "" { + if input.Type == "anthropic" { + req.Header.Set("x-api-key", input.APIKey) + req.Header.Set("anthropic-version", "2023-06-01") + } else { + req.Header.Set("Authorization", "Bearer "+input.APIKey) + } + } + + resp, err := client.Do(req) + latency := time.Since(start).Milliseconds() + + if err != nil { + response.RespondWithSuccess(c, http.StatusOK, llmprovider.ProviderTestResult{ + Success: false, + Message: fmt.Sprintf("Connection failed: %s", err.Error()), + }) + return + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + response.RespondWithSuccess(c, http.StatusOK, llmprovider.ProviderTestResult{ + Success: false, + Message: fmt.Sprintf("Server returned HTTP %d", resp.StatusCode), + }) + return + } + + response.RespondWithSuccess(c, http.StatusOK, llmprovider.ProviderTestResult{ + Success: true, + Message: "Connection successful", + LatencyMs: latency, + }) +} + +// handleLLMGet returns the aggregated LLM configuration page data. +// GET /setting/llm +func handleLLMGet(c *gin.Context) { + info := authorized.GetInfo(c) + + if llmprovider.Global == nil { + respondError(c, http.StatusInternalServerError, "LLM provider registry not initialized") + return + } + + llmEnsureEncKey() + + owner := llmOwner(info) + filter := &llmprovider.ProviderFilter{ + Owner: owner, + Source: llmprovider.ProviderSourceAll, + } + providers, err := llmprovider.Global.List(filter) + if err != nil { + providers = []llmprovider.Provider{} + } + + enriched := make([]interface{}, 0, len(providers)) + for i := range providers { + enriched = append(enriched, enrichProvider(&providers[i])) + } + + var roles map[string]interface{} + if setting.Global != nil { + roles, _ = setting.Global.GetMerged(info.UserID, info.TeamID, llmRolesNS) + } + if roles == nil { + roles = make(map[string]interface{}) + } + + presetList := llmprovider.GetPresets() + presetIface := make([]interface{}, len(presetList)) + for i, p := range presetList { + raw, _ := json.Marshal(p) + var m map[string]interface{} + json.Unmarshal(raw, &m) + presetIface[i] = m + } + + response.RespondWithSuccess(c, http.StatusOK, LLMPageData{ + Providers: enriched, + Roles: roles, + PresetProviders: presetIface, + }) +} + +// handleLLMRoles saves the role assignment (default models). +// PUT /setting/llm/roles +func handleLLMRoles(c *gin.Context) { + if !guardOwner(c) { + return + } + info := authorized.GetInfo(c) + scope := llmScope(info) + + var body map[string]interface{} + if err := c.ShouldBindJSON(&body); err != nil { + respondError(c, http.StatusBadRequest, "invalid request body") + return + } + + if _, ok := body["default"]; !ok { + respondError(c, http.StatusBadRequest, "\"default\" role is required") + return + } + + if llmprovider.Global == nil { + respondError(c, http.StatusInternalServerError, "LLM provider registry not initialized") + return + } + + llmEnsureEncKey() + + for roleName, target := range body { + targetMap, ok := target.(map[string]interface{}) + if !ok { + respondError(c, http.StatusBadRequest, fmt.Sprintf("invalid target for role \"%s\"", roleName)) + return + } + + providerKey, _ := targetMap["provider"].(string) + modelID, _ := targetMap["model"].(string) + if providerKey == "" || modelID == "" { + respondError(c, http.StatusBadRequest, fmt.Sprintf("role \"%s\" requires provider and model", roleName)) + return + } + + p, err := llmprovider.Global.Get(providerKey) + if err != nil { + respondError(c, http.StatusBadRequest, fmt.Sprintf("provider \"%s\" not found", providerKey)) + return + } + if !p.Enabled { + respondError(c, http.StatusBadRequest, fmt.Sprintf("provider \"%s\" is not enabled", providerKey)) + return + } + if err := llmCheckOwnership(p, info); err != nil { + respondError(c, http.StatusBadRequest, fmt.Sprintf("provider \"%s\" not found", providerKey)) + return + } + + modelFound := false + for _, m := range p.Models { + if m.ID == modelID { + modelFound = true + break + } + } + if !modelFound { + respondError(c, http.StatusBadRequest, fmt.Sprintf("model \"%s\" not found in provider \"%s\"", modelID, providerKey)) + return + } + } + + if setting.Global == nil { + respondError(c, http.StatusInternalServerError, "setting registry not initialized") + return + } + + if _, err := setting.Global.Set(scope, llmRolesNS, body); err != nil { + respondError(c, http.StatusInternalServerError, err.Error()) + return + } + + response.RespondWithSuccess(c, http.StatusOK, body) +} + +// handleLLMProviderCreate creates a new LLM provider (preset or custom). +// POST /setting/llm/providers +func handleLLMProviderCreate(c *gin.Context) { + if !guardOwner(c) { + return + } + info := authorized.GetInfo(c) + + if llmprovider.Global == nil { + respondError(c, http.StatusInternalServerError, "LLM provider registry not initialized") + return + } + + llmEnsureEncKey() + + var body map[string]interface{} + if err := c.ShouldBindJSON(&body); err != nil { + respondError(c, http.StatusBadRequest, "invalid request body") + return + } + + var provider llmprovider.Provider + owner := llmOwner(info) + provider.Owner = *owner + provider.Source = llmprovider.ProviderSourceDynamic + provider.Enabled = true + + presetKey, _ := body["preset_key"].(string) + + if presetKey != "" { + preset := llmprovider.GetPreset(presetKey) + if preset == nil { + respondError(c, http.StatusBadRequest, fmt.Sprintf("unknown preset: %s", presetKey)) + return + } + + provider.Key = presetKey + provider.Name = preset.Name + provider.Type = preset.Type + provider.APIURL = preset.APIURL + provider.RequireKey = preset.RequireKey + provider.PresetKey = presetKey + + if v, ok := body["api_url"].(string); ok && v != "" { + provider.APIURL = v + } + if v, ok := body["api_key"].(string); ok && v != "" { + provider.APIKey = v + } + if v, ok := body["name"].(string); ok && v != "" { + provider.Name = v + } + + modelIDs, hasModelIDs := body["model_ids"].([]interface{}) + if hasModelIDs && len(modelIDs) > 0 { + idSet := make(map[string]bool, len(modelIDs)) + for _, id := range modelIDs { + if s, ok := id.(string); ok { + idSet[s] = true + } + } + for _, m := range preset.DefaultModels { + if idSet[m.ID] { + provider.Models = append(provider.Models, m) + } + } + } else { + provider.Models = make([]llmprovider.ModelInfo, len(preset.DefaultModels)) + copy(provider.Models, preset.DefaultModels) + } + } else { + provider.IsCustom = true + + key, _ := body["key"].(string) + if key == "" { + respondError(c, http.StatusBadRequest, "key is required for custom provider") + return + } + provider.Key = key + + name, _ := body["name"].(string) + if name == "" { + respondError(c, http.StatusBadRequest, "name is required") + return + } + provider.Name = name + + typ, _ := body["type"].(string) + if typ == "" { + typ = "openai" + } + provider.Type = typ + + provider.APIURL, _ = body["api_url"].(string) + provider.APIKey, _ = body["api_key"].(string) + + if modelsRaw, ok := body["models"]; ok { + raw, _ := json.Marshal(modelsRaw) + var models []llmprovider.ModelInfo + if err := json.Unmarshal(raw, &models); err == nil { + provider.Models = models + } + } + + if v, ok := body["require_key"].(bool); ok { + provider.RequireKey = v + } + } + + if provider.Models == nil { + provider.Models = []llmprovider.ModelInfo{} + } + + if provider.RequireKey && provider.APIKey != "" && provider.APIURL != "" { + if err := llmValidateKey(provider.Type, provider.APIURL, provider.APIKey); err != nil { + respondError(c, http.StatusBadRequest, fmt.Sprintf("API key validation failed: %s", err.Error())) + return + } + } + + created, err := llmprovider.Global.Create(&provider) + if err != nil { + if strings.Contains(err.Error(), "already exists") { + respondError(c, http.StatusConflict, err.Error()) + } else { + respondError(c, http.StatusInternalServerError, err.Error()) + } + return + } + + masked, err := llmprovider.Global.GetMasked(created.Key) + if err != nil { + created.APIKey = "" + response.RespondWithSuccess(c, http.StatusCreated, enrichProvider(created)) + return + } + response.RespondWithSuccess(c, http.StatusCreated, enrichProvider(masked)) +} + +// handleLLMProviderUpdate replaces a provider's configuration. +// Full replacement: api_key empty string preserves existing value. +// PUT /setting/llm/providers/:key +func handleLLMProviderUpdate(c *gin.Context) { + if !guardOwner(c) { + return + } + info := authorized.GetInfo(c) + key := c.Param("key") + + if llmprovider.Global == nil { + respondError(c, http.StatusInternalServerError, "LLM provider registry not initialized") + return + } + + llmEnsureEncKey() + + existing, err := llmprovider.Global.Get(key) + if err != nil { + respondError(c, http.StatusNotFound, fmt.Sprintf("provider \"%s\" not found", key)) + return + } + if err := llmCheckOwnership(existing, info); err != nil { + respondError(c, http.StatusNotFound, err.Error()) + return + } + + var body map[string]interface{} + if err := c.ShouldBindJSON(&body); err != nil { + respondError(c, http.StatusBadRequest, "invalid request body") + return + } + + var provider llmprovider.Provider + provider.Key = key + provider.Owner = existing.Owner + provider.Source = existing.Source + provider.ConnectorID = existing.ConnectorID + provider.PresetKey = existing.PresetKey + provider.IsCustom = existing.IsCustom + + if v, ok := body["name"].(string); ok { + provider.Name = v + } else { + provider.Name = existing.Name + } + if v, ok := body["type"].(string); ok { + provider.Type = v + } else { + provider.Type = existing.Type + } + if v, ok := body["api_url"].(string); ok { + provider.APIURL = v + } else { + provider.APIURL = existing.APIURL + } + + if v, ok := body["api_key"].(string); ok && v != "" { + provider.APIKey = v + } else { + provider.APIKey = existing.APIKey + } + + if v, ok := body["enabled"].(bool); ok { + provider.Enabled = v + } else { + provider.Enabled = existing.Enabled + } + if v, ok := body["require_key"].(bool); ok { + provider.RequireKey = v + } else { + provider.RequireKey = existing.RequireKey + } + if v, ok := body["status"].(string); ok { + provider.Status = v + } else { + provider.Status = existing.Status + } + + if modelsRaw, ok := body["models"]; ok { + raw, _ := json.Marshal(modelsRaw) + var models []llmprovider.ModelInfo + if err := json.Unmarshal(raw, &models); err == nil { + provider.Models = models + } + } else { + provider.Models = existing.Models + } + if provider.Models == nil { + provider.Models = []llmprovider.ModelInfo{} + } + + if _, err = llmprovider.Global.Update(key, &provider); err != nil { + respondError(c, http.StatusInternalServerError, err.Error()) + return + } + + masked, err := llmprovider.Global.GetMasked(key) + if err != nil { + provider.APIKey = "" + response.RespondWithSuccess(c, http.StatusOK, enrichProvider(&provider)) + return + } + response.RespondWithSuccess(c, http.StatusOK, enrichProvider(masked)) +} + +// handleLLMProviderDelete removes a provider and cleans up role references. +// DELETE /setting/llm/providers/:key +func handleLLMProviderDelete(c *gin.Context) { + if !guardOwner(c) { + return + } + info := authorized.GetInfo(c) + key := c.Param("key") + + if llmprovider.Global == nil { + respondError(c, http.StatusInternalServerError, "LLM provider registry not initialized") + return + } + + llmEnsureEncKey() + + existing, err := llmprovider.Global.Get(key) + if err != nil { + respondError(c, http.StatusNotFound, fmt.Sprintf("provider \"%s\" not found", key)) + return + } + if err := llmCheckOwnership(existing, info); err != nil { + respondError(c, http.StatusNotFound, err.Error()) + return + } + + var warning string + if setting.Global != nil { + scope := llmScope(info) + roles, _ := setting.Global.Get(scope, llmRolesNS) + if roles != nil { + cleaned := false + for roleName, target := range roles { + if targetMap, ok := target.(map[string]interface{}); ok { + if provKey, _ := targetMap["provider"].(string); provKey == key { + delete(roles, roleName) + cleaned = true + } + } + } + if cleaned { + setting.Global.Set(scope, llmRolesNS, roles) + warning = fmt.Sprintf("roles referencing provider \"%s\" have been cleared", key) + } + } + } + + if err := llmprovider.Global.Delete(key); err != nil { + respondError(c, http.StatusInternalServerError, err.Error()) + return + } + + result := map[string]interface{}{"success": true} + if warning != "" { + result["warning"] = warning + } + response.RespondWithSuccess(c, http.StatusOK, result) +} + +// handleLLMProviderTest tests connectivity for a provider and writes back status. +// POST /setting/llm/providers/:key/test +func handleLLMProviderTest(c *gin.Context) { + if !guardOwner(c) { + return + } + info := authorized.GetInfo(c) + key := c.Param("key") + + if llmprovider.Global == nil { + respondError(c, http.StatusInternalServerError, "LLM provider registry not initialized") + return + } + + llmEnsureEncKey() + + p, err := llmprovider.Global.Get(key) + if err != nil { + respondError(c, http.StatusNotFound, fmt.Sprintf("provider \"%s\" not found", key)) + return + } + if err := llmCheckOwnership(p, info); err != nil { + respondError(c, http.StatusNotFound, err.Error()) + return + } + + start := time.Now() + err = llmValidateKey(p.Type, p.APIURL, p.APIKey) + latency := time.Since(start).Milliseconds() + + var testResult llmprovider.ProviderTestResult + if err != nil { + testResult = llmprovider.ProviderTestResult{ + Success: false, + Message: err.Error(), + } + p.Status = "disconnected" + } else { + testResult = llmprovider.ProviderTestResult{ + Success: true, + Message: "Connection successful", + LatencyMs: latency, + } + p.Status = "connected" + llmprovider.Global.Update(key, p) + } + + response.RespondWithSuccess(c, http.StatusOK, testResult) +} diff --git a/openapi/setting/setting.go b/openapi/setting/setting.go index 094490f9..a4a1c7c7 100644 --- a/openapi/setting/setting.go +++ b/openapi/setting/setting.go @@ -36,6 +36,15 @@ func Attach(group *gin.RouterGroup, oauth oauthTypes.OAuth) { cloud.GET("", handleCloudGet) cloud.PUT("", handleCloudUpdate) cloud.POST("/test", handleCloudTest) + + llm := group.Group("/llm") + llm.GET("", handleLLMGet) + llm.PUT("/roles", handleLLMRoles) + llm.POST("/test", handleLLMTest) + llm.POST("/providers", handleLLMProviderCreate) + llm.PUT("/providers/:key", handleLLMProviderUpdate) + llm.DELETE("/providers/:key", handleLLMProviderDelete) + llm.POST("/providers/:key/test", handleLLMProviderTest) } // requireOwner checks that the current user is the team owner. diff --git a/openapi/setting/types.go b/openapi/setting/types.go index 80f288a7..fef5b56b 100644 --- a/openapi/setting/types.go +++ b/openapi/setting/types.go @@ -81,3 +81,14 @@ type CloudTestResult struct { Message string `json:"message"` LatencyMs int64 `json:"latency_ms,omitempty"` } + +// --------------------------------------------------------------------------- +// LLM Providers +// --------------------------------------------------------------------------- + +// LLMPageData is the aggregated response for GET /setting/llm. +type LLMPageData struct { + Providers []interface{} `json:"providers"` + Roles map[string]interface{} `json:"roles"` + PresetProviders []interface{} `json:"preset_providers"` +} diff --git a/openapi/tests/setting/llm_test.go b/openapi/tests/setting/llm_test.go new file mode 100644 index 00000000..993b3cfe --- /dev/null +++ b/openapi/tests/setting/llm_test.go @@ -0,0 +1,453 @@ +package setting_test + +import ( + "bytes" + "encoding/json" + "io" + "net/http" + "os" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/config" + "github.com/yaoapp/yao/llmprovider" + "github.com/yaoapp/yao/openapi/tests/testutils" +) + +func requireOpenAIKey(t *testing.T) string { + t.Helper() + key := os.Getenv("OPENAI_TEST_KEY") + if key == "" { + t.Skip("OPENAI_TEST_KEY not set") + } + return key +} + +func initLLMRegistry(t *testing.T) { + t.Helper() + if err := llmprovider.Init(); err != nil { + t.Fatalf("llmprovider.Init: %v", err) + } + if config.Conf.DB.AESKey != "" { + llmprovider.Global.SetEncryptionKey(config.Conf.DB.AESKey) + } +} + +func llmURL(serverURL, path string) string { + return serverURL + baseURL() + "/setting/llm" + path +} + +func llmGet(t *testing.T, url, token string) *http.Response { + t.Helper() + req, _ := http.NewRequest("GET", url, nil) + req.Header.Set("Authorization", "Bearer "+token) + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + return resp +} + +func llmPost(t *testing.T, url, token string, payload interface{}) *http.Response { + t.Helper() + var body io.Reader + if payload != nil { + raw, _ := json.Marshal(payload) + body = bytes.NewReader(raw) + } + req, _ := http.NewRequest("POST", url, body) + req.Header.Set("Authorization", "Bearer "+token) + if payload != nil { + req.Header.Set("Content-Type", "application/json") + } + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + return resp +} + +func llmPut(t *testing.T, url, token string, payload interface{}) *http.Response { + t.Helper() + raw, _ := json.Marshal(payload) + req, _ := http.NewRequest("PUT", url, bytes.NewReader(raw)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + return resp +} + +func llmDelete(t *testing.T, url, token string) *http.Response { + t.Helper() + req, _ := http.NewRequest("DELETE", url, nil) + req.Header.Set("Authorization", "Bearer "+token) + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + return resp +} + +func llmBody(t *testing.T, resp *http.Response) map[string]interface{} { + t.Helper() + var body map[string]interface{} + json.NewDecoder(resp.Body).Decode(&body) + return body +} + +func createTestOpenAI(t *testing.T, serverURL, token string) { + t.Helper() + apiKey := requireOpenAIKey(t) + llmprovider.Global.Delete("openai") + payload := map[string]interface{}{ + "preset_key": "openai", + "api_key": apiKey, + "model_ids": []string{"gpt-4o", "gpt-4o-mini"}, + } + resp := llmPost(t, llmURL(serverURL, "/providers"), token, payload) + resp.Body.Close() + assert.Equal(t, http.StatusCreated, resp.StatusCode, "createTestOpenAI should succeed") + t.Cleanup(func() { llmprovider.Global.Delete("openai") }) +} + +// ----------- Functional tests ----------- + +func TestLLMGetPageData(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + token := obtainToken(t, serverURL) + + createTestOpenAI(t, serverURL, token) + + resp := llmGet(t, llmURL(serverURL, ""), token) + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) + + body := llmBody(t, resp) + assert.Contains(t, body, "providers") + assert.Contains(t, body, "roles") + assert.Contains(t, body, "preset_providers") + + providers, ok := body["providers"].([]interface{}) + assert.True(t, ok, "providers should be an array") + assert.GreaterOrEqual(t, len(providers), 1) + + if len(providers) > 0 { + p := providers[0].(map[string]interface{}) + assert.Contains(t, p, "key") + assert.Contains(t, p, "name") + assert.Contains(t, p, "models") + assert.NotContains(t, p, "connector_id", "internal field should be stripped") + assert.NotContains(t, p, "source", "internal field should be stripped") + assert.NotContains(t, p, "owner", "internal field should be stripped") + } + + presets, ok := body["preset_providers"].([]interface{}) + assert.True(t, ok) + assert.Equal(t, 5, len(presets), "should have 5 presets") +} + +func TestLLMGetUnauthenticated(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + req, _ := http.NewRequest("GET", llmURL(serverURL, ""), nil) + resp, err := http.DefaultClient.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) +} + +func TestLLMProviderCreate(t *testing.T) { + realKey := requireOpenAIKey(t) + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + token := obtainToken(t, serverURL) + llmprovider.Global.Delete("openai") + + payload := map[string]interface{}{ + "preset_key": "openai", + "api_key": realKey, + "model_ids": []string{"gpt-4o"}, + } + resp := llmPost(t, llmURL(serverURL, "/providers"), token, payload) + defer resp.Body.Close() + assert.Equal(t, http.StatusCreated, resp.StatusCode) + t.Cleanup(func() { llmprovider.Global.Delete("openai") }) + + body := llmBody(t, resp) + assert.Equal(t, "openai", body["key"]) + assert.Equal(t, "OpenAI", body["name"]) + assert.Equal(t, "openai", body["type"]) + + apiKey, _ := body["api_key"].(string) + assert.NotEqual(t, realKey, apiKey, "API key should be masked") + assert.NotEmpty(t, apiKey) + + models, _ := body["models"].([]interface{}) + assert.Equal(t, 1, len(models)) +} + +func TestLLMProviderCreateCustom(t *testing.T) { + realKey := requireOpenAIKey(t) + mirror := os.Getenv("TEST_MOAPI_MIRROR") + if mirror == "" { + mirror = "https://api.openai.com" + } + + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + token := obtainToken(t, serverURL) + llmprovider.Global.Delete("my-custom-llm") + + payload := map[string]interface{}{ + "key": "my-custom-llm", + "name": "My Custom LLM", + "type": "openai", + "api_url": mirror, + "api_key": realKey, + "models": []map[string]interface{}{ + {"id": "custom-model", "name": "Custom Model", "capabilities": []string{"streaming"}, "enabled": true}, + }, + "require_key": true, + } + resp := llmPost(t, llmURL(serverURL, "/providers"), token, payload) + defer resp.Body.Close() + assert.Equal(t, http.StatusCreated, resp.StatusCode) + t.Cleanup(func() { llmprovider.Global.Delete("my-custom-llm") }) + + body := llmBody(t, resp) + assert.Equal(t, "my-custom-llm", body["key"]) + assert.Equal(t, "My Custom LLM", body["name"]) + assert.Equal(t, true, body["is_custom"]) + + models, _ := body["models"].([]interface{}) + assert.Equal(t, 1, len(models)) +} + +func TestLLMProviderUpdate(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + token := obtainToken(t, serverURL) + + createTestOpenAI(t, serverURL, token) + + updatePayload := map[string]interface{}{ + "name": "Updated OpenAI", + "api_url": "https://api.openai.com/v2", + "models": []map[string]interface{}{ + {"id": "gpt-4o", "name": "GPT-4o Updated", "capabilities": []string{"vision", "tool_calls"}, "enabled": true}, + }, + } + resp := llmPut(t, llmURL(serverURL, "/providers/openai"), token, updatePayload) + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) + + body := llmBody(t, resp) + assert.Equal(t, "Updated OpenAI", body["name"]) + assert.Equal(t, "https://api.openai.com/v2", body["api_url"]) + + apiKey, _ := body["api_key"].(string) + assert.NotEmpty(t, apiKey, "API key should be preserved when not sent") +} + +func TestLLMProviderDelete(t *testing.T) { + anthropicKey := os.Getenv("ANTHROPIC_API_KEY") + if anthropicKey == "" { + t.Skip("ANTHROPIC_API_KEY not set") + } + + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + token := obtainToken(t, serverURL) + llmprovider.Global.Delete("anthropic") + + createPayload := map[string]interface{}{ + "preset_key": "anthropic", + "api_key": anthropicKey, + } + createResp := llmPost(t, llmURL(serverURL, "/providers"), token, createPayload) + createResp.Body.Close() + assert.Equal(t, http.StatusCreated, createResp.StatusCode) + + rolesPayload := map[string]interface{}{ + "default": map[string]interface{}{ + "provider": "anthropic", + "model": "claude-sonnet-4-20250514", + }, + } + rolesResp := llmPut(t, llmURL(serverURL, "/roles"), token, rolesPayload) + rolesResp.Body.Close() + assert.Equal(t, http.StatusOK, rolesResp.StatusCode) + + deleteResp := llmDelete(t, llmURL(serverURL, "/providers/anthropic"), token) + defer deleteResp.Body.Close() + assert.Equal(t, http.StatusOK, deleteResp.StatusCode) + + body := llmBody(t, deleteResp) + assert.Equal(t, true, body["success"]) + assert.NotEmpty(t, body["warning"], "should warn about cleared roles") + + getResp := llmGet(t, llmURL(serverURL, ""), token) + defer getResp.Body.Close() + getBody := llmBody(t, getResp) + roles, _ := getBody["roles"].(map[string]interface{}) + assert.NotContains(t, roles, "default", "role referencing deleted provider should be cleared") +} + +func TestLLMProviderDeleteForbidden(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + token := obtainToken(t, serverURL) + + llmprovider.Global.Delete("other-team-provider") + otherProvider := &llmprovider.Provider{ + Key: "other-team-provider", + Name: "Other Team's Provider", + Type: "openai", + APIURL: "https://api.example.com", + Models: []llmprovider.ModelInfo{}, + Enabled: true, + Source: llmprovider.ProviderSourceDynamic, + Owner: llmprovider.ProviderOwner{Type: "user", UserID: "some-other-user-999"}, + } + llmprovider.Global.Create(otherProvider) + t.Cleanup(func() { llmprovider.Global.Delete("other-team-provider") }) + + resp := llmDelete(t, llmURL(serverURL, "/providers/other-team-provider"), token) + defer resp.Body.Close() + assert.Equal(t, http.StatusNotFound, resp.StatusCode, "should not be able to delete another user's provider") +} + +func TestLLMProviderTest(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + token := obtainToken(t, serverURL) + + createTestOpenAI(t, serverURL, token) + + resp := llmPost(t, llmURL(serverURL, "/providers/openai/test"), token, nil) + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) + + body := llmBody(t, resp) + assert.Equal(t, true, body["success"]) + assert.NotEmpty(t, body["message"]) + assert.NotNil(t, body["latency_ms"]) +} + +func TestLLMRoles(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + token := obtainToken(t, serverURL) + createTestOpenAI(t, serverURL, token) + + rolesPayload := map[string]interface{}{ + "default": map[string]interface{}{ + "provider": "openai", + "model": "gpt-4o", + }, + "vision": map[string]interface{}{ + "provider": "openai", + "model": "gpt-4o", + }, + } + resp := llmPut(t, llmURL(serverURL, "/roles"), token, rolesPayload) + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) + + body := llmBody(t, resp) + assert.Contains(t, body, "default") + assert.Contains(t, body, "vision") + + getResp := llmGet(t, llmURL(serverURL, ""), token) + defer getResp.Body.Close() + getBody := llmBody(t, getResp) + roles, _ := getBody["roles"].(map[string]interface{}) + assert.Contains(t, roles, "default") + assert.Contains(t, roles, "vision") +} + +func TestLLMRolesValidation(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + token := obtainToken(t, serverURL) + + resp1 := llmPut(t, llmURL(serverURL, "/roles"), token, map[string]interface{}{ + "vision": map[string]interface{}{ + "provider": "openai", + "model": "gpt-4o", + }, + }) + defer resp1.Body.Close() + assert.Equal(t, http.StatusBadRequest, resp1.StatusCode, "should require 'default' role") + + resp2 := llmPut(t, llmURL(serverURL, "/roles"), token, map[string]interface{}{ + "default": map[string]interface{}{ + "provider": "nonexistent-provider", + "model": "some-model", + }, + }) + defer resp2.Body.Close() + assert.Equal(t, http.StatusBadRequest, resp2.StatusCode, "should reject non-existent provider") + + createTestOpenAI(t, serverURL, token) + + resp3 := llmPut(t, llmURL(serverURL, "/roles"), token, map[string]interface{}{ + "default": map[string]interface{}{ + "provider": "openai", + "model": "nonexistent-model", + }, + }) + defer resp3.Body.Close() + assert.Equal(t, http.StatusBadRequest, resp3.StatusCode, "should reject non-existent model") +} + +// ----------- ACL permission tests ----------- + +func TestLLMACL_ReadOnlyScopeCannotWrite(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + + readToken := obtainRestrictedToken(t, serverURL, "setting:llm:read:all") + + resp := llmGet(t, llmURL(serverURL, ""), readToken) + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode, "read-only scope should allow GET") + + createPayload := map[string]interface{}{ + "preset_key": "openai", + "api_key": "sk-test", + } + resp2 := llmPost(t, llmURL(serverURL, "/providers"), readToken, createPayload) + defer resp2.Body.Close() + assert.Equal(t, http.StatusForbidden, resp2.StatusCode, "read-only scope should deny POST") +} + +func TestLLMACL_NoScopeCannotAccess(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + initSettingRegistry(t) + initLLMRegistry(t) + + noSettingToken := obtainRestrictedToken(t, serverURL, "kb:collections:read:all") + + resp := llmGet(t, llmURL(serverURL, ""), noSettingToken) + defer resp.Body.Close() + assert.Equal(t, http.StatusForbidden, resp.StatusCode, "token without setting scope should be denied") +}