diff options
Diffstat (limited to 'internal/provider')
| -rw-r--r-- | internal/provider/forwarder.go | 19 | ||||
| -rw-r--r-- | internal/provider/forwarder_test.go | 39 |
2 files changed, 57 insertions, 1 deletions
diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go index a9d5734..0d40bb1 100644 --- a/internal/provider/forwarder.go +++ b/internal/provider/forwarder.go @@ -49,7 +49,7 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque if err := ctx.Err(); err != nil { return Result{Attempts: i}, err } - body, err := rewriteModel(originalBody, route.UpstreamModel) + body, err := rewriteRequest(originalBody, route.UpstreamModel, protocol) if err != nil { return Result{Attempts: i}, err } @@ -80,12 +80,29 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque } func rewriteModel(body []byte, upstreamModel string) ([]byte, error) { + return rewriteRequest(body, upstreamModel, "") +} + +func rewriteRequest(body []byte, upstreamModel string, protocol domain.Protocol) ([]byte, error) { var object map[string]json.RawMessage if err := json.Unmarshal(body, &object); err != nil { return nil, fmt.Errorf("decode request body: %w", err) } encoded, _ := json.Marshal(upstreamModel) object["model"] = encoded + if protocol == domain.ProtocolOpenAI { + var stream bool + _ = json.Unmarshal(object["stream"], &stream) + if stream { + var options map[string]json.RawMessage + _ = json.Unmarshal(object["stream_options"], &options) + if options == nil { + options = map[string]json.RawMessage{} + } + options["include_usage"] = json.RawMessage("true") + object["stream_options"], _ = json.Marshal(options) + } + } result, err := json.Marshal(object) if err != nil { return nil, fmt.Errorf("encode upstream request: %w", err) diff --git a/internal/provider/forwarder_test.go b/internal/provider/forwarder_test.go new file mode 100644 index 0000000..2e9b4a1 --- /dev/null +++ b/internal/provider/forwarder_test.go @@ -0,0 +1,39 @@ +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) + } +} |
