summaryrefslogtreecommitdiff
path: root/internal/provider
diff options
context:
space:
mode:
Diffstat (limited to 'internal/provider')
-rw-r--r--internal/provider/forwarder.go19
-rw-r--r--internal/provider/forwarder_test.go39
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)
+ }
+}