mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-30 19:41:30 +02:00
endpoint model discovery and proxy integration
This commit is contained in:
337
proxy/internal/modeldiscovery/discovery.go
Normal file
337
proxy/internal/modeldiscovery/discovery.go
Normal file
@@ -0,0 +1,337 @@
|
||||
// Package modeldiscovery fetches and normalizes model catalogs from
|
||||
// Agent Network provider endpoints. Discovery always runs on the proxy so
|
||||
// it observes the same network path as inference traffic.
|
||||
package modeldiscovery
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"path"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/roundtrip"
|
||||
)
|
||||
|
||||
const (
|
||||
SourceOpenAIV1Models = "openai_v1_models"
|
||||
SourceOllamaAPITags = "ollama_api_tags"
|
||||
|
||||
defaultTimeout = 5 * time.Second
|
||||
maxResponseBytes = 1 << 20 // 1 MiB, after HTTP decompression.
|
||||
maxUpstreamURLBytes = 4096
|
||||
maxHeaderNameBytes = 256
|
||||
maxHeaderValueBytes = 64 << 10
|
||||
maxModels = 500
|
||||
maxModelIDBytes = 512
|
||||
)
|
||||
|
||||
// Request contains the provider-owned values management resolved from the
|
||||
// persisted provider record. Callers must not populate these fields from a
|
||||
// dashboard-supplied URL or credential.
|
||||
type Request struct {
|
||||
UpstreamURL string
|
||||
AuthHeaderName string
|
||||
AuthHeaderValue string
|
||||
SkipTLSVerify bool
|
||||
AllowOllamaFallback bool
|
||||
}
|
||||
|
||||
// Model is the deliberately small response surface returned to management.
|
||||
// Arbitrary fields supplied by an upstream never cross the control channel.
|
||||
type Model struct {
|
||||
ID string
|
||||
Label string
|
||||
}
|
||||
|
||||
// Result is a normalized model catalog and the endpoint shape that supplied
|
||||
// it.
|
||||
type Result struct {
|
||||
Models []Model
|
||||
Source string
|
||||
}
|
||||
|
||||
// Discoverer owns the HTTP client used for provider probes.
|
||||
type Discoverer struct {
|
||||
client *http.Client
|
||||
timeout time.Duration
|
||||
bodyLimit int64
|
||||
}
|
||||
|
||||
// New constructs a direct-upstream discoverer. The transport is the same
|
||||
// host-network transport family used by Agent Network inference routes.
|
||||
func New(logger *log.Logger) *Discoverer {
|
||||
return newWithTransport(roundtrip.NewDirectOnly(logger))
|
||||
}
|
||||
|
||||
func newWithTransport(transport http.RoundTripper) *Discoverer {
|
||||
return &Discoverer{
|
||||
client: &http.Client{
|
||||
Transport: transport,
|
||||
// Redirects could move a credentialed request away from the
|
||||
// persisted provider origin. Discovery never follows them.
|
||||
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
},
|
||||
timeout: defaultTimeout,
|
||||
bodyLimit: maxResponseBytes,
|
||||
}
|
||||
}
|
||||
|
||||
// Discover queries the OpenAI-compatible model-list endpoint. Ollama's native
|
||||
// tags endpoint is attempted only when explicitly enabled and the primary
|
||||
// endpoint reports that the route does not exist.
|
||||
func (d *Discoverer) Discover(ctx context.Context, in Request) (Result, error) {
|
||||
if d == nil || d.client == nil {
|
||||
return Result{}, errors.New("model discovery client is unavailable")
|
||||
}
|
||||
if err := validateRequest(in); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
|
||||
timeout := d.timeout
|
||||
if timeout <= 0 || timeout > defaultTimeout {
|
||||
timeout = defaultTimeout
|
||||
}
|
||||
probeCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
models, err := d.fetchOpenAIModels(probeCtx, in)
|
||||
if err == nil {
|
||||
return Result{Models: models, Source: SourceOpenAIV1Models}, nil
|
||||
}
|
||||
if !in.AllowOllamaFallback || !isMissingEndpoint(err) {
|
||||
return Result{}, err
|
||||
}
|
||||
|
||||
models, err = d.fetchOllamaTags(probeCtx, in)
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
return Result{Models: models, Source: SourceOllamaAPITags}, nil
|
||||
}
|
||||
|
||||
func validateRequest(in Request) error {
|
||||
rawURL := strings.TrimSpace(in.UpstreamURL)
|
||||
if rawURL == "" {
|
||||
return errors.New("model discovery upstream URL is required")
|
||||
}
|
||||
if len(rawURL) > maxUpstreamURLBytes {
|
||||
return errors.New("model discovery upstream URL is too long")
|
||||
}
|
||||
if len(in.AuthHeaderName) > maxHeaderNameBytes || len(in.AuthHeaderValue) > maxHeaderValueBytes {
|
||||
return errors.New("model discovery authentication header is too large")
|
||||
}
|
||||
if (in.AuthHeaderName == "") != (in.AuthHeaderValue == "") {
|
||||
return errors.New("model discovery authentication header is incomplete")
|
||||
}
|
||||
if strings.ContainsAny(in.AuthHeaderName, "\r\n") || strings.ContainsAny(in.AuthHeaderValue, "\r\n") {
|
||||
return errors.New("model discovery authentication header is invalid")
|
||||
}
|
||||
if in.AuthHeaderName != "" && !strings.EqualFold(in.AuthHeaderName, "Authorization") {
|
||||
return errors.New("model discovery authentication header is unsupported")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *Discoverer) fetchOpenAIModels(ctx context.Context, in Request) ([]Model, error) {
|
||||
body, err := d.fetch(ctx, in, "v1/models")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var payload struct {
|
||||
Data *[]struct {
|
||||
ID string `json:"id"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &payload); err != nil || payload.Data == nil {
|
||||
return nil, errors.New("upstream returned invalid OpenAI model-list JSON")
|
||||
}
|
||||
|
||||
ids := make([]string, 0, len(*payload.Data))
|
||||
for _, model := range *payload.Data {
|
||||
ids = append(ids, model.ID)
|
||||
}
|
||||
return normalize(ids)
|
||||
}
|
||||
|
||||
func (d *Discoverer) fetchOllamaTags(ctx context.Context, in Request) ([]Model, error) {
|
||||
body, err := d.fetch(ctx, in, "api/tags")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var payload struct {
|
||||
Models *[]struct {
|
||||
Name string `json:"name"`
|
||||
Model string `json:"model"`
|
||||
} `json:"models"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &payload); err != nil || payload.Models == nil {
|
||||
return nil, errors.New("upstream returned invalid Ollama tags JSON")
|
||||
}
|
||||
|
||||
ids := make([]string, 0, len(*payload.Models))
|
||||
for _, model := range *payload.Models {
|
||||
id := model.Model
|
||||
if strings.TrimSpace(id) == "" {
|
||||
id = model.Name
|
||||
}
|
||||
ids = append(ids, id)
|
||||
}
|
||||
return normalize(ids)
|
||||
}
|
||||
|
||||
func (d *Discoverer) fetch(ctx context.Context, in Request, endpointPath string) ([]byte, error) {
|
||||
endpoint, err := buildEndpointURL(in.UpstreamURL, endpointPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
reqCtx := roundtrip.WithDirectUpstream(ctx)
|
||||
if in.SkipTLSVerify {
|
||||
reqCtx = roundtrip.WithSkipTLSVerify(reqCtx)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, endpoint.String(), nil)
|
||||
if err != nil {
|
||||
return nil, errors.New("create model discovery request")
|
||||
}
|
||||
req.Header.Set("Accept", "application/json")
|
||||
if in.AuthHeaderName != "" {
|
||||
// Phase 3 is Ollama-only. Canonicalizing the sole catalog-owned
|
||||
// credential header keeps the control message from becoming a generic
|
||||
// arbitrary-header primitive.
|
||||
req.Header.Set("Authorization", in.AuthHeaderValue)
|
||||
}
|
||||
|
||||
resp, err := d.client.Do(req)
|
||||
if err != nil {
|
||||
if errors.Is(ctx.Err(), context.DeadlineExceeded) {
|
||||
return nil, errors.New("model discovery timed out")
|
||||
}
|
||||
if errors.Is(ctx.Err(), context.Canceled) {
|
||||
return nil, context.Canceled
|
||||
}
|
||||
// Deliberately omit the underlying error: net/http errors include the
|
||||
// internal URL, which should not be reflected through the public API.
|
||||
return nil, errors.New("model discovery request failed")
|
||||
}
|
||||
defer func() {
|
||||
_ = resp.Body.Close()
|
||||
}()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, &upstreamStatusError{statusCode: resp.StatusCode}
|
||||
}
|
||||
|
||||
limit := d.bodyLimit
|
||||
if limit <= 0 || limit > maxResponseBytes {
|
||||
limit = maxResponseBytes
|
||||
}
|
||||
if resp.ContentLength > limit {
|
||||
return nil, errors.New("model discovery response is too large")
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, limit+1))
|
||||
if err != nil {
|
||||
return nil, errors.New("read model discovery response")
|
||||
}
|
||||
if int64(len(body)) > limit {
|
||||
return nil, errors.New("model discovery response is too large")
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
func buildEndpointURL(rawURL, endpointPath string) (*url.URL, error) {
|
||||
parsed, err := url.Parse(strings.TrimSpace(rawURL))
|
||||
if err != nil || parsed.Host == "" || parsed.Hostname() == "" || parsed.Opaque != "" {
|
||||
return nil, errors.New("model discovery upstream URL is invalid")
|
||||
}
|
||||
switch strings.ToLower(parsed.Scheme) {
|
||||
case "http":
|
||||
parsed.Scheme = "http"
|
||||
case "https":
|
||||
parsed.Scheme = "https"
|
||||
default:
|
||||
return nil, errors.New("model discovery upstream URL must use http or https")
|
||||
}
|
||||
if parsed.User != nil {
|
||||
return nil, errors.New("model discovery upstream URL must not contain credentials")
|
||||
}
|
||||
|
||||
// Match Agent Network routing semantics: the static discovery path is
|
||||
// appended to any persisted base path. Queries and fragments on a provider
|
||||
// URL are not forwarded to inference and are likewise excluded here.
|
||||
parsed.Path = "/" + strings.TrimPrefix(path.Join(parsed.Path, endpointPath), "/")
|
||||
parsed.RawPath = ""
|
||||
parsed.RawQuery = ""
|
||||
parsed.ForceQuery = false
|
||||
parsed.Fragment = ""
|
||||
return parsed, nil
|
||||
}
|
||||
|
||||
func normalize(ids []string) ([]Model, error) {
|
||||
if len(ids) > maxModels {
|
||||
return nil, errors.New("upstream returned too many models")
|
||||
}
|
||||
|
||||
seen := make(map[string]struct{}, len(ids))
|
||||
normalized := make([]string, 0, len(ids))
|
||||
for _, raw := range ids {
|
||||
id := strings.TrimSpace(raw)
|
||||
if !validModelID(id) {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
normalized = append(normalized, id)
|
||||
}
|
||||
sort.Strings(normalized)
|
||||
|
||||
models := make([]Model, 0, len(normalized))
|
||||
for _, id := range normalized {
|
||||
models = append(models, Model{ID: id, Label: id})
|
||||
}
|
||||
return models, nil
|
||||
}
|
||||
|
||||
func validModelID(id string) bool {
|
||||
if id == "" || len(id) > maxModelIDBytes || !utf8.ValidString(id) {
|
||||
return false
|
||||
}
|
||||
for _, r := range id {
|
||||
if unicode.IsControl(r) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
type upstreamStatusError struct {
|
||||
statusCode int
|
||||
}
|
||||
|
||||
func (e *upstreamStatusError) Error() string {
|
||||
return fmt.Sprintf("upstream returned HTTP %d", e.statusCode)
|
||||
}
|
||||
|
||||
func isMissingEndpoint(err error) bool {
|
||||
var statusErr *upstreamStatusError
|
||||
if !errors.As(err, &statusErr) {
|
||||
return false
|
||||
}
|
||||
return statusErr.statusCode == http.StatusNotFound || statusErr.statusCode == http.StatusMethodNotAllowed
|
||||
}
|
||||
315
proxy/internal/modeldiscovery/discovery_test.go
Normal file
315
proxy/internal/modeldiscovery/discovery_test.go
Normal file
@@ -0,0 +1,315 @@
|
||||
package modeldiscovery
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestDiscoverOpenAIModels(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
assert.Equal(t, http.MethodGet, r.Method)
|
||||
assert.Equal(t, "/gateway/v1/models", r.URL.Path)
|
||||
assert.Empty(t, r.URL.RawQuery)
|
||||
assert.Equal(t, "application/json", r.Header.Get("Accept"))
|
||||
assert.Equal(t, "Bearer secret", r.Header.Get("Authorization"))
|
||||
_, _ = io.WriteString(w, `{
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"id": " qwen2.5:latest ", "owned_by": "ignored"},
|
||||
{"id": "llama3.2:latest"},
|
||||
{"id": "llama3.2:latest"},
|
||||
{"id": ""},
|
||||
{"id": "bad\u0000id"}
|
||||
]
|
||||
}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
discoverer := newWithTransport(http.DefaultTransport)
|
||||
result, err := discoverer.Discover(context.Background(), Request{
|
||||
UpstreamURL: server.URL + "/gateway/?ignored=true#fragment",
|
||||
AuthHeaderName: "authorization",
|
||||
AuthHeaderValue: "Bearer secret",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, SourceOpenAIV1Models, result.Source)
|
||||
assert.Equal(t, []Model{
|
||||
{ID: "llama3.2:latest", Label: "llama3.2:latest"},
|
||||
{ID: "qwen2.5:latest", Label: "qwen2.5:latest"},
|
||||
}, result.Models)
|
||||
}
|
||||
|
||||
func TestDiscoverOllamaFallback(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var primaryCalls atomic.Int32
|
||||
var fallbackCalls atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/v1/models":
|
||||
primaryCalls.Add(1)
|
||||
http.NotFound(w, r)
|
||||
case "/api/tags":
|
||||
fallbackCalls.Add(1)
|
||||
_, _ = io.WriteString(w, `{
|
||||
"models": [
|
||||
{"name": "ignored-name", "model": "gemma3:latest"},
|
||||
{"name": "llama3.2:latest", "model": ""}
|
||||
]
|
||||
}`)
|
||||
default:
|
||||
http.Error(w, "unexpected path", http.StatusInternalServerError)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
result, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), Request{
|
||||
UpstreamURL: server.URL,
|
||||
AllowOllamaFallback: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int32(1), primaryCalls.Load())
|
||||
assert.Equal(t, int32(1), fallbackCalls.Load())
|
||||
assert.Equal(t, SourceOllamaAPITags, result.Source)
|
||||
assert.Equal(t, []Model{
|
||||
{ID: "gemma3:latest", Label: "gemma3:latest"},
|
||||
{ID: "llama3.2:latest", Label: "llama3.2:latest"},
|
||||
}, result.Models)
|
||||
}
|
||||
|
||||
func TestDiscoverDoesNotFallbackOnAuthenticationFailure(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var fallbackCalls atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/api/tags" {
|
||||
fallbackCalls.Add(1)
|
||||
}
|
||||
http.Error(w, "secret upstream body", http.StatusUnauthorized)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
_, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), Request{
|
||||
UpstreamURL: server.URL,
|
||||
AllowOllamaFallback: true,
|
||||
})
|
||||
require.EqualError(t, err, "upstream returned HTTP 401")
|
||||
assert.Zero(t, fallbackCalls.Load())
|
||||
assert.NotContains(t, err.Error(), "secret upstream body")
|
||||
assert.NotContains(t, err.Error(), server.URL)
|
||||
}
|
||||
|
||||
func TestDiscoverDoesNotFollowRedirectsOrLeakAuth(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var redirectedCalls atomic.Int32
|
||||
redirected := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
redirectedCalls.Add(1)
|
||||
assert.Empty(t, r.Header.Get("Authorization"))
|
||||
_, _ = io.WriteString(w, `{"data":[]}`)
|
||||
}))
|
||||
defer redirected.Close()
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, redirected.URL+"/captured", http.StatusFound)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
_, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), Request{
|
||||
UpstreamURL: server.URL,
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer must-not-leak",
|
||||
})
|
||||
require.EqualError(t, err, "upstream returned HTTP 302")
|
||||
assert.Zero(t, redirectedCalls.Load())
|
||||
}
|
||||
|
||||
func TestDiscoverEnforcesResponseLimit(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = io.WriteString(w, `{"data":[{"id":"llama3.2:latest"}]}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
discoverer := newWithTransport(http.DefaultTransport)
|
||||
discoverer.bodyLimit = 16
|
||||
_, err := discoverer.Discover(context.Background(), Request{UpstreamURL: server.URL})
|
||||
require.EqualError(t, err, "model discovery response is too large")
|
||||
}
|
||||
|
||||
func TestDiscoverEnforcesTimeout(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
select {
|
||||
case <-r.Context().Done():
|
||||
case <-time.After(time.Second):
|
||||
_, _ = io.WriteString(w, `{"data":[]}`)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
discoverer := newWithTransport(http.DefaultTransport)
|
||||
discoverer.timeout = 30 * time.Millisecond
|
||||
_, err := discoverer.Discover(context.Background(), Request{UpstreamURL: server.URL})
|
||||
require.EqualError(t, err, "model discovery timed out")
|
||||
}
|
||||
|
||||
func TestDiscoverHonorsSkipTLSVerification(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = io.WriteString(w, `{"data":[]}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
discoverer := New(nil)
|
||||
_, err := discoverer.Discover(context.Background(), Request{UpstreamURL: server.URL})
|
||||
require.EqualError(t, err, "model discovery request failed")
|
||||
|
||||
result, err := discoverer.Discover(context.Background(), Request{
|
||||
UpstreamURL: server.URL,
|
||||
SkipTLSVerify: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, result.Models)
|
||||
}
|
||||
|
||||
func TestDiscoverRejectsUnsafeRequestValues(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
request Request
|
||||
wantErr string
|
||||
}{
|
||||
{
|
||||
name: "unsupported scheme",
|
||||
request: Request{UpstreamURL: "file:///etc/passwd"},
|
||||
wantErr: "model discovery upstream URL is invalid",
|
||||
},
|
||||
{
|
||||
name: "URL credentials",
|
||||
request: Request{UpstreamURL: "http://user:pass@example.com"},
|
||||
wantErr: "model discovery upstream URL must not contain credentials",
|
||||
},
|
||||
{
|
||||
name: "missing hostname",
|
||||
request: Request{UpstreamURL: "http://:11434"},
|
||||
wantErr: "model discovery upstream URL is invalid",
|
||||
},
|
||||
{
|
||||
name: "incomplete auth header",
|
||||
request: Request{
|
||||
UpstreamURL: "http://example.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
},
|
||||
wantErr: "model discovery authentication header is incomplete",
|
||||
},
|
||||
{
|
||||
name: "header injection",
|
||||
request: Request{
|
||||
UpstreamURL: "http://example.com",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer safe\r\nX-Evil: yes",
|
||||
},
|
||||
wantErr: "model discovery authentication header is invalid",
|
||||
},
|
||||
{
|
||||
name: "unsupported auth header",
|
||||
request: Request{
|
||||
UpstreamURL: "http://example.com",
|
||||
AuthHeaderName: "X-API-Key",
|
||||
AuthHeaderValue: "secret",
|
||||
},
|
||||
wantErr: "model discovery authentication header is unsupported",
|
||||
},
|
||||
{
|
||||
name: "oversized URL",
|
||||
request: Request{
|
||||
UpstreamURL: "http://example.com/" + strings.Repeat("a", maxUpstreamURLBytes),
|
||||
},
|
||||
wantErr: "model discovery upstream URL is too long",
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), test.request)
|
||||
require.EqualError(t, err, test.wantErr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverRejectsInvalidPayloads(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
status int
|
||||
body string
|
||||
wantErr string
|
||||
}{
|
||||
{
|
||||
name: "malformed JSON",
|
||||
status: http.StatusOK,
|
||||
body: `{"data":`,
|
||||
wantErr: "upstream returned invalid OpenAI model-list JSON",
|
||||
},
|
||||
{
|
||||
name: "missing data",
|
||||
status: http.StatusOK,
|
||||
body: `{}`,
|
||||
wantErr: "upstream returned invalid OpenAI model-list JSON",
|
||||
},
|
||||
{
|
||||
name: "too many models",
|
||||
status: http.StatusOK,
|
||||
body: modelListJSON(maxModels + 1),
|
||||
wantErr: "upstream returned too many models",
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(test.status)
|
||||
_, _ = io.WriteString(w, test.body)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
_, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), Request{
|
||||
UpstreamURL: server.URL,
|
||||
})
|
||||
require.EqualError(t, err, test.wantErr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func modelListJSON(count int) string {
|
||||
var body strings.Builder
|
||||
body.WriteString(`{"data":[`)
|
||||
for i := range count {
|
||||
if i > 0 {
|
||||
body.WriteByte(',')
|
||||
}
|
||||
_, _ = fmt.Fprintf(&body, `{"id":"model-%d"}`, i)
|
||||
}
|
||||
body.WriteString(`]}`)
|
||||
return body.String()
|
||||
}
|
||||
@@ -271,6 +271,10 @@ func (c *testProxyController) SendServiceUpdateToCluster(_ context.Context, _ st
|
||||
// noop
|
||||
}
|
||||
|
||||
func (c *testProxyController) DiscoverModels(_ context.Context, _, _ string, _ *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
|
||||
return nil, nbproxy.ErrModelDiscoveryUnavailable
|
||||
}
|
||||
|
||||
func (c *testProxyController) GetOIDCValidationConfig() nbproxy.OIDCValidationConfig {
|
||||
return nbproxy.OIDCValidationConfig{}
|
||||
}
|
||||
|
||||
311
proxy/model_discovery_sync_test.go
Normal file
311
proxy/model_discovery_sync_test.go
Normal file
@@ -0,0 +1,311 @@
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/crowdsec"
|
||||
"github.com/netbirdio/netbird/proxy/internal/modeldiscovery"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
type stubModelDiscoverer struct {
|
||||
started chan modeldiscovery.Request
|
||||
release <-chan struct{}
|
||||
result modeldiscovery.Result
|
||||
err error
|
||||
}
|
||||
|
||||
func (s *stubModelDiscoverer) Discover(ctx context.Context, request modeldiscovery.Request) (modeldiscovery.Result, error) {
|
||||
if s.started != nil {
|
||||
select {
|
||||
case s.started <- request:
|
||||
case <-ctx.Done():
|
||||
return modeldiscovery.Result{}, ctx.Err()
|
||||
}
|
||||
}
|
||||
if s.release != nil {
|
||||
select {
|
||||
case <-s.release:
|
||||
case <-ctx.Done():
|
||||
return modeldiscovery.Result{}, ctx.Err()
|
||||
}
|
||||
}
|
||||
return s.result, s.err
|
||||
}
|
||||
|
||||
type modelDiscoverySyncStream struct {
|
||||
grpc.ClientStream
|
||||
ctx context.Context
|
||||
recv chan *proto.SyncMappingsResponse
|
||||
sent chan *proto.SyncMappingsRequest
|
||||
sendWait time.Duration
|
||||
sending atomic.Int32
|
||||
overlap atomic.Bool
|
||||
}
|
||||
|
||||
func newModelDiscoverySyncStream(ctx context.Context) *modelDiscoverySyncStream {
|
||||
return &modelDiscoverySyncStream{
|
||||
ctx: ctx,
|
||||
recv: make(chan *proto.SyncMappingsResponse, 16),
|
||||
sent: make(chan *proto.SyncMappingsRequest, 16),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *modelDiscoverySyncStream) Send(message *proto.SyncMappingsRequest) error {
|
||||
if s.sending.Add(1) != 1 {
|
||||
s.overlap.Store(true)
|
||||
}
|
||||
defer s.sending.Add(-1)
|
||||
if s.sendWait > 0 {
|
||||
time.Sleep(s.sendWait)
|
||||
}
|
||||
select {
|
||||
case s.sent <- message:
|
||||
return nil
|
||||
case <-s.ctx.Done():
|
||||
return s.ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *modelDiscoverySyncStream) Recv() (*proto.SyncMappingsResponse, error) {
|
||||
select {
|
||||
case message, ok := <-s.recv:
|
||||
if !ok {
|
||||
return nil, io.EOF
|
||||
}
|
||||
return message, nil
|
||||
case <-s.ctx.Done():
|
||||
return nil, s.ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *modelDiscoverySyncStream) Context() context.Context {
|
||||
return s.ctx
|
||||
}
|
||||
|
||||
func TestProxyCapabilitiesAdvertiseModelDiscovery(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := &Server{
|
||||
crowdsecRegistry: crowdsec.NewRegistry("", "", log.New().WithField("test", true)),
|
||||
}
|
||||
capabilities := server.proxyCapabilities()
|
||||
require.NotNil(t, capabilities.SupportsModelDiscovery)
|
||||
assert.True(t, capabilities.GetSupportsModelDiscovery())
|
||||
}
|
||||
|
||||
func TestExecuteModelDiscoveryMapsControlShape(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
discoverer := &stubModelDiscoverer{
|
||||
result: modeldiscovery.Result{
|
||||
Source: modeldiscovery.SourceOpenAIV1Models,
|
||||
Models: []modeldiscovery.Model{
|
||||
{ID: "llama3.2:latest", Label: "Llama 3.2"},
|
||||
},
|
||||
},
|
||||
}
|
||||
request := &proto.ModelDiscoveryRequest{
|
||||
RequestId: "request-1",
|
||||
UpstreamUrl: "http://ollama.internal:11434",
|
||||
AuthHeaderName: "Authorization",
|
||||
AuthHeaderValue: "Bearer secret",
|
||||
SkipTlsVerify: true,
|
||||
OllamaFallback: true,
|
||||
}
|
||||
|
||||
result := executeModelDiscovery(context.Background(), discoverer, request)
|
||||
assert.Equal(t, "request-1", result.GetRequestId())
|
||||
assert.Equal(t, modeldiscovery.SourceOpenAIV1Models, result.GetSource())
|
||||
require.Len(t, result.GetModels(), 1)
|
||||
assert.Equal(t, "llama3.2:latest", result.GetModels()[0].GetId())
|
||||
assert.Equal(t, "Llama 3.2", result.GetModels()[0].GetLabel())
|
||||
|
||||
discoverer.started = make(chan modeldiscovery.Request, 1)
|
||||
_ = executeModelDiscovery(context.Background(), discoverer, request)
|
||||
received := <-discoverer.started
|
||||
assert.Equal(t, request.GetUpstreamUrl(), received.UpstreamURL)
|
||||
assert.Equal(t, request.GetAuthHeaderName(), received.AuthHeaderName)
|
||||
assert.Equal(t, request.GetAuthHeaderValue(), received.AuthHeaderValue)
|
||||
assert.True(t, received.SkipTLSVerify)
|
||||
assert.True(t, received.AllowOllamaFallback)
|
||||
}
|
||||
|
||||
func TestExecuteModelDiscoveryReturnsSanitizedError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
result := executeModelDiscovery(context.Background(), &stubModelDiscoverer{
|
||||
err: errors.New("model discovery request failed"),
|
||||
}, &proto.ModelDiscoveryRequest{RequestId: "request-error"})
|
||||
assert.Equal(t, "request-error", result.GetRequestId())
|
||||
assert.Equal(t, "model discovery request failed", result.GetError())
|
||||
assert.Empty(t, result.GetModels())
|
||||
assert.Empty(t, result.GetSource())
|
||||
}
|
||||
|
||||
func TestHandleSyncMappingsStreamRunsDiscoveryOutOfBand(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
release := make(chan struct{})
|
||||
started := make(chan modeldiscovery.Request, 1)
|
||||
server := &Server{
|
||||
Logger: log.New(),
|
||||
routerReady: closedChan(),
|
||||
modelDiscoverer: &stubModelDiscoverer{
|
||||
started: started,
|
||||
release: release,
|
||||
result: modeldiscovery.Result{
|
||||
Source: modeldiscovery.SourceOpenAIV1Models,
|
||||
Models: []modeldiscovery.Model{{ID: "model-a", Label: "model-a"}},
|
||||
},
|
||||
},
|
||||
}
|
||||
stream := newModelDiscoverySyncStream(ctx)
|
||||
stream.sendWait = 10 * time.Millisecond
|
||||
|
||||
done := make(chan error, 1)
|
||||
initialSyncDone := true
|
||||
go func() {
|
||||
done <- server.handleSyncMappingsStream(ctx, stream, &initialSyncDone, time.Time{})
|
||||
}()
|
||||
|
||||
stream.recv <- &proto.SyncMappingsResponse{
|
||||
ModelDiscoveryRequest: &proto.ModelDiscoveryRequest{
|
||||
RequestId: "request-1",
|
||||
UpstreamUrl: "http://ollama.internal:11434",
|
||||
},
|
||||
}
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("model discovery did not start")
|
||||
}
|
||||
|
||||
// A normal mapping batch must still be acknowledged while the HTTP probe
|
||||
// is in flight.
|
||||
stream.recv <- &proto.SyncMappingsResponse{}
|
||||
select {
|
||||
case sent := <-stream.sent:
|
||||
assert.NotNil(t, sent.GetAck())
|
||||
assert.Nil(t, sent.GetModelDiscoveryResult())
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("mapping ack was blocked by model discovery")
|
||||
}
|
||||
|
||||
close(release)
|
||||
select {
|
||||
case sent := <-stream.sent:
|
||||
result := sent.GetModelDiscoveryResult()
|
||||
require.NotNil(t, result)
|
||||
assert.Equal(t, "request-1", result.GetRequestId())
|
||||
assert.Equal(t, modeldiscovery.SourceOpenAIV1Models, result.GetSource())
|
||||
assert.Nil(t, sent.GetAck())
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("model discovery result was not sent")
|
||||
}
|
||||
assert.False(t, stream.overlap.Load(), "acks and discovery results must use one serialized sender")
|
||||
|
||||
close(stream.recv)
|
||||
require.NoError(t, <-done)
|
||||
}
|
||||
|
||||
func TestHandleSyncMappingsStreamBoundsConcurrentDiscovery(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
release := make(chan struct{})
|
||||
started := make(chan modeldiscovery.Request, 8)
|
||||
server := &Server{
|
||||
Logger: log.New(),
|
||||
routerReady: closedChan(),
|
||||
modelDiscoverer: &stubModelDiscoverer{
|
||||
started: started,
|
||||
release: release,
|
||||
},
|
||||
}
|
||||
stream := newModelDiscoverySyncStream(ctx)
|
||||
|
||||
done := make(chan error, 1)
|
||||
initialSyncDone := true
|
||||
go func() {
|
||||
done <- server.handleSyncMappingsStream(ctx, stream, &initialSyncDone, time.Time{})
|
||||
}()
|
||||
|
||||
for i := range 5 {
|
||||
stream.recv <- &proto.SyncMappingsResponse{
|
||||
ModelDiscoveryRequest: &proto.ModelDiscoveryRequest{
|
||||
RequestId: fmt.Sprintf("request-%d", i),
|
||||
UpstreamUrl: "http://ollama.internal:11434",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
for range 4 {
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected four concurrent model discoveries")
|
||||
}
|
||||
}
|
||||
select {
|
||||
case sent := <-stream.sent:
|
||||
result := sent.GetModelDiscoveryResult()
|
||||
require.NotNil(t, result)
|
||||
assert.Equal(t, "model discovery is busy", result.GetError())
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("fifth discovery did not fail fast")
|
||||
}
|
||||
|
||||
close(release)
|
||||
for range 4 {
|
||||
select {
|
||||
case sent := <-stream.sent:
|
||||
require.NotNil(t, sent.GetModelDiscoveryResult())
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("in-flight model discovery did not complete")
|
||||
}
|
||||
}
|
||||
close(stream.recv)
|
||||
require.NoError(t, <-done)
|
||||
}
|
||||
|
||||
func TestHandleSyncMappingsStreamRejectsMixedDiscoveryMessage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
server := &Server{
|
||||
Logger: log.New(),
|
||||
routerReady: closedChan(),
|
||||
modelDiscoverer: &stubModelDiscoverer{},
|
||||
}
|
||||
stream := newModelDiscoverySyncStream(ctx)
|
||||
stream.recv <- &proto.SyncMappingsResponse{
|
||||
Mapping: []*proto.ProxyMapping{{Id: "mapping-1"}},
|
||||
ModelDiscoveryRequest: &proto.ModelDiscoveryRequest{
|
||||
RequestId: "request-1",
|
||||
},
|
||||
}
|
||||
close(stream.recv)
|
||||
|
||||
initialSyncDone := true
|
||||
err := server.handleSyncMappingsStream(ctx, stream, &initialSyncDone, time.Time{})
|
||||
require.EqualError(t, err, "model discovery message must not include mapping data")
|
||||
}
|
||||
126
proxy/server.go
126
proxy/server.go
@@ -57,6 +57,7 @@ import (
|
||||
proxymetrics "github.com/netbirdio/netbird/proxy/internal/metrics"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
mwbuiltin "github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
|
||||
"github.com/netbirdio/netbird/proxy/internal/modeldiscovery"
|
||||
"github.com/netbirdio/netbird/proxy/internal/netutil"
|
||||
"github.com/netbirdio/netbird/proxy/internal/proxy"
|
||||
"github.com/netbirdio/netbird/proxy/internal/restrict"
|
||||
@@ -78,6 +79,10 @@ type portRouter struct {
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
type providerModelDiscoverer interface {
|
||||
Discover(context.Context, modeldiscovery.Request) (modeldiscovery.Result, error)
|
||||
}
|
||||
|
||||
type Server struct {
|
||||
ctx context.Context
|
||||
mgmtClient proto.ProxyServiceClient
|
||||
@@ -100,16 +105,20 @@ type Server struct {
|
||||
// middlewareRegistry is the source of registered middleware factories.
|
||||
// Concrete middlewares register themselves through init().
|
||||
middlewareRegistry *middleware.Registry
|
||||
mainRouter *nbtcp.Router
|
||||
mainPort uint16
|
||||
udpMu sync.Mutex
|
||||
udpRelays map[types.ServiceID]*udprelay.Relay
|
||||
udpRelayWg sync.WaitGroup
|
||||
portMu sync.RWMutex
|
||||
portRouters map[uint16]*portRouter
|
||||
svcPorts map[types.ServiceID][]uint16
|
||||
lastMappings map[types.ServiceID]*proto.ProxyMapping
|
||||
portRouterWg sync.WaitGroup
|
||||
// modelDiscoverer executes explicit provider model-list probes on the
|
||||
// proxy's host network. Lazily constructed for normal servers; injectable
|
||||
// in focused control-stream tests.
|
||||
modelDiscoverer providerModelDiscoverer
|
||||
mainRouter *nbtcp.Router
|
||||
mainPort uint16
|
||||
udpMu sync.Mutex
|
||||
udpRelays map[types.ServiceID]*udprelay.Relay
|
||||
udpRelayWg sync.WaitGroup
|
||||
portMu sync.RWMutex
|
||||
portRouters map[uint16]*portRouter
|
||||
svcPorts map[types.ServiceID][]uint16
|
||||
lastMappings map[types.ServiceID]*proto.ProxyMapping
|
||||
portRouterWg sync.WaitGroup
|
||||
|
||||
// hijackTracker tracks hijacked connections (e.g. WebSocket upgrades)
|
||||
// so they can be closed during graceful shutdown, since http.Server.Shutdown
|
||||
@@ -1277,12 +1286,16 @@ func (s *Server) proxyCapabilities() *proto.ProxyCapabilities {
|
||||
privateCapability := s.Private
|
||||
// Always true: this build enforces ProxyMapping.private via the auth middleware.
|
||||
supportsPrivateService := true
|
||||
// Model discovery is handled only on SyncMappings. Management also gates
|
||||
// dispatch on the live connection using this capability.
|
||||
supportsModelDiscovery := true
|
||||
return &proto.ProxyCapabilities{
|
||||
SupportsCustomPorts: &s.SupportsCustomPorts,
|
||||
RequireSubdomain: &s.RequireSubdomain,
|
||||
SupportsCrowdsec: &supportsCrowdSec,
|
||||
Private: &privateCapability,
|
||||
SupportsPrivateService: &supportsPrivateService,
|
||||
SupportsModelDiscovery: &supportsModelDiscovery,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1349,7 +1362,8 @@ func isSyncUnimplemented(err error) bool {
|
||||
// handleSyncMappingsStream consumes batches from a bidirectional SyncMappings
|
||||
// stream, sending an ack after each batch is fully processed. Management waits
|
||||
// for the ack before sending the next batch, providing application-level
|
||||
// back-pressure.
|
||||
// back-pressure. Model discovery commands are out-of-band: they run with
|
||||
// bounded concurrency and return a correlated result instead of an ack.
|
||||
func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.ProxyService_SyncMappingsClient, initialSyncDone *bool, connectTime time.Time) error {
|
||||
select {
|
||||
case <-s.routerReady:
|
||||
@@ -1358,6 +1372,29 @@ func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.Prox
|
||||
}
|
||||
|
||||
tracker := s.newSnapshotTracker(initialSyncDone, connectTime)
|
||||
discoverer := s.modelDiscoverer
|
||||
if discoverer == nil {
|
||||
discoverer = modeldiscovery.New(s.Logger)
|
||||
}
|
||||
|
||||
const maxConcurrentDiscoveries = 4
|
||||
discoverySlots := make(chan struct{}, maxConcurrentDiscoveries)
|
||||
discoveryCtx, cancelDiscoveries := context.WithCancel(ctx)
|
||||
var discoveryWG sync.WaitGroup
|
||||
var sendMu sync.Mutex
|
||||
defer func() {
|
||||
cancelDiscoveries()
|
||||
discoveryWG.Wait()
|
||||
}()
|
||||
|
||||
// gRPC permits one concurrent sender and one concurrent receiver, but not
|
||||
// multiple senders. Mapping acks and asynchronous discovery results share
|
||||
// this serialized send path.
|
||||
send := func(message *proto.SyncMappingsRequest) error {
|
||||
sendMu.Lock()
|
||||
defer sendMu.Unlock()
|
||||
return stream.Send(message)
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
@@ -1372,6 +1409,46 @@ func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.Prox
|
||||
return fmt.Errorf("receive msg: %w", err)
|
||||
}
|
||||
|
||||
if discovery := msg.GetModelDiscoveryRequest(); discovery != nil {
|
||||
if len(msg.GetMapping()) != 0 || msg.GetInitialSyncComplete() {
|
||||
return errors.New("model discovery message must not include mapping data")
|
||||
}
|
||||
|
||||
select {
|
||||
case discoverySlots <- struct{}{}:
|
||||
discoveryWG.Add(1)
|
||||
go func() {
|
||||
defer discoveryWG.Done()
|
||||
defer func() { <-discoverySlots }()
|
||||
|
||||
result := executeModelDiscovery(discoveryCtx, discoverer, discovery)
|
||||
if discoveryCtx.Err() != nil {
|
||||
return
|
||||
}
|
||||
if err := send(&proto.SyncMappingsRequest{
|
||||
Msg: &proto.SyncMappingsRequest_ModelDiscoveryResult{
|
||||
ModelDiscoveryResult: result,
|
||||
},
|
||||
}); err != nil {
|
||||
s.Logger.WithError(err).Debug("failed to send model discovery result")
|
||||
}
|
||||
}()
|
||||
default:
|
||||
result := &proto.ModelDiscoveryResult{
|
||||
RequestId: discovery.GetRequestId(),
|
||||
Error: "model discovery is busy",
|
||||
}
|
||||
if err := send(&proto.SyncMappingsRequest{
|
||||
Msg: &proto.SyncMappingsRequest_ModelDiscoveryResult{
|
||||
ModelDiscoveryResult: result,
|
||||
},
|
||||
}); err != nil {
|
||||
return fmt.Errorf("send model discovery busy result: %w", err)
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
batchStart := time.Now()
|
||||
s.Logger.Debug("Received mapping update, starting processing")
|
||||
if err := s.processMappingsGuarded(ctx, msg.GetMapping()); err != nil {
|
||||
@@ -1380,7 +1457,7 @@ func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.Prox
|
||||
s.Logger.Debug("Processing mapping update completed")
|
||||
tracker.recordBatch(ctx, s, msg.GetMapping(), msg.GetInitialSyncComplete(), batchStart)
|
||||
|
||||
if err := stream.Send(&proto.SyncMappingsRequest{
|
||||
if err := send(&proto.SyncMappingsRequest{
|
||||
Msg: &proto.SyncMappingsRequest_Ack{Ack: &proto.SyncMappingsAck{}},
|
||||
}); err != nil {
|
||||
return fmt.Errorf("send ack: %w", err)
|
||||
@@ -1389,6 +1466,31 @@ func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.Prox
|
||||
}
|
||||
}
|
||||
|
||||
func executeModelDiscovery(ctx context.Context, discoverer providerModelDiscoverer, request *proto.ModelDiscoveryRequest) *proto.ModelDiscoveryResult {
|
||||
result := &proto.ModelDiscoveryResult{RequestId: request.GetRequestId()}
|
||||
discovered, err := discoverer.Discover(ctx, modeldiscovery.Request{
|
||||
UpstreamURL: request.GetUpstreamUrl(),
|
||||
AuthHeaderName: request.GetAuthHeaderName(),
|
||||
AuthHeaderValue: request.GetAuthHeaderValue(),
|
||||
SkipTLSVerify: request.GetSkipTlsVerify(),
|
||||
AllowOllamaFallback: request.GetOllamaFallback(),
|
||||
})
|
||||
if err != nil {
|
||||
result.Error = err.Error()
|
||||
return result
|
||||
}
|
||||
|
||||
result.Source = discovered.Source
|
||||
result.Models = make([]*proto.ModelDiscoveryModel, 0, len(discovered.Models))
|
||||
for _, model := range discovered.Models {
|
||||
result.Models = append(result.Models, &proto.ModelDiscoveryModel{
|
||||
Id: model.ID,
|
||||
Label: model.Label,
|
||||
})
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// snapshotTracker accumulates service IDs during the initial snapshot and
|
||||
// finalises sync state when the complete flag arrives. Used by both
|
||||
// handleMappingStream and handleSyncMappingsStream so metric emission and
|
||||
|
||||
Reference in New Issue
Block a user