summaryrefslogtreecommitdiff
path: root/internal/provider/forwarder_test.go
blob: 2ae81afa7cb8b2c01c2cfa8e92709ddbb2a59306 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
package provider

import (
	"encoding/json"
	"testing"

	"aigw/internal/domain"
)

func TestRewriteRequestForcesOpenAIStreamUsage(t *testing.T) {
	result, err := rewriteRequest([]byte(`{"model":"public/model","stream":true,"stream_options":{"other":true}}`), "upstream/model", domain.ProtocolOpenAI)
	if err != nil {
		t.Fatal(err)
	}
	var body struct {
		Model         string         `json:"model"`
		StreamOptions map[string]any `json:"stream_options"`
	}
	if err := json.Unmarshal(result, &body); err != nil {
		t.Fatal(err)
	}
	if body.Model != "upstream/model" || body.StreamOptions["include_usage"] != true || body.StreamOptions["other"] != true {
		t.Fatalf("unexpected rewritten body: %s", result)
	}
}

func TestRewriteRequestDoesNotAddStreamOptionsToAnthropic(t *testing.T) {
	result, err := rewriteRequest([]byte(`{"model":"public/model","stream":true}`), "upstream/model", domain.ProtocolAnthropic)
	if err != nil {
		t.Fatal(err)
	}
	var body map[string]json.RawMessage
	if err := json.Unmarshal(result, &body); err != nil {
		t.Fatal(err)
	}
	if _, exists := body["stream_options"]; exists {
		t.Fatalf("unexpected OpenAI stream options in Anthropic request: %s", result)
	}
}

func TestResponsesWireAPIUsesResponsesEndpointWithoutChatStreamOptions(t *testing.T) {
	result, err := rewriteRequestWithWireAPI([]byte(`{"model":"public/model","input":"hello","stream":true}`), "gpt-upstream", domain.ProtocolOpenAIResponses, "responses")
	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"]) != `"gpt-upstream"` {
		t.Fatalf("model was not rewritten: %s", result)
	}
	if _, exists := body["stream_options"]; exists {
		t.Fatalf("Responses request contains Chat Completions stream options: %s", result)
	}
	provider := domain.Provider{BaseURL: "https://example.test", Protocol: domain.ProtocolOpenAI, WireAPI: "responses"}
	if got := endpointURL(provider, domain.ProtocolOpenAIResponses); got != "https://example.test/responses" {
		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)
	}
}