fix: re-validate provider per request and classify reloads (#25766)

Refactors the `aibridgeproxyd` provider reload mechanism which was unnecessarily complex.

Also ensures that providers are evaluated on each CONNECT request to prevent interception of requests to (newly) disabled providers; in this case the requests will passthrough unencrypted, by design.
This commit is contained in:
Danny Kopping
2026-05-28 13:22:38 +02:00
committed by GitHub
parent 673709bd34
commit a9f5ed7644
7 changed files with 856 additions and 629 deletions
+56 -6
View File
@@ -5,7 +5,9 @@ package cli
import (
"context"
"io"
"net/url"
"path/filepath"
"strings"
"github.com/prometheus/client_golang/prometheus"
"golang.org/x/xerrors"
@@ -86,19 +88,67 @@ func newAIBridgeProxyDaemon(coderAPI *coderd.API) (io.Closer, error) {
}, nil
}
// refreshProxyProviders classifies every ai_providers row as enabled,
// disabled, or error so the proxy router and any observers see the full
// configured set. Disabled rows are excluded from routing; errored rows
// are excluded from routing and surface their failure reason for
// metrics and logs.
func refreshProxyProviders(db database.Store) aibridgeproxyd.RefreshProvidersFunc {
return func(ctx context.Context) ([]aibridgeproxyd.ProviderRoute, error) {
return func(ctx context.Context) (aibridgeproxyd.ProviderReload, error) {
//nolint:gocritic // AsAIProviderMetadataReader is the correct subject for routing-only access.
rows, err := db.GetAIProviders(dbauthz.AsAIProviderMetadataReader(ctx), database.GetAIProvidersParams{
IncludeDisabled: false,
IncludeDisabled: true,
})
if err != nil {
return nil, xerrors.Errorf("load ai providers: %w", err)
return aibridgeproxyd.ProviderReload{}, xerrors.Errorf("load ai providers: %w", err)
}
out := make([]aibridgeproxyd.ProviderRoute, 0, len(rows))
reload := aibridgeproxyd.ProviderReload{
Providers: make([]aibridgeproxyd.ReloadedProvider, 0, len(rows)),
}
seenHost := make(map[string]string, len(rows))
for _, row := range rows {
out = append(out, aibridgeproxyd.ProviderRoute{Name: row.Name, BaseURL: row.BaseUrl})
reload.Providers = append(reload.Providers, classifyProviderRow(row, seenHost))
}
return out, nil
return reload, nil
}
}
// classifyProviderRow evaluates a single ai_providers row for routing.
// seenHost is mutated to track the first provider that claimed each
// hostname so later duplicates can be flagged as errors.
func classifyProviderRow(row database.AIProvider, seenHost map[string]string) aibridgeproxyd.ReloadedProvider {
out := aibridgeproxyd.ReloadedProvider{
Name: row.Name,
Type: string(row.Type),
}
if !row.Enabled {
out.Status = aibridgeproxyd.ProviderStatusDisabled
return out
}
if strings.TrimSpace(row.BaseUrl) == "" {
out.Status = aibridgeproxyd.ProviderStatusError
out.Err = xerrors.New("base url is empty")
return out
}
u, err := url.Parse(row.BaseUrl)
if err != nil {
out.Status = aibridgeproxyd.ProviderStatusError
out.Err = xerrors.Errorf("invalid base url %q: %w", row.BaseUrl, err)
return out
}
host := strings.ToLower(u.Hostname())
if host == "" {
out.Status = aibridgeproxyd.ProviderStatusError
out.Err = xerrors.Errorf("base url %q has no hostname", row.BaseUrl)
return out
}
if claimedBy, taken := seenHost[host]; taken {
out.Status = aibridgeproxyd.ProviderStatusError
out.Err = xerrors.Errorf("hostname %q already claimed by provider %q", host, claimedBy)
return out
}
seenHost[host] = row.Name
out.Host = host
out.Status = aibridgeproxyd.ProviderStatusEnabled
return out
}
@@ -0,0 +1,105 @@
//go:build !slim
package cli
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/enterprise/aibridgeproxyd"
)
// TestClassifyProviderRow covers every branch of the classifier so the
// disabled, error, and enabled paths are exercised through the
// production code instead of relying on classifyRaw, the test mirror in
// reload_test.go.
func TestClassifyProviderRow(t *testing.T) {
t.Parallel()
enabledRow := func(name, baseURL string) database.AIProvider {
return database.AIProvider{
Name: name,
Type: database.AiProviderTypeOpenai,
Enabled: true,
BaseUrl: baseURL,
}
}
t.Run("Enabled", func(t *testing.T) {
t.Parallel()
seen := map[string]string{}
got := classifyProviderRow(enabledRow("openai", "https://api.openai.com/v1"), seen)
assert.Equal(t, "openai", got.Name)
assert.Equal(t, string(database.AiProviderTypeOpenai), got.Type)
assert.Equal(t, aibridgeproxyd.ProviderStatusEnabled, got.Status)
assert.Equal(t, "api.openai.com", got.Host)
assert.NoError(t, got.Err)
assert.Equal(t, "openai", seen["api.openai.com"])
})
t.Run("DisabledRow", func(t *testing.T) {
t.Parallel()
seen := map[string]string{}
row := enabledRow("off", "https://api.off.example.com/v1")
row.Enabled = false
got := classifyProviderRow(row, seen)
assert.Equal(t, aibridgeproxyd.ProviderStatusDisabled, got.Status)
assert.Empty(t, got.Host, "disabled provider must not claim a host")
assert.NoError(t, got.Err)
assert.Empty(t, seen, "disabled provider must not occupy a host slot")
})
t.Run("EmptyBaseURL", func(t *testing.T) {
t.Parallel()
seen := map[string]string{}
got := classifyProviderRow(enabledRow("no-url", " "), seen)
assert.Equal(t, aibridgeproxyd.ProviderStatusError, got.Status)
assert.Empty(t, got.Host)
assert.ErrorContains(t, got.Err, "base url is empty")
})
t.Run("MalformedBaseURL", func(t *testing.T) {
t.Parallel()
seen := map[string]string{}
got := classifyProviderRow(enabledRow("bad", "://not-a-url"), seen)
assert.Equal(t, aibridgeproxyd.ProviderStatusError, got.Status)
assert.ErrorContains(t, got.Err, "invalid base url")
})
t.Run("BaseURLWithoutHostname", func(t *testing.T) {
t.Parallel()
seen := map[string]string{}
got := classifyProviderRow(enabledRow("no-host", "https://"), seen)
assert.Equal(t, aibridgeproxyd.ProviderStatusError, got.Status)
assert.ErrorContains(t, got.Err, "no hostname")
})
t.Run("DuplicateHostnameFirstWins", func(t *testing.T) {
t.Parallel()
seen := map[string]string{}
first := classifyProviderRow(enabledRow("first", "https://shared.example.com/v1"), seen)
assert.Equal(t, aibridgeproxyd.ProviderStatusEnabled, first.Status)
second := classifyProviderRow(enabledRow("second", "https://shared.example.com/v2"), seen)
assert.Equal(t, aibridgeproxyd.ProviderStatusError, second.Status)
assert.ErrorContains(t, second.Err, "already claimed by provider \"first\"")
assert.Equal(t, "first", seen["shared.example.com"], "first wins must not be overwritten")
})
t.Run("HostnameLowercased", func(t *testing.T) {
t.Parallel()
seen := map[string]string{}
got := classifyProviderRow(enabledRow("mixed", "https://API.Example.COM/v1"), seen)
assert.Equal(t, aibridgeproxyd.ProviderStatusEnabled, got.Status)
assert.Equal(t, "api.example.com", got.Host)
})
}