Files
oci-portal/internal/oci/genai_responses_test.go
2026-07-22 16:51:23 +08:00

171 lines
5.1 KiB
Go

package oci
import (
"bytes"
"context"
"errors"
"io"
"net/http"
"strings"
"testing"
"time"
"github.com/oracle/oci-go-sdk/v65/common"
)
// staticDispatcher 是非 *http.Client 的自定义 dispatcher,用于降级分支。
type staticDispatcher struct{}
func (staticDispatcher) Do(*http.Request) (*http.Response, error) { return nil, nil }
func TestDispatcherWithTimeout(t *testing.T) {
tr := &http.Transport{}
tests := []struct {
name string
in common.HTTPRequestDispatcher
timeout time.Duration
check func(t *testing.T, out common.HTTPRequestDispatcher)
}{
{
name: "http.Client 换总超时并保留 Transport", in: &http.Client{Transport: tr, Timeout: 60 * time.Second},
timeout: 300 * time.Second,
check: func(t *testing.T, out common.HTTPRequestDispatcher) {
hc, ok := out.(*http.Client)
if !ok || hc.Timeout != 300*time.Second || hc.Transport != tr {
t.Fatalf("期望拷贝 client 且 Timeout=300s、Transport 保留, 得到 %#v", out)
}
},
},
{
name: "timeout=0 表示无总超时", in: &http.Client{Timeout: 60 * time.Second}, timeout: 0,
check: func(t *testing.T, out common.HTTPRequestDispatcher) {
if hc := out.(*http.Client); hc.Timeout != 0 {
t.Fatalf("期望 Timeout=0, 得到 %v", hc.Timeout)
}
},
},
{
name: "非 http.Client 原样返回", in: staticDispatcher{}, timeout: 300 * time.Second,
check: func(t *testing.T, out common.HTTPRequestDispatcher) {
if _, ok := out.(staticDispatcher); !ok {
t.Fatalf("期望原样返回自定义 dispatcher, 得到 %#v", out)
}
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { tt.check(t, dispatcherWithTimeout(tt.in, tt.timeout)) })
}
}
// ctxReader 模拟真实 HTTP body:请求 ctx 取消后读即失败。
type ctxReader struct {
ctx context.Context
r io.Reader
}
func (c *ctxReader) Read(p []byte) (int, error) {
if err := c.ctx.Err(); err != nil {
return 0, err
}
return c.r.Read(p)
}
// fakeCaller 在 delay 后返回响应;期间 ctx 取消则按真实行为返回 ctx 错误。
type fakeCaller struct {
delay time.Duration
body string
gotCtx context.Context
failErr error
}
func (f *fakeCaller) Call(ctx context.Context, _ *http.Request) (*http.Response, error) {
f.gotCtx = ctx
if f.failErr != nil {
return nil, f.failErr
}
select {
case <-time.After(f.delay):
body := &ctxReader{ctx: ctx, r: strings.NewReader(f.body)}
return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(body)}, nil
case <-ctx.Done():
return nil, ctx.Err()
}
}
func TestCallWithHeaderBudget(t *testing.T) {
req, _ := http.NewRequest(http.MethodPost, "/actions/v1/responses", nil)
t.Run("预算内返回响应头后长读不受预算影响", func(t *testing.T) {
fc := &fakeCaller{delay: 0, body: "data: hello"}
stream, err := callWithHeaderBudget(context.Background(), fc, req, 30*time.Millisecond)
if err != nil {
t.Fatalf("callWithHeaderBudget = %v", err)
}
defer stream.Close()
time.Sleep(90 * time.Millisecond) // 远超预算,验证响应头到达后预算已解除
payload, err := io.ReadAll(stream)
if err != nil || string(payload) != "data: hello" {
t.Fatalf("预算解除后读流 = %q, %v; 期望完整 body", payload, err)
}
})
t.Run("超预算未返回响应头即取消", func(t *testing.T) {
fc := &fakeCaller{delay: time.Minute}
start := time.Now()
_, err := callWithHeaderBudget(context.Background(), fc, req, 30*time.Millisecond)
if !errors.Is(err, context.Canceled) {
t.Fatalf("期望 context.Canceled, 得到 %v", err)
}
if elapsed := time.Since(start); elapsed > 5*time.Second {
t.Fatalf("取消耗时 %v, 未受预算约束", elapsed)
}
})
t.Run("关闭流时取消派生 ctx", func(t *testing.T) {
fc := &fakeCaller{delay: 0, body: "x"}
stream, err := callWithHeaderBudget(context.Background(), fc, req, time.Minute)
if err != nil {
t.Fatalf("callWithHeaderBudget = %v", err)
}
stream.Close()
if fc.gotCtx.Err() == nil {
t.Fatal("Close 后派生 ctx 应已取消")
}
})
t.Run("建立失败时同样取消派生 ctx", func(t *testing.T) {
fc := &fakeCaller{failErr: errors.New("boom")}
if _, err := callWithHeaderBudget(context.Background(), fc, req, time.Minute); err == nil {
t.Fatal("期望建立失败")
}
if fc.gotCtx.Err() == nil {
t.Fatal("失败路径派生 ctx 应已取消")
}
})
}
// TestReadCompatBodyRejectsOversize 锁定上游响应上限:恰好达限通过,超限报错,
// 不允许静默截断配 200 让下游当完整成功。
func TestReadCompatBodyRejectsOversize(t *testing.T) {
cases := []struct {
name string
size int64
wantErr bool
}{
{"under limit", 15, false},
{"exactly limit", 16, false},
{"over limit", 17, true},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
payload, err := readCompatBody(bytes.NewReader(make([]byte, tc.size)), 16, "test")
if tc.wantErr {
if err == nil || !strings.Contains(err.Error(), "exceeds") {
t.Fatalf("err = %v, want 超限错误", err)
}
return
}
if err != nil || int64(len(payload)) != tc.size {
t.Fatalf("payload = %d bytes, %v; want %d", len(payload), err, tc.size)
}
})
}
}