diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/provider/forwarder_test.go | 18 |
1 files changed, 18 insertions, 0 deletions
diff --git a/internal/provider/forwarder_test.go b/internal/provider/forwarder_test.go index e9751ce..2ae81af 100644 --- a/internal/provider/forwarder_test.go +++ b/internal/provider/forwarder_test.go @@ -58,3 +58,21 @@ func TestResponsesWireAPIUsesResponsesEndpointWithoutChatStreamOptions(t *testin t.Fatalf("endpoint URL = %q", got) } } + +func TestEmbeddingsWireAPIUsesEmbeddingsEndpoint(t *testing.T) { + result, err := rewriteRequestWithWireAPI([]byte(`{"model":"public/model","input":["one","two"]}`), "embedding-upstream", domain.ProtocolOpenAIEmbeddings, "embeddings") + if err != nil { + t.Fatal(err) + } + var body map[string]json.RawMessage + if err := json.Unmarshal(result, &body); err != nil { + t.Fatal(err) + } + if string(body["model"]) != `"embedding-upstream"` || string(body["input"]) != `["one","two"]` { + t.Fatalf("unexpected rewritten body: %s", result) + } + provider := domain.Provider{BaseURL: "https://example.test/v1", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings"} + if got := endpointURL(provider, domain.ProtocolOpenAIEmbeddings); got != "https://example.test/v1/embeddings" { + t.Fatalf("endpoint URL = %q", got) + } +} |
