- 为图像协议放宽 JSON 响应上限并显式处理读取与超限错误 - 合并基础模型与平台能力,避免运行时丢失分辨率和比例约束 - 固化稳定 Gemini 图像模型的能力、计价和 preview 兼容别名
126 lines
4.6 KiB
Go
126 lines
4.6 KiB
Go
package clients
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"io"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
|
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
|
)
|
|
|
|
func TestDecodeHTTPResponseForProtocolCapturesOfficialErrorWire(t *testing.T) {
|
|
response := &http.Response{
|
|
StatusCode: http.StatusTooManyRequests,
|
|
Status: "429 Too Many Requests",
|
|
Header: http.Header{
|
|
"Content-Type": {"application/json"},
|
|
"X-Request-Id": {"req_1"},
|
|
"X-Ratelimit-Reset-Requests": {"1s"},
|
|
"Set-Cookie": {"secret=must-not-leak"},
|
|
"Connection": {"keep-alive"},
|
|
},
|
|
Body: io.NopCloser(strings.NewReader(`{"error":{"message":"slow down","future_field":true}}`)),
|
|
}
|
|
_, wire, err := decodeHTTPResponseForProtocol(response, ProtocolOpenAIResponses)
|
|
if err == nil {
|
|
t.Fatal("expected upstream error")
|
|
}
|
|
if wire == nil || wire.StatusCode != http.StatusTooManyRequests || wire.Protocol != ProtocolOpenAIResponses {
|
|
t.Fatalf("unexpected wire response: %+v", wire)
|
|
}
|
|
if wire.Headers["X-Request-Id"][0] != "req_1" || wire.Headers["X-Ratelimit-Reset-Requests"][0] != "1s" {
|
|
t.Fatalf("official response headers were lost: %+v", wire.Headers)
|
|
}
|
|
if _, ok := wire.Headers["Set-Cookie"]; ok {
|
|
t.Fatalf("sensitive header leaked: %+v", wire.Headers)
|
|
}
|
|
if _, ok := wire.Headers["Connection"]; ok {
|
|
t.Fatalf("hop-by-hop header leaked: %+v", wire.Headers)
|
|
}
|
|
if ErrorWireResponse(err) != wire {
|
|
t.Fatal("client error did not retain its wire response")
|
|
}
|
|
}
|
|
|
|
func TestDecodeHTTPResponseForProtocolAllowsLargeImagePayload(t *testing.T) {
|
|
payload := `{"data":"` + strings.Repeat("a", int(defaultMaxJSONResponseBytes)) + `"}`
|
|
response := &http.Response{
|
|
StatusCode: http.StatusOK,
|
|
Status: "200 OK",
|
|
Header: http.Header{"Content-Type": {"application/json"}},
|
|
Body: io.NopCloser(strings.NewReader(payload)),
|
|
}
|
|
|
|
result, wire, err := decodeHTTPResponseForProtocol(response, ProtocolGeminiGenerateContent)
|
|
if err != nil {
|
|
t.Fatalf("large Gemini image response failed: %v", err)
|
|
}
|
|
if got, _ := result["data"].(string); len(got) != int(defaultMaxJSONResponseBytes) {
|
|
t.Fatalf("large Gemini image payload length = %d", len(got))
|
|
}
|
|
if wire == nil || len(wire.RawJSON) != len(payload) {
|
|
t.Fatalf("large Gemini wire payload was truncated: %+v", wire)
|
|
}
|
|
}
|
|
|
|
func TestDecodeHTTPResponseForProtocolRejectsOversizedDefaultPayload(t *testing.T) {
|
|
payload := `{"data":"` + strings.Repeat("a", int(defaultMaxJSONResponseBytes)) + `"}`
|
|
response := &http.Response{
|
|
StatusCode: http.StatusOK,
|
|
Status: "200 OK",
|
|
Header: http.Header{"Content-Type": {"application/json"}},
|
|
Body: io.NopCloser(strings.NewReader(payload)),
|
|
}
|
|
|
|
_, _, err := decodeHTTPResponseForProtocol(response, ProtocolOpenAIResponses)
|
|
if err == nil {
|
|
t.Fatal("expected oversized default response to fail")
|
|
}
|
|
clientErr, ok := err.(*ClientError)
|
|
if !ok || clientErr.Code != "response_too_large" {
|
|
t.Fatalf("unexpected oversized response error: %T %v", err, err)
|
|
}
|
|
}
|
|
|
|
func TestGeminiNativeStreamPreservesOfficialEvents(t *testing.T) {
|
|
var requestedPath string
|
|
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
requestedPath = r.URL.RequestURI()
|
|
w.Header().Set("Content-Type", "text/event-stream")
|
|
w.Header().Set("X-Goog-Request-Id", "goog-stream-1")
|
|
_, _ = io.WriteString(w, "data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"hello\"}]}}],\"futureOfficialField\":true}\n\n")
|
|
}))
|
|
defer upstream.Close()
|
|
|
|
var events []StreamDeltaEvent
|
|
response, err := (GeminiClient{HTTPClient: upstream.Client()}).Run(context.Background(), Request{
|
|
Kind: "chat.completions", ModelType: "text", Model: "gemini-test", Stream: true,
|
|
Body: map[string]any{"prompt": "hello"},
|
|
Candidate: store.RuntimeModelCandidate{
|
|
Provider: "gemini", BaseURL: upstream.URL, ProviderModelName: "gemini-test",
|
|
Credentials: map[string]any{"apiKey": "test-key"},
|
|
},
|
|
StreamDelta: func(event StreamDeltaEvent) error {
|
|
events = append(events, event)
|
|
return nil
|
|
},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("Gemini stream failed: %v", err)
|
|
}
|
|
if !strings.Contains(requestedPath, ":streamGenerateContent") || !strings.Contains(requestedPath, "alt=sse") {
|
|
t.Fatalf("unexpected Gemini stream endpoint: %s", requestedPath)
|
|
}
|
|
if len(events) != 1 || events[0].WireProtocol != ProtocolGeminiGenerateContent || events[0].Event["futureOfficialField"] != true {
|
|
raw, _ := json.Marshal(events)
|
|
t.Fatalf("official Gemini event was not preserved: %s", raw)
|
|
}
|
|
if response.Wire == nil || response.Wire.Headers["X-Goog-Request-Id"][0] != "goog-stream-1" {
|
|
t.Fatalf("Gemini stream wire metadata was lost: %+v", response.Wire)
|
|
}
|
|
}
|