mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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)
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user