mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-13 04:13:35 +08:00
Honor explicit model URL overrides for NVIDIA endpoints (#18152)
Some NVIDIA hosted models (e.g. meta/llama-3.2-11b-vision-instruct )
expose a full
endpoint URL per model that does not follow the normal base_url +
url_suffix
assembly. Previously the Go driver always called {base}/chat/completions
, so chat
requests for these vision models hit the wrong endpoint and failed.
This PR adds an optional per-model url field in conf/models/nvidia.json
. When
present, every NVIDIA driver request (chat, streaming chat, embedding,
rerank, model
listing) uses it directly; otherwise the standard assembly is unchanged.
This commit is contained in:
@@ -72,6 +72,7 @@
|
||||
"name": "meta/llama-3.2-11b-vision-instruct",
|
||||
"content_length": 131072,
|
||||
"max_output": 8192,
|
||||
"url":"https://integrate.api.nvidia.com/v1/meta/llama-3.2-11b-vision-instruct",
|
||||
"model_types": [
|
||||
"chat",
|
||||
"vision"
|
||||
@@ -97,6 +98,7 @@
|
||||
"name": "meta/llama-3.2-90b-vision-instruct",
|
||||
"content_length": 131072,
|
||||
"max_output": 8192,
|
||||
"url":"https://integrate.api.nvidia.com/v1/meta/llama-3.2-90b-vision-instruct",
|
||||
"model_types": [
|
||||
"chat",
|
||||
"vision"
|
||||
|
||||
@@ -166,6 +166,7 @@ type Model struct {
|
||||
Thinking *ModelThinking `json:"thinking"`
|
||||
Tools *ModelTools `json:"tools"`
|
||||
Class *string `json:"class"`
|
||||
URL string `json:"url"`
|
||||
MaxDimension *int `json:"max_dimension"` // used by embedding models
|
||||
MaxBatchSize *int `json:"max_batch_size"` // used by embedding models
|
||||
Dimensions []int `json:"dimensions"`
|
||||
|
||||
@@ -64,6 +64,30 @@ func (n *NvidiaModel) Name() string {
|
||||
return "nvidia"
|
||||
}
|
||||
|
||||
// resolveEndpoint returns the endpoint URL for modelName. When the
|
||||
// model's preset config in conf/models/nvidia.json carries an explicit
|
||||
// url, that full endpoint is used as-is because it does not follow the
|
||||
// standard baseURL + suffix assembly; otherwise the usual assembly
|
||||
// applies. Every endpoint (chat, embedding, rerank, models list) goes
|
||||
// through this path so the override is honored uniformly.
|
||||
func (n *NvidiaModel) resolveEndpoint(modelName string, apiConfig *APIConfig, suffix string) (string, error) {
|
||||
if modelName != "" {
|
||||
if pm := GetProviderManager(); pm != nil {
|
||||
if provider := pm.FindProvider("NVIDIA"); provider != nil {
|
||||
if model := pm.FindModel(provider, modelName); model != nil && model.URL != "" {
|
||||
return strings.TrimSuffix(model.URL, "/"), nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return fmt.Sprintf("%s/%s", strings.TrimRight(resolvedBaseURL, "/"), strings.TrimLeft(suffix, "/")), nil
|
||||
}
|
||||
|
||||
func (n *NvidiaModel) ChatWithMessages(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) {
|
||||
if err := n.baseModel.APIConfigCheck(apiConfig); err != nil {
|
||||
return nil, err
|
||||
@@ -73,11 +97,10 @@ func (n *NvidiaModel) ChatWithMessages(ctx context.Context, modelName string, me
|
||||
return nil, fmt.Errorf("messages is empty")
|
||||
}
|
||||
|
||||
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
|
||||
baseURL, err := n.resolveEndpoint(modelName, apiConfig, n.baseModel.URLSuffix.Chat)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
baseURL := fmt.Sprintf("%s/%s", resolvedBaseURL, n.baseModel.URLSuffix.Chat)
|
||||
reqBody := buildRequestBody(chatModelConfig, modelName, messages, false)
|
||||
|
||||
if chatModelConfig != nil {
|
||||
@@ -108,11 +131,10 @@ func (n *NvidiaModel) ChatStreamlyWithSender(ctx context.Context, modelName stri
|
||||
return fmt.Errorf("messages is empty")
|
||||
}
|
||||
|
||||
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
|
||||
baseURL, err := n.resolveEndpoint(modelName, apiConfig, n.baseModel.URLSuffix.Chat)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
baseURL := fmt.Sprintf("%s/%s", resolvedBaseURL, n.baseModel.URLSuffix.Chat)
|
||||
reqBody := buildRequestBody(modelConfig, modelName, messages, true)
|
||||
|
||||
if modelConfig != nil {
|
||||
@@ -152,13 +174,11 @@ func (n *NvidiaModel) Embed(ctx context.Context, modelName *string, texts []stri
|
||||
return nil, fmt.Errorf("model name is required")
|
||||
}
|
||||
|
||||
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
|
||||
baseURL, err := n.resolveEndpoint(*modelName, apiConfig, n.baseModel.URLSuffix.Embedding)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
baseURL := fmt.Sprintf("%s/%s", strings.TrimSuffix(resolvedBaseURL, "/"), n.baseModel.URLSuffix.Embedding)
|
||||
|
||||
reqBody := map[string]interface{}{
|
||||
"model": *modelName,
|
||||
"input": texts,
|
||||
@@ -264,13 +284,11 @@ func (n *NvidiaModel) Rerank(ctx context.Context, modelName *string, query strin
|
||||
return nil, fmt.Errorf("model name is required")
|
||||
}
|
||||
|
||||
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
|
||||
baseURL, err := n.resolveEndpoint(*modelName, apiConfig, n.baseModel.URLSuffix.Rerank)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
baseURL := fmt.Sprintf("%s/%s", strings.TrimSuffix(resolvedBaseURL, "/"), n.baseModel.URLSuffix.Rerank)
|
||||
|
||||
topN := len(documents)
|
||||
if rerankConfig != nil && rerankConfig.TopN > 0 && rerankConfig.TopN < topN {
|
||||
topN = rerankConfig.TopN
|
||||
@@ -376,11 +394,10 @@ func (n *NvidiaModel) ListModels(ctx context.Context, apiConfig *APIConfig) ([]L
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
|
||||
baseURL, err := n.resolveEndpoint("", apiConfig, n.baseModel.URLSuffix.Models)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
baseURL := fmt.Sprintf("%s/%s", strings.TrimRight(resolvedBaseURL, "/"), strings.TrimLeft(n.baseModel.URLSuffix.Models, "/"))
|
||||
|
||||
modelListCtx, cancel := context.WithTimeout(ctx, nonStreamCallTimeout)
|
||||
defer cancel()
|
||||
@@ -422,7 +439,7 @@ func (n *NvidiaModel) ListModels(ctx context.Context, apiConfig *APIConfig) ([]L
|
||||
provider = pm.FindProvider("NVIDIA")
|
||||
}
|
||||
models := parseNvidiaModelList(modelList, provider)
|
||||
if n.usesHostedCatalog(resolvedBaseURL) {
|
||||
if n.usesHostedCatalog(baseURL) {
|
||||
catalogCtx, catalogCancel := context.WithTimeout(ctx, nonStreamCallTimeout)
|
||||
catalog, catalogErr := n.fetchHostedCatalog(catalogCtx)
|
||||
catalogCancel()
|
||||
|
||||
@@ -229,6 +229,162 @@ func TestNvidiaCatalogResourceRejectsMalformedDeprecation(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func withNvidiaProviderManager(t *testing.T, models []*Model) {
|
||||
t.Helper()
|
||||
saved := providerManager
|
||||
providerManager = &ProviderManager{
|
||||
Providers: []Provider{{
|
||||
Name: "NVIDIA",
|
||||
URL: map[string]string{"default": "https://integrate.api.nvidia.com/v1"},
|
||||
Models: models,
|
||||
}},
|
||||
}
|
||||
t.Cleanup(func() { providerManager = saved })
|
||||
}
|
||||
|
||||
func TestNvidiaResolveEndpointUsesModelURL(t *testing.T) {
|
||||
const modelURL = "https://integrate.api.nvidia.com/v1/meta/llama-3.2-11b-vision-instruct"
|
||||
withNvidiaProviderManager(t, []*Model{{
|
||||
Name: "meta/llama-3.2-11b-vision-instruct",
|
||||
URL: modelURL,
|
||||
}})
|
||||
|
||||
driver := NewNvidiaModel(
|
||||
map[string]string{"default": "https://integrate.api.nvidia.com/v1"},
|
||||
URLSuffix{Chat: "chat/completions"},
|
||||
)
|
||||
got, err := driver.resolveEndpoint("meta/llama-3.2-11b-vision-instruct", &APIConfig{}, "chat/completions")
|
||||
if err != nil {
|
||||
t.Fatalf("resolveEndpoint() error = %v", err)
|
||||
}
|
||||
if got != modelURL {
|
||||
t.Fatalf("resolveEndpoint() = %q, want %q", got, modelURL)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNvidiaResolveEndpointFallsBackToAssembly(t *testing.T) {
|
||||
withNvidiaProviderManager(t, []*Model{{Name: "meta/llama-3.1-8b-instruct"}})
|
||||
|
||||
driver := NewNvidiaModel(
|
||||
map[string]string{"default": "https://integrate.api.nvidia.com/v1/"},
|
||||
URLSuffix{Chat: "chat/completions", Embedding: "embeddings"},
|
||||
)
|
||||
got, err := driver.resolveEndpoint("meta/llama-3.1-8b-instruct", &APIConfig{}, "chat/completions")
|
||||
if err != nil {
|
||||
t.Fatalf("resolveEndpoint() error = %v", err)
|
||||
}
|
||||
if want := "https://integrate.api.nvidia.com/v1/chat/completions"; got != want {
|
||||
t.Fatalf("resolveEndpoint() = %q, want %q", got, want)
|
||||
}
|
||||
|
||||
got, err = driver.resolveEndpoint("meta/llama-3.1-8b-instruct", &APIConfig{}, "embeddings")
|
||||
if err != nil {
|
||||
t.Fatalf("resolveEndpoint() error = %v", err)
|
||||
}
|
||||
if want := "https://integrate.api.nvidia.com/v1/embeddings"; got != want {
|
||||
t.Fatalf("resolveEndpoint() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNvidiaResolveEndpointIgnoresManagerWhenNil(t *testing.T) {
|
||||
saved := providerManager
|
||||
providerManager = nil
|
||||
defer func() { providerManager = saved }()
|
||||
|
||||
driver := NewNvidiaModel(
|
||||
map[string]string{"default": "https://integrate.api.nvidia.com/v1"},
|
||||
URLSuffix{},
|
||||
)
|
||||
got, err := driver.resolveEndpoint("meta/llama-3.2-11b-vision-instruct", &APIConfig{}, "chat/completions")
|
||||
if err != nil {
|
||||
t.Fatalf("resolveEndpoint() error = %v", err)
|
||||
}
|
||||
if want := "https://integrate.api.nvidia.com/v1/chat/completions"; got != want {
|
||||
t.Fatalf("resolveEndpoint() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNvidiaChatUsesModelSpecificURL(t *testing.T) {
|
||||
withSSRFBypass(t)
|
||||
const apiKey = "nvapi-test"
|
||||
const modelName = "meta/llama-3.2-11b-vision-instruct"
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/v1/meta/llama-3.2-11b-vision-instruct" {
|
||||
t.Fatalf("path = %s, want model-specific chat endpoint", r.URL.Path)
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"choices": []map[string]interface{}{{
|
||||
"message": map[string]interface{}{"role": "assistant", "content": "pong"},
|
||||
"finish_reason": "stop",
|
||||
}},
|
||||
"usage": map[string]interface{}{"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
withNvidiaProviderManager(t, []*Model{{
|
||||
Name: modelName,
|
||||
URL: server.URL + "/v1/meta/llama-3.2-11b-vision-instruct",
|
||||
}})
|
||||
driver := NewNvidiaModel(
|
||||
map[string]string{"default": server.URL + "/v1"},
|
||||
URLSuffix{Chat: "chat/completions"},
|
||||
)
|
||||
resp, err := driver.ChatWithMessages(
|
||||
context.Background(),
|
||||
modelName,
|
||||
[]Message{{Role: "user", Content: "hi"}},
|
||||
&APIConfig{ApiKey: ptr(apiKey)},
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("ChatWithMessages() error = %v", err)
|
||||
}
|
||||
if resp.Answer == nil || *resp.Answer != "pong" {
|
||||
t.Fatalf("answer = %v, want pong", resp.Answer)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNvidiaEmbedUsesModelSpecificURL(t *testing.T) {
|
||||
withSSRFBypass(t)
|
||||
const apiKey = "nvapi-test"
|
||||
const modelName = "nvidia/nv-embed-v1"
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/v1/embeddings/nvidia/nv-embed-v1" {
|
||||
t.Fatalf("path = %s, want model-specific embed endpoint", r.URL.Path)
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"data": []map[string]interface{}{{"index": 0, "embedding": []float64{0.1, 0.2}}},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
withNvidiaProviderManager(t, []*Model{{
|
||||
Name: modelName,
|
||||
URL: server.URL + "/v1/embeddings/nvidia/nv-embed-v1",
|
||||
}})
|
||||
driver := NewNvidiaModel(
|
||||
map[string]string{"default": server.URL + "/v1"},
|
||||
URLSuffix{Embedding: "embeddings"},
|
||||
)
|
||||
namePtr := modelName
|
||||
got, err := driver.Embed(
|
||||
context.Background(),
|
||||
&namePtr,
|
||||
[]string{"hello"},
|
||||
&APIConfig{ApiKey: ptr(apiKey)},
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("Embed() error = %v", err)
|
||||
}
|
||||
if len(got) != 1 || len(got[0].Embedding) != 2 {
|
||||
t.Fatalf("embedding = %#v, want 1 vector of length 2", got)
|
||||
}
|
||||
}
|
||||
|
||||
func ptr[T any](value T) *T {
|
||||
return &value
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user