From cb117f156965d97b223fc19ce7cbb1c09a27f846 Mon Sep 17 00:00:00 2001 From: Reidho Satria Date: Tue, 15 Sep 2026 12:28:01 +0700 Subject: [PATCH 1/9] fix(setup): persist custom providers under their named id --- cmd/antares/setup.go | 2 +- cmd/antares/setup_test.go | 102 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 103 insertions(+), 1 deletion(-) create mode 100644 cmd/antares/setup_test.go diff --git a/cmd/antares/setup.go b/cmd/antares/setup.go index 0f41577..bd4f949 100644 --- a/cmd/antares/setup.go +++ b/cmd/antares/setup.go @@ -324,7 +324,7 @@ func runTerminalSetup(ctx context.Context, rt *runtimeServices) error { entry.APIKey = key } } - cfg.Providers[chosen.id] = entry + cfg.Providers[cfg.Model.Provider] = entry // 4. Model, verified against the provider when possible fmt.Println() diff --git a/cmd/antares/setup_test.go b/cmd/antares/setup_test.go new file mode 100644 index 0000000..70e201b --- /dev/null +++ b/cmd/antares/setup_test.go @@ -0,0 +1,102 @@ +package main + +import ( + "bufio" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + + "github.com/enowdev/antares/internal/agent" + "github.com/enowdev/antares/internal/config" +) + +func TestTerminalSetupPersistsNamedProvider(t *testing.T) { + t.Setenv("ANTARES_HOME", t.TempDir()) + + fixture := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/v1/models": + _ = json.NewEncoder(w).Encode(map[string]any{ + "object": "list", + "data": []map[string]string{{"id": "header-model", "owned_by": "test"}}, + }) + case "/v1/chat/completions": + _ = json.NewEncoder(w).Encode(map[string]any{ + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "header-model", + "choices": []map[string]any{{ + "index": 0, + "message": map[string]string{"role": "assistant", "content": "pong"}, + "finish_reason": "stop", + }}, + }) + default: + http.NotFound(w, r) + } + })) + t.Cleanup(fixture.Close) + + cfg := config.Default() + cfg.Agent.Workspace = filepath.Join(t.TempDir(), "workspace") + cfg.Model.MaxRetries = -1 + cfg.Providers["custom"] = config.Provider{ + Enabled: true, + Kind: "openai-compatible", + BaseURL: "https://legacy.example/v1", + Label: "Legacy custom", + } + if err := config.Save(cfg); err != nil { + t.Fatal(err) + } + + a := agent.New(cfg, nil, nil, nil, nil) + rt := &runtimeServices{cfg: cfg, agent: a} + oldReader := stdinReader + stdinReader = bufio.NewReader(strings.NewReader(strings.Join([]string{ + "7", // Custom provider + "Named Provider", // provider name + fixture.URL + "/v1", // endpoint + "1", // first live model or manual fallback + "", // workspace + "", // PostgreSQL + "", // RAG + "", // Telegram + "", // dashboard password + }, "\n") + "\n")) + t.Cleanup(func() { stdinReader = oldReader }) + + var setupErr error + output := captureProviderStdout(t, func() { + setupErr = runTerminalSetup(context.Background(), rt) + }) + if setupErr != nil { + t.Fatalf("run terminal setup: %v\noutput:\n%s", setupErr, output) + } + after, err := config.Reload() + if err != nil { + t.Fatal(err) + } + if after.Model.Provider != "named-provider" { + t.Fatalf("model provider = %q, want named-provider", after.Model.Provider) + } + got, ok := after.Providers[after.Model.Provider] + if !ok { + t.Fatalf("named provider %q was not saved", after.Model.Provider) + } + if got.BaseURL != fixture.URL+"/v1" { + t.Fatalf("named provider endpoint = %q, want %q", got.BaseURL, fixture.URL+"/v1") + } + legacy := after.Providers["custom"] + if legacy.BaseURL != "https://legacy.example/v1" || legacy.Label != "Legacy custom" { + t.Fatalf("legacy custom provider changed: %#v", legacy) + } + if _, resolved := after.ResolveProvider(after.Model.Provider); resolved.BaseURL != fixture.URL+"/v1" { + t.Fatalf("resolved provider endpoint = %q, want %q", resolved.BaseURL, fixture.URL+"/v1") + } +} From 06d761204ad8b3939057786202eb56f7493a187c Mon Sep 17 00:00:00 2001 From: Reidho Satria Date: Tue, 15 Sep 2026 12:46:39 +0700 Subject: [PATCH 2/9] feat(providers): accept and retain custom request headers --- internal/config/provider_headers.go | 44 ++++++ internal/config/provider_headers_test.go | 52 +++++++ internal/server/handlers_config.go | 36 +++-- internal/server/handlers_config_test.go | 82 +++++++++-- internal/server/handlers_providers.go | 29 ++-- internal/server/handlers_setup.go | 46 ++++-- internal/server/provider_headers_test.go | 179 +++++++++++++++++++++++ 7 files changed, 426 insertions(+), 42 deletions(-) create mode 100644 internal/config/provider_headers.go create mode 100644 internal/config/provider_headers_test.go create mode 100644 internal/server/provider_headers_test.go diff --git a/internal/config/provider_headers.go b/internal/config/provider_headers.go new file mode 100644 index 0000000..bb8d5d0 --- /dev/null +++ b/internal/config/provider_headers.go @@ -0,0 +1,44 @@ +package config + +import ( + "errors" + "net/http" + "strings" + + "golang.org/x/net/http/httpguts" +) + +var ( + ErrInvalidHeaderName = errors.New("invalid header name") + ErrInvalidHeaderValue = errors.New("invalid header value") + ErrDuplicateHeaderName = errors.New("duplicate header name") +) + +// NormalizeProviderHeaders validates an HTTP header map, canonicalizes names, +// and returns an independent map. A nil map remains nil; an allocated empty map +// remains allocated so callers can distinguish an omitted value from a clear. +func NormalizeProviderHeaders(headers map[string]string) (map[string]string, error) { + if headers == nil { + return nil, nil + } + + normalized := make(map[string]string, len(headers)) + seen := make(map[string]struct{}, len(headers)) + for name, value := range headers { + if !httpguts.ValidHeaderFieldValue(value) { + return nil, ErrInvalidHeaderValue + } + name = strings.Trim(name, " \t") + if name == "" || !httpguts.ValidHeaderFieldName(name) { + return nil, ErrInvalidHeaderName + } + canonical := http.CanonicalHeaderKey(name) + folded := strings.ToLower(canonical) + if _, duplicate := seen[folded]; duplicate { + return nil, ErrDuplicateHeaderName + } + seen[folded] = struct{}{} + normalized[canonical] = strings.Trim(value, " \t") + } + return normalized, nil +} diff --git a/internal/config/provider_headers_test.go b/internal/config/provider_headers_test.go new file mode 100644 index 0000000..30771d2 --- /dev/null +++ b/internal/config/provider_headers_test.go @@ -0,0 +1,52 @@ +package config + +import ( + "errors" + "reflect" + "testing" +) + +func TestNormalizeProviderHeaders(t *testing.T) { + input := map[string]string{" x-tenant ": " team=a=b ", "Empty": ""} + got, err := NormalizeProviderHeaders(input) + if err != nil { + t.Fatal(err) + } + want := map[string]string{"X-Tenant": "team=a=b", "Empty": ""} + if !reflect.DeepEqual(got, want) { + t.Fatalf("normalized = %#v, want %#v", got, want) + } + got["X-Tenant"] = "changed" + if input[" x-tenant "] != " team=a=b " { + t.Fatalf("normalization mutated input: %#v", input) + } + + nilHeaders, err := NormalizeProviderHeaders(nil) + if err != nil || nilHeaders != nil { + t.Fatalf("nil headers = %#v, %v; want nil, nil", nilHeaders, err) + } + emptyHeaders, err := NormalizeProviderHeaders(map[string]string{}) + if err != nil || emptyHeaders == nil || len(emptyHeaders) != 0 { + t.Fatalf("empty headers = %#v, %v; want allocated empty map", emptyHeaders, err) + } +} + +func TestNormalizeProviderHeadersRejectsMalformedInput(t *testing.T) { + cases := []struct { + name string + headers map[string]string + want error + }{ + {"invalid name", map[string]string{"bad header": "value"}, ErrInvalidHeaderName}, + {"control value", map[string]string{"X-Test": "bad\nvalue"}, ErrInvalidHeaderValue}, + {"duplicate casing", map[string]string{"X-Test": "one", "x-test": "two"}, ErrDuplicateHeaderName}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := NormalizeProviderHeaders(tc.headers) + if !errors.Is(err, tc.want) { + t.Fatalf("error = %v, want %v", err, tc.want) + } + }) + } +} diff --git a/internal/server/handlers_config.go b/internal/server/handlers_config.go index 6fb6ffe..bd37b2e 100644 --- a/internal/server/handlers_config.go +++ b/internal/server/handlers_config.go @@ -156,15 +156,16 @@ func (s *Server) handleModelOptions(w http.ResponseWriter, r *http.Request) { BaseURL string `json:"base_url"` Active bool `json:"active"` // Setup metadata, so the connect form can render the right fields. - Hint string `json:"hint,omitempty"` - KeyHint string `json:"key_hint,omitempty"` - KeyURL string `json:"key_url,omitempty"` - KeyLabel string `json:"key_label,omitempty"` - Note string `json:"note,omitempty"` - NeedsRegion bool `json:"needs_region,omitempty"` - NeedsAPIVersion bool `json:"needs_api_version,omitempty"` - NeedsBaseURL bool `json:"needs_base_url,omitempty"` - TimeoutSecs int `json:"timeout_seconds,omitempty"` + Hint string `json:"hint,omitempty"` + KeyHint string `json:"key_hint,omitempty"` + KeyURL string `json:"key_url,omitempty"` + KeyLabel string `json:"key_label,omitempty"` + Note string `json:"note,omitempty"` + NeedsRegion bool `json:"needs_region,omitempty"` + NeedsAPIVersion bool `json:"needs_api_version,omitempty"` + NeedsBaseURL bool `json:"needs_base_url,omitempty"` + TimeoutSecs int `json:"timeout_seconds,omitempty"` + Headers map[string]string `json:"headers,omitempty"` // Custom marks a user-defined provider. Customs always group under // "API key" — even a localhost endpoint is a configured service, not // one of the built-in local runtimes. @@ -177,6 +178,7 @@ func (s *Server) handleModelOptions(w http.ResponseWriter, r *http.Request) { // installations may still use that id for their real custom provider. seen := map[string]bool{} providerList := make([]providerInfo, 0) + exposeHeaders := cfg.Server.DashboardLocked() || s.bearerAuthorizedOrQuery(r) for _, sp := range setupProviderCatalogue(cfg) { p := cfg.Providers[sp.ID] if sp.Custom && !legacyCustomProviderInUse(cfg, p) { @@ -188,14 +190,18 @@ func (s *Server) handleModelOptions(w http.ResponseWriter, r *http.Request) { label = firstNonEmpty(p.Label, sp.Label) kind = firstNonEmpty(p.Kind, sp.Kind) } - providerList = append(providerList, providerInfo{ + info := providerInfo{ ID: sp.ID, Label: label, Kind: kind, Enabled: p.Enabled, HasKey: p.APIKey != "", Local: sp.Local, BaseURL: firstNonEmpty(p.BaseURL, sp.BaseURL), Active: sp.ID == cfg.Model.Provider, Hint: sp.Hint, KeyHint: sp.KeyHint, KeyURL: sp.KeyURL, KeyLabel: sp.KeyLabel, Note: sp.Note, NeedsRegion: sp.NeedsRegion, NeedsAPIVersion: sp.NeedsAPIVersion, NeedsBaseURL: sp.NeedsBaseURL, TimeoutSecs: p.TimeoutSecs, Custom: sp.Custom, - }) + } + if sp.Custom && exposeHeaders { + info.Headers = p.Headers + } + providerList = append(providerList, info) seen[sp.ID] = true } names := make([]string, 0, len(cfg.Providers)) @@ -207,12 +213,16 @@ func (s *Server) handleModelOptions(w http.ResponseWriter, r *http.Request) { sort.Strings(names) for _, name := range names { p := cfg.Providers[name] - providerList = append(providerList, providerInfo{ + info := providerInfo{ ID: name, Label: firstNonEmpty(p.Label, name), Kind: p.Kind, Enabled: p.Enabled, HasKey: p.APIKey != "", BaseURL: p.BaseURL, Active: name == cfg.Model.Provider, TimeoutSecs: p.TimeoutSecs, Custom: true, NeedsBaseURL: true, - }) + } + if exposeHeaders { + info.Headers = p.Headers + } + providerList = append(providerList, info) } writeJSON(w, http.StatusOK, map[string]any{ diff --git a/internal/server/handlers_config_test.go b/internal/server/handlers_config_test.go index 5d3db93..3333524 100644 --- a/internal/server/handlers_config_test.go +++ b/internal/server/handlers_config_test.go @@ -4,21 +4,23 @@ import ( "encoding/json" "net/http" "net/http/httptest" + "strings" "testing" "github.com/enowdev/antares/internal/config" ) type modelOptionsProvider struct { - ID string `json:"id"` - Label string `json:"label"` - Kind string `json:"kind"` - Enabled bool `json:"enabled"` - HasKey bool `json:"has_key"` - BaseURL string `json:"base_url"` - Active bool `json:"active"` - Custom bool `json:"custom"` - NeedsBaseURL bool `json:"needs_base_url"` + ID string `json:"id"` + Label string `json:"label"` + Kind string `json:"kind"` + Enabled bool `json:"enabled"` + HasKey bool `json:"has_key"` + BaseURL string `json:"base_url"` + Active bool `json:"active"` + Custom bool `json:"custom"` + NeedsBaseURL bool `json:"needs_base_url"` + Headers map[string]string `json:"headers"` } type modelOptionsResponse struct { @@ -90,3 +92,65 @@ func TestModelOptionsHidesUnusedLegacyCustomPlaceholder(t *testing.T) { t.Fatalf("unused legacy custom provider cards = %d, want 0", count) } } + +func TestModelOptionsProtectsCustomHeaderValues(t *testing.T) { + cfg := config.Default() + cfg.Providers["custom"] = config.Provider{ + Kind: "openai-compatible", BaseURL: "https://legacy.example/v1", Enabled: true, + Headers: map[string]string{"X-Legacy": "legacy-secret"}, + } + cfg.Providers["named"] = config.Provider{ + Kind: "openai-compatible", BaseURL: "https://named.example/v1", Enabled: true, + Headers: map[string]string{"X-Named": "named-secret"}, + } + cfg.Providers["openai"] = config.Provider{Headers: map[string]string{"X-Builtin": "must-not-leak"}} + + for _, provider := range modelOptions(t, cfg).Providers { + if len(provider.Headers) != 0 { + t.Fatalf("unprotected response exposed headers for %s: %#v", provider.ID, provider.Headers) + } + } + + cfg.Server.DashboardPasswordHash = "test-hash" + cfg.Server.AuthToken = "header-test-token" + s := New(Options{Config: cfg}) + unauthorized := httptest.NewRecorder() + s.Handler().ServeHTTP(unauthorized, httptest.NewRequest(http.MethodGet, "/api/model/options", nil)) + if unauthorized.Code != http.StatusUnauthorized { + t.Fatalf("locked unauthenticated status = %d, want 401", unauthorized.Code) + } + + authorizedRequest := httptest.NewRequest(http.MethodGet, "/api/model/options", nil) + authorizedRequest.Header.Set("Authorization", "Bearer header-test-token") + authorized := httptest.NewRecorder() + s.Handler().ServeHTTP(authorized, authorizedRequest) + if authorized.Code != http.StatusOK { + t.Fatalf("authorized status = %d, body = %s", authorized.Code, authorized.Body.String()) + } + var response modelOptionsResponse + if err := json.Unmarshal(authorized.Body.Bytes(), &response); err != nil { + t.Fatal(err) + } + for _, provider := range response.Providers { + switch provider.ID { + case "custom": + if provider.Headers["X-Legacy"] != "legacy-secret" { + t.Fatalf("legacy headers = %#v", provider.Headers) + } + case "named": + if provider.Headers["X-Named"] != "named-secret" { + t.Fatalf("named headers = %#v", provider.Headers) + } + case "openai": + if len(provider.Headers) != 0 { + t.Fatalf("built-in headers leaked: %#v", provider.Headers) + } + } + } + + setupStatus := httptest.NewRecorder() + s.Handler().ServeHTTP(setupStatus, httptest.NewRequest(http.MethodGet, "/api/setup/status", nil)) + if strings.Contains(setupStatus.Body.String(), "legacy-secret") || strings.Contains(setupStatus.Body.String(), "named-secret") { + t.Fatalf("setup status disclosed headers: %s", setupStatus.Body.String()) + } +} diff --git a/internal/server/handlers_providers.go b/internal/server/handlers_providers.go index 51e7b9b..535655d 100644 --- a/internal/server/handlers_providers.go +++ b/internal/server/handlers_providers.go @@ -188,6 +188,14 @@ func (s *Server) handleProviderSettings(w http.ResponseWriter, r *http.Request) return } p := cfg.Providers[id] + if body.Headers != nil { + headers, err := config.NormalizeProviderHeaders(body.Headers) + if err != nil { + writeError(w, http.StatusBadRequest, err) + return + } + p.Headers = headers + } // Custom providers (user-named entries, plus the legacy "custom" slot) may // point at loopback or LAN addresses; built-ins keep their catalogue rule. sp := lookupSetupProvider(cfg, id) @@ -211,9 +219,6 @@ func (s *Server) handleProviderSettings(w http.ResponseWriter, r *http.Request) if body.TimeoutSecs != nil { p.TimeoutSecs = *body.TimeoutSecs } - if body.Headers != nil { - p.Headers = body.Headers - } cfg.Providers[id] = p if err := config.Save(cfg); err != nil { @@ -236,14 +241,20 @@ func (s *Server) handleCreateProvider(w http.ResponseWriter, r *http.Request) { return } var body struct { - Name string `json:"name"` - BaseURL string `json:"base_url"` - APIKey string `json:"api_key"` + Name string `json:"name"` + BaseURL string `json:"base_url"` + APIKey string `json:"api_key"` + Headers map[string]string `json:"headers"` } if err := decodeBody(r, &body); err != nil { writeError(w, http.StatusBadRequest, err) return } + headers, err := config.NormalizeProviderHeaders(body.Headers) + if err != nil { + writeError(w, http.StatusBadRequest, err) + return + } name := strings.TrimSpace(body.Name) if name == "" { writeError(w, http.StatusBadRequest, errors.New("a name is required")) @@ -269,9 +280,9 @@ func (s *Server) handleCreateProvider(w http.ResponseWriter, r *http.Request) { // Verify the pair now so a bad endpoint or key surfaces at creation time // rather than on the first turn. A keyless service is allowed. key := strings.TrimSpace(body.APIKey) - if key != "" { + if key != "" || len(headers) > 0 { client, err := llm.New(llm.Options{ - Kind: "openai-compatible", BaseURL: baseURL, APIKey: key, + Kind: "openai-compatible", BaseURL: baseURL, APIKey: key, Headers: headers, ProviderID: id, Timeout: 30 * time.Second, }) if err != nil { @@ -295,7 +306,7 @@ func (s *Server) handleCreateProvider(w http.ResponseWriter, r *http.Request) { } cfg.Providers[id] = config.Provider{ - Kind: "openai-compatible", BaseURL: baseURL, APIKey: key, + Kind: "openai-compatible", BaseURL: baseURL, APIKey: key, Headers: headers, Enabled: true, Label: name, } if err := config.Save(cfg); err != nil { diff --git a/internal/server/handlers_setup.go b/internal/server/handlers_setup.go index 80ef917..407e80c 100644 --- a/internal/server/handlers_setup.go +++ b/internal/server/handlers_setup.go @@ -202,9 +202,10 @@ func (s *Server) handleSetupTest(w http.ResponseWriter, r *http.Request) { return } var body struct { - Provider string `json:"provider"` - BaseURL string `json:"base_url"` - APIKey string `json:"api_key"` + Provider string `json:"provider"` + BaseURL string `json:"base_url"` + APIKey string `json:"api_key"` + Headers map[string]string `json:"headers"` } if err := decodeBody(r, &body); err != nil { writeError(w, http.StatusBadRequest, err) @@ -244,6 +245,15 @@ func (s *Server) handleSetupTest(w http.ResponseWriter, r *http.Request) { apiKey = p.APIKey } } + headers := cfg.Providers[body.Provider].Headers + if chosen.Custom && body.Headers != nil { + var err error + headers, err = config.NormalizeProviderHeaders(body.Headers) + if err != nil { + writeError(w, http.StatusBadRequest, err) + return + } + } // A keyless custom service on a LAN is legitimate; everything else needs // a credential unless the endpoint is local. if apiKey == "" && !chosen.Custom && !isLocalEndpoint(baseURL) { @@ -254,7 +264,7 @@ func (s *Server) handleSetupTest(w http.ResponseWriter, r *http.Request) { } client, err := llm.New(llm.Options{ - Kind: chosen.Kind, BaseURL: baseURL, APIKey: apiKey, + Kind: chosen.Kind, BaseURL: baseURL, APIKey: apiKey, Headers: headers, ProviderID: body.Provider, Timeout: 30 * time.Second, }) if err != nil { @@ -305,12 +315,13 @@ func (s *Server) handleSetupComplete(w http.ResponseWriter, r *http.Request) { } var body struct { - Provider string `json:"provider"` - Name string `json:"name"` - BaseURL string `json:"base_url"` - APIKey string `json:"api_key"` - Model string `json:"model"` - Workspace string `json:"workspace"` + Provider string `json:"provider"` + Name string `json:"name"` + BaseURL string `json:"base_url"` + APIKey string `json:"api_key"` + Headers map[string]string `json:"headers"` + Model string `json:"model"` + Workspace string `json:"workspace"` Database struct { Driver string `json:"driver"` DSN string `json:"dsn"` @@ -379,6 +390,19 @@ func (s *Server) handleSetupComplete(w http.ResponseWriter, r *http.Request) { if baseURL != "" { entry.BaseURL = baseURL } + if chosen.Custom { + if entry.Headers == nil { + entry.Headers = cfg.Providers[body.Provider].Headers + } + if body.Headers != nil { + headers, err := config.NormalizeProviderHeaders(body.Headers) + if err != nil { + writeError(w, http.StatusBadRequest, err) + return + } + entry.Headers = headers + } + } if key := strings.TrimSpace(body.APIKey); key != "" && !strings.Contains(key, "••••") { entry.APIKey = key } @@ -551,7 +575,7 @@ func (s *Server) handleSetProviderKey(w http.ResponseWriter, r *http.Request) { // Reject a bad key here rather than saving it and failing on the next turn. client, err := llm.New(llm.Options{ - Kind: entry.Kind, BaseURL: baseURL, APIKey: key, + Kind: entry.Kind, BaseURL: baseURL, APIKey: key, Headers: entry.Headers, Region: region, APIVersion: apiVersion, ProviderID: id, Timeout: 30 * time.Second, }) diff --git a/internal/server/provider_headers_test.go b/internal/server/provider_headers_test.go new file mode 100644 index 0000000..6f6a010 --- /dev/null +++ b/internal/server/provider_headers_test.go @@ -0,0 +1,179 @@ +package server + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + + "github.com/enowdev/antares/internal/agent" + "github.com/enowdev/antares/internal/config" + "github.com/enowdev/antares/internal/store" +) + +func headerProviderFixture(t *testing.T) (*httptest.Server, *string) { + t.Helper() + expected := "team=a=b" + fixture := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/models" || r.Header.Get("X-Tenant") != expected { + w.WriteHeader(http.StatusUnauthorized) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"header-model","owned_by":"test"}]}`)) + })) + t.Cleanup(fixture.Close) + return fixture, &expected +} + +func providerJSONRequest(method, path, body string) *http.Request { + r := httptest.NewRequest(method, path, strings.NewReader(body)) + r.Header.Set("Content-Type", "application/json") + return r +} + +func assertProviderOK(t *testing.T, rr *httptest.ResponseRecorder) map[string]any { + t.Helper() + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, body = %s", rr.Code, rr.Body.String()) + } + var response map[string]any + if err := json.Unmarshal(rr.Body.Bytes(), &response); err != nil { + t.Fatalf("decode response: %v", err) + } + if ok, _ := response["ok"].(bool); !ok { + t.Fatalf("response ok = %v, body = %s", response["ok"], rr.Body.String()) + } + return response +} + +func TestProviderHeadersCreateAndReconnect(t *testing.T) { + fixture, expected := headerProviderFixture(t) + s := newProviderKeyServer(t, func(*config.Config) {}) + + rr := httptest.NewRecorder() + s.handleCreateProvider(rr, providerJSONRequest(http.MethodPost, "/api/providers", `{"name":"Header Gateway","base_url":"`+fixture.URL+`/v1","headers":{"X-Tenant":"team=a=b"}}`)) + response := assertProviderOK(t, rr) + id := response["id"].(string) + + reloaded, err := config.Reload() + if err != nil { + t.Fatal(err) + } + if got := reloaded.Providers[id].Headers["X-Tenant"]; got != *expected { + t.Fatalf("saved header = %q, want %q", got, *expected) + } + + rr = postProviderKey(s, id, `{"api_key":"","base_url":"`+fixture.URL+`/v1"}`) + assertProviderOK(t, rr) +} + +func TestProviderHeadersRejectInvalidAndPreserveSettings(t *testing.T) { + fixture, _ := headerProviderFixture(t) + s := newProviderKeyServer(t, func(cfg *config.Config) { + cfg.Providers["gateway"] = config.Provider{ + Kind: "openai-compatible", BaseURL: fixture.URL + "/v1", Enabled: true, + Headers: map[string]string{"X-Old": "old"}, + } + }) + + for _, body := range []string{ + `{"headers":{"bad header":"value"}}`, + "{\"headers\":{\"X-Test\":\"bad\\nvalue\"}}", + `{"headers":{"X-Test":"one","x-test":"two"}}`, + } { + rr := httptest.NewRecorder() + r := providerJSONRequest(http.MethodPatch, "/api/providers/gateway/settings", body) + r.SetPathValue("id", "gateway") + s.handleProviderSettings(rr, r) + if rr.Code != http.StatusBadRequest { + t.Fatalf("invalid headers status = %d, body = %s", rr.Code, rr.Body.String()) + } + reloaded, err := config.Reload() + if err != nil { + t.Fatal(err) + } + if got := reloaded.Providers["gateway"].Headers["X-Old"]; got != "old" { + t.Fatalf("invalid headers changed saved headers: %#v", reloaded.Providers["gateway"].Headers) + } + } +} + +func TestProviderHeadersSettingsReplacePreserveAndClear(t *testing.T) { + s := newProviderKeyServer(t, func(cfg *config.Config) { + cfg.Providers["gateway"] = config.Provider{Headers: map[string]string{"X-Old": "old"}} + }) + request := func(body string) { + t.Helper() + rr := httptest.NewRecorder() + r := providerJSONRequest(http.MethodPatch, "/api/providers/gateway/settings", body) + r.SetPathValue("id", "gateway") + s.handleProviderSettings(rr, r) + assertProviderOK(t, rr) + } + + request(`{"label":"Gateway"}`) + reloaded, err := config.Reload() + if err != nil { + t.Fatal(err) + } + if got := reloaded.Providers["gateway"].Headers["X-Old"]; got != "old" { + t.Fatalf("omitted headers = %#v, want existing map", reloaded.Providers["gateway"].Headers) + } + + request(`{"headers":{"x-new":" new "}}`) + reloaded, err = config.Reload() + if err != nil { + t.Fatal(err) + } + if got := reloaded.Providers["gateway"].Headers; len(got) != 1 || got["X-New"] != "new" { + t.Fatalf("replacement headers = %#v", got) + } + + request(`{"headers":{}}`) + reloaded, err = config.Reload() + if err != nil { + t.Fatal(err) + } + if got := reloaded.Providers["gateway"].Headers; got == nil || len(got) != 0 { + t.Fatalf("cleared headers = %#v, want allocated empty map", got) + } +} + +func TestProviderHeadersSetupTestAndComplete(t *testing.T) { + fixture, expected := headerProviderFixture(t) + home := t.TempDir() + t.Setenv("ANTARES_HOME", home) + cfg := config.Default() + cfg.Model.Default = "" + if err := config.Save(cfg); err != nil { + t.Fatal(err) + } + db, err := store.Open(t.Context(), "sqlite", filepath.Join(t.TempDir(), "test.db"), 1, 5000, false) + if err != nil { + t.Fatal(err) + } + s := &Server{cfg: cfg, db: db, agent: agent.New(cfg, db, nil, nil, nil)} + + testRequest := providerJSONRequest(http.MethodPost, "/api/setup/test", `{"provider":"custom","base_url":"`+fixture.URL+`/v1","headers":{"X-Tenant":"team=a=b"}}`) + testRequest.RemoteAddr = "127.0.0.1:12345" + rr := httptest.NewRecorder() + s.handleSetupTest(rr, testRequest) + assertProviderOK(t, rr) + + completeRequest := providerJSONRequest(http.MethodPost, "/api/setup/complete", `{"provider":"custom","name":"Header Setup","base_url":"`+fixture.URL+`/v1","headers":{"X-Tenant":"team=a=b"},"model":"header-model","workspace":"`+filepath.Join(t.TempDir(), "workspace")+`"}`) + completeRequest.RemoteAddr = "127.0.0.1:12345" + rr = httptest.NewRecorder() + s.handleSetupComplete(rr, completeRequest) + assertProviderOK(t, rr) + + reloaded, err := config.Reload() + if err != nil { + t.Fatal(err) + } + if got := reloaded.Providers["header-setup"].Headers["X-Tenant"]; got != *expected { + t.Fatalf("completed headers = %q, want %q", got, *expected) + } +} From 64b551ad06a06aa7974a0559ca80470ca19a1026 Mon Sep 17 00:00:00 2001 From: Reidho Satria Date: Tue, 15 Sep 2026 12:51:21 +0700 Subject: [PATCH 3/9] fix(providers): recognize header-authenticated custom providers --- cmd/antares/setup.go | 12 +--- internal/server/handlers_chat.go | 3 +- internal/server/handlers_config.go | 7 +- internal/server/handlers_setup.go | 12 +++- internal/server/provider_headers_test.go | 84 ++++++++++++++++++++++++ 5 files changed, 100 insertions(+), 18 deletions(-) diff --git a/cmd/antares/setup.go b/cmd/antares/setup.go index bd4f949..6e094a4 100644 --- a/cmd/antares/setup.go +++ b/cmd/antares/setup.go @@ -24,17 +24,7 @@ import ( ) // needsSetup reports whether Antares has enough configuration to answer at all. -func needsSetup(cfg *config.Config) bool { - if strings.TrimSpace(cfg.Model.Default) == "" { - return true - } - _, p := cfg.ResolveProvider(cfg.Model.Provider) - // A local endpoint needs no credential; everything else does. - if p.APIKey == "" && !isLocalEndpoint(p.BaseURL) { - return true - } - return false -} +func needsSetup(cfg *config.Config) bool { return server.NeedsSetup(cfg) } func isLocalEndpoint(url string) bool { l := strings.ToLower(url) diff --git a/internal/server/handlers_chat.go b/internal/server/handlers_chat.go index 1532cda..58ebf94 100644 --- a/internal/server/handlers_chat.go +++ b/internal/server/handlers_chat.go @@ -34,8 +34,7 @@ func (s *Server) handleStatus(w http.ResponseWriter, r *http.Request) { // Provider readiness is a config check, not a live call: the status pill // polls every ten seconds and must not bill the user for pings. - _, provider := cfg.ResolveProvider(cfg.Model.Provider) - ready := cfg.Model.Default != "" && (provider.APIKey != "" || isLocalEndpoint(provider.BaseURL)) + ready := !NeedsSetup(cfg) writeJSON(w, http.StatusOK, map[string]any{ "ok": true, diff --git a/internal/server/handlers_config.go b/internal/server/handlers_config.go index bd37b2e..0449195 100644 --- a/internal/server/handlers_config.go +++ b/internal/server/handlers_config.go @@ -251,7 +251,7 @@ func (s *Server) handleModelList(w http.ResponseWriter, r *http.Request) { // Calling a provider we know has no credential just turns a known state // into an opaque 401. Report the missing key instead. id, p := cfg.ResolveProvider(provider) - if p.APIKey == "" && !isLocalEndpoint(p.BaseURL) { + if p.APIKey == "" && !isLocalEndpoint(p.BaseURL) && !customProviderHasHeaders(cfg, id, p) { writeJSON(w, http.StatusOK, map[string]any{ "models": []any{}, "needs_key": true, "provider": id, }) @@ -286,7 +286,8 @@ func (s *Server) handleModelListAll(w http.ResponseWriter, r *http.Request) { } cfg := s.config() - // Which providers are worth calling: a stored key, a set key-env, or local. + // Which providers are worth calling: a stored key, a set key-env, local, or + // a configured custom provider with explicit request headers. type target struct { id, label string } @@ -298,7 +299,7 @@ func (s *Server) handleModelListAll(w http.ResponseWriter, r *http.Request) { } p := cfg.Providers[id] keyed := p.APIKey != "" || (p.APIKeyEnv != "" && os.Getenv(p.APIKeyEnv) != "") - if keyed || isLocalEndpoint(p.BaseURL) { + if keyed || isLocalEndpoint(p.BaseURL) || customProviderHasHeaders(cfg, id, p) { targets = append(targets, target{id: id, label: firstNonEmpty(p.Label, label, id)}) seen[id] = true } diff --git a/internal/server/handlers_setup.go b/internal/server/handlers_setup.go index 407e80c..c1feab1 100644 --- a/internal/server/handlers_setup.go +++ b/internal/server/handlers_setup.go @@ -175,8 +175,16 @@ func NeedsSetup(cfg *config.Config) bool { if strings.TrimSpace(cfg.Model.Default) == "" { return true } - _, p := cfg.ResolveProvider(cfg.Model.Provider) - return p.APIKey == "" && !isLocalEndpoint(p.BaseURL) + id, p := cfg.ResolveProvider(cfg.Model.Provider) + return p.APIKey == "" && !isLocalEndpoint(p.BaseURL) && !customProviderHasHeaders(cfg, id, p) +} + +func customProviderHasHeaders(cfg *config.Config, id string, p config.Provider) bool { + if strings.TrimSpace(p.BaseURL) == "" || len(p.Headers) == 0 { + return false + } + provider := lookupSetupProvider(cfg, id) + return provider == nil || provider.Custom } func (s *Server) handleSetupStatus(w http.ResponseWriter, r *http.Request) { diff --git a/internal/server/provider_headers_test.go b/internal/server/provider_headers_test.go index 6f6a010..73d0c49 100644 --- a/internal/server/provider_headers_test.go +++ b/internal/server/provider_headers_test.go @@ -1,10 +1,13 @@ package server import ( + "context" "encoding/json" + "net" "net/http" "net/http/httptest" "path/filepath" + "strconv" "strings" "testing" @@ -177,3 +180,84 @@ func TestProviderHeadersSetupTestAndComplete(t *testing.T) { t.Fatalf("completed headers = %q, want %q", got, *expected) } } + +func TestHeaderAuthenticatedCustomProviderIsReadyAndDiscoverable(t *testing.T) { + fixture, expected := headerProviderFixture(t) + aliasURL := "http://127.0.0.2:" + strconv.Itoa(fixture.Listener.Addr().(*net.TCPAddr).Port) + "/v1" + + originalTransport := http.DefaultTransport + transport := originalTransport.(*http.Transport).Clone() + transport.Proxy = nil + transport.DialContext = func(ctx context.Context, network, address string) (net.Conn, error) { + return (&net.Dialer{}).DialContext(ctx, network, fixture.Listener.Addr().String()) + } + http.DefaultTransport = transport + t.Cleanup(func() { http.DefaultTransport = originalTransport }) + + cfg := config.Default() + cfg.Server.DashboardPasswordHash = "test-hash" + cfg.Model.Provider = "header-gateway" + cfg.Model.Default = "header-model" + cfg.Providers["header-gateway"] = config.Provider{ + Kind: "openai-compatible", BaseURL: aliasURL, Enabled: true, + Headers: map[string]string{"X-Tenant": *expected}, + } + if NeedsSetup(cfg) { + t.Fatal("header-authenticated custom provider still needs setup") + } + if NeedsSetup(&config.Config{Model: config.Model{Provider: "header-gateway", Default: "header-model"}, Providers: map[string]config.Provider{"header-gateway": {BaseURL: aliasURL}}}) == false { + t.Fatal("headerless remote provider did not need setup") + } + if NeedsSetup(&config.Config{Model: config.Model{Provider: "openai", Default: "gpt-5"}, Providers: map[string]config.Provider{"openai": {BaseURL: "https://api.openai.com/v1", Headers: map[string]string{"X-Tenant": *expected}}}}) == false { + t.Fatal("built-in provider headers bypassed setup") + } + + db, err := store.Open(t.Context(), "memory", "", 1, 5000, false) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + s := &Server{cfg: cfg, db: db, agent: agent.New(cfg, db, nil, nil, nil)} + + status := httptest.NewRecorder() + s.handleStatus(status, httptest.NewRequest(http.MethodGet, "/api/status", nil)) + var statusBody struct { + ProviderReady bool `json:"provider_ready"` + NeedsSetup bool `json:"needs_setup"` + } + if err := json.Unmarshal(status.Body.Bytes(), &statusBody); err != nil { + t.Fatal(err) + } + if !statusBody.ProviderReady || statusBody.NeedsSetup { + t.Fatalf("status = %#v, want ready and complete", statusBody) + } + + list := httptest.NewRecorder() + s.handleModelList(list, httptest.NewRequest(http.MethodGet, "/api/models?provider=header-gateway", nil)) + var listBody struct { + Models []struct { + ID string `json:"id"` + } `json:"models"` + } + if err := json.Unmarshal(list.Body.Bytes(), &listBody); err != nil { + t.Fatal(err) + } + if len(listBody.Models) != 1 || listBody.Models[0].ID != "header-model" { + t.Fatalf("single provider models = %#v", listBody.Models) + } + + all := httptest.NewRecorder() + s.handleModelListAll(all, httptest.NewRequest(http.MethodGet, "/api/models/all", nil)) + var allBody struct { + Models []struct { + ID string `json:"id"` + Provider string `json:"provider"` + } `json:"models"` + } + if err := json.Unmarshal(all.Body.Bytes(), &allBody); err != nil { + t.Fatal(err) + } + if len(allBody.Models) != 1 || allBody.Models[0].ID != "header-model" || allBody.Models[0].Provider != "header-gateway" { + t.Fatalf("all provider models = %#v", allBody.Models) + } +} From a1efaedd4e384f880ada8a795c422c066cb0bbe7 Mon Sep 17 00:00:00 2001 From: Reidho Satria Date: Tue, 15 Sep 2026 12:55:58 +0700 Subject: [PATCH 4/9] feat(web): edit custom provider headers in provider dialogs --- .../providers/ProviderHeadersField.tsx | 28 +++++++++++++ web/src/lib/i18n.tsx | 3 ++ web/src/lib/providerHeaders.test.mjs | 35 ++++++++++++++++ web/src/lib/providerHeaders.ts | 38 +++++++++++++++++ web/src/pages/ProvidersPage.tsx | 42 ++++++++++++++++--- 5 files changed, 140 insertions(+), 6 deletions(-) create mode 100644 web/src/components/providers/ProviderHeadersField.tsx create mode 100644 web/src/lib/providerHeaders.test.mjs create mode 100644 web/src/lib/providerHeaders.ts diff --git a/web/src/components/providers/ProviderHeadersField.tsx b/web/src/components/providers/ProviderHeadersField.tsx new file mode 100644 index 0000000..11aa772 --- /dev/null +++ b/web/src/components/providers/ProviderHeadersField.tsx @@ -0,0 +1,28 @@ +import { useI18n } from '@/lib/i18n' +import { Label, Textarea } from '@/components/ui/primitives' + +export function ProviderHeadersField({ + id, + value, + onChange, +}: { + id: string + value: string + onChange: (value: string) => void +}) { + const { t } = useI18n() + return ( +
+ +