mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-21 14:19:18 +08:00
Merge pull request #4072 from jianjianai/codex/fix-upstream-review
修复 tool_search 参数解析与 compact writer 生命周期
This commit is contained in:
@@ -1,9 +1,15 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) {
|
||||
@@ -57,3 +63,28 @@ func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) {
|
||||
assert.Nil(t, p)
|
||||
})
|
||||
}
|
||||
|
||||
func TestOpsCaptureWriter_CompactKeepaliveRestoresOriginalWriter(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
outerStatus := -1
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Next()
|
||||
outerStatus = c.Writer.Status()
|
||||
})
|
||||
router.Use(OpsErrorLoggerMiddleware(nil))
|
||||
router.GET("/compact", func(c *gin.Context) {
|
||||
service.MarkOpenAICompactClientStream(c)
|
||||
stop := service.StartOpenAICompactSSEKeepalive(c, time.Hour)
|
||||
defer stop()
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodGet, "/compact", nil)
|
||||
require.NotPanics(t, func() {
|
||||
router.ServeHTTP(recorder, request)
|
||||
})
|
||||
require.Equal(t, http.StatusOK, outerStatus)
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
}
|
||||
|
||||
@@ -131,3 +131,56 @@ func TestWire_UnknownEventFallsBackToDefault(t *testing.T) {
|
||||
})
|
||||
require.Contains(t, m, "response")
|
||||
}
|
||||
|
||||
func TestResponsesOutputUnmarshal_ToolSearchObjectArguments(t *testing.T) {
|
||||
var item ResponsesOutput
|
||||
require.NoError(t, json.Unmarshal([]byte(`{
|
||||
"type":"tool_search_call",
|
||||
"id":"item_1",
|
||||
"call_id":"call_1",
|
||||
"execution":"client",
|
||||
"arguments":{"query":"gmail","limit":2}
|
||||
}`), &item))
|
||||
require.Equal(t, "tool_search_call", item.Type)
|
||||
require.Equal(t, `{"query":"gmail","limit":2}`, item.Arguments)
|
||||
|
||||
wire, err := json.Marshal(item)
|
||||
require.NoError(t, err)
|
||||
var decoded map[string]any
|
||||
require.NoError(t, json.Unmarshal(wire, &decoded))
|
||||
args, ok := decoded["arguments"].(map[string]any)
|
||||
require.True(t, ok, "tool_search_call arguments must remain an object")
|
||||
require.Equal(t, "gmail", args["query"])
|
||||
}
|
||||
|
||||
func TestResponsesResponseUnmarshal_ToolSearchObjectArguments(t *testing.T) {
|
||||
var response ResponsesResponse
|
||||
require.NoError(t, json.Unmarshal([]byte(`{
|
||||
"id":"response_1",
|
||||
"object":"response",
|
||||
"status":"completed",
|
||||
"output":[{
|
||||
"type":"tool_search_call",
|
||||
"id":"item_1",
|
||||
"call_id":"call_1",
|
||||
"arguments":{"query":"gmail"}
|
||||
}]
|
||||
}`), &response))
|
||||
require.Len(t, response.Output, 1)
|
||||
require.Equal(t, `{"query":"gmail"}`, response.Output[0].Arguments)
|
||||
}
|
||||
|
||||
func TestResponsesStreamEventUnmarshal_ToolSearchObjectArguments(t *testing.T) {
|
||||
var event ResponsesStreamEvent
|
||||
require.NoError(t, json.Unmarshal([]byte(`{
|
||||
"type":"response.output_item.done",
|
||||
"item":{
|
||||
"type":"tool_search_call",
|
||||
"id":"item_1",
|
||||
"call_id":"call_1",
|
||||
"arguments":{"query":"gmail"}
|
||||
}
|
||||
}`), &event))
|
||||
require.NotNil(t, event.Item)
|
||||
require.Equal(t, `{"query":"gmail"}`, event.Item.Arguments)
|
||||
}
|
||||
|
||||
@@ -353,6 +353,56 @@ func (o ResponsesOutput) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(m)
|
||||
}
|
||||
|
||||
// UnmarshalJSON accepts both the Responses function-call string form and the
|
||||
// tool_search_call object form for arguments. The bridge stores arguments as a
|
||||
// string internally, so object arguments are retained as their raw JSON.
|
||||
func (o *ResponsesOutput) UnmarshalJSON(data []byte) error {
|
||||
type responsesOutputAlias ResponsesOutput
|
||||
|
||||
var kind struct {
|
||||
Type string `json:"type"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &kind); err != nil {
|
||||
return err
|
||||
}
|
||||
if kind.Type != "tool_search_call" {
|
||||
var decoded responsesOutputAlias
|
||||
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||
return err
|
||||
}
|
||||
*o = ResponsesOutput(decoded)
|
||||
return nil
|
||||
}
|
||||
|
||||
var fields map[string]json.RawMessage
|
||||
if err := json.Unmarshal(data, &fields); err != nil {
|
||||
return err
|
||||
}
|
||||
arguments, hasArguments := fields["arguments"]
|
||||
delete(fields, "arguments")
|
||||
normalized, err := json.Marshal(fields)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var decoded responsesOutputAlias
|
||||
if err := json.Unmarshal(normalized, &decoded); err != nil {
|
||||
return err
|
||||
}
|
||||
*o = ResponsesOutput(decoded)
|
||||
if !hasArguments || string(arguments) == "null" {
|
||||
return nil
|
||||
}
|
||||
|
||||
var argumentString string
|
||||
if err := json.Unmarshal(arguments, &argumentString); err == nil {
|
||||
o.Arguments = argumentString
|
||||
} else {
|
||||
o.Arguments = string(arguments)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// WebSearchAction describes the search action in a web_search_call output item.
|
||||
type WebSearchAction struct {
|
||||
Type string `json:"type,omitempty"` // "search"
|
||||
|
||||
@@ -43,12 +43,14 @@ func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func
|
||||
if c == nil || c.Writer == nil || interval <= 0 || !openAICompactClientWantsStream(c) {
|
||||
return func() {}
|
||||
}
|
||||
originalWriter := c.Writer
|
||||
k := &openAICompactSSEKeepalive{
|
||||
writer: c.Writer,
|
||||
writer: originalWriter,
|
||||
stop: make(chan struct{}),
|
||||
}
|
||||
c.Set(openAICompactSSEKeepaliveKey, k)
|
||||
c.Writer = &openAICompactKeepaliveWriter{ResponseWriter: c.Writer, k: k}
|
||||
wrappedWriter := &openAICompactKeepaliveWriter{ResponseWriter: originalWriter, k: k}
|
||||
c.Writer = wrappedWriter
|
||||
|
||||
var reqDone <-chan struct{}
|
||||
if c.Request != nil {
|
||||
@@ -71,7 +73,14 @@ func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func
|
||||
timer.Reset(interval)
|
||||
}
|
||||
}()
|
||||
return k.Stop
|
||||
return func() {
|
||||
k.Stop()
|
||||
// Do not leave a pooled middleware writer reachable through the compact
|
||||
// wrapper after the request finishes.
|
||||
if current, ok := c.Writer.(*openAICompactKeepaliveWriter); ok && current == wrappedWriter {
|
||||
c.Writer = originalWriter
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// beat 在锁内提交(首次)响应头并写出一条 SSE 注释行;返回 false 表示心跳已
|
||||
|
||||
Reference in New Issue
Block a user