mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge branch 'main' into feat/grok-video-edit-extension
This commit is contained in:
@@ -1 +1 @@
|
||||
0.1.151
|
||||
0.1.152
|
||||
|
||||
@@ -1454,13 +1454,14 @@ func (h *GatewayHandler) usageUnrestricted(c *gin.Context, ctx context.Context,
|
||||
remaining := h.calculateSubscriptionRemaining(apiKey.Group, subscription)
|
||||
resp["remaining"] = remaining
|
||||
resp["subscription"] = gin.H{
|
||||
"daily_usage_usd": subscription.DailyUsageUSD,
|
||||
"weekly_usage_usd": subscription.WeeklyUsageUSD,
|
||||
"monthly_usage_usd": subscription.MonthlyUsageUSD,
|
||||
"daily_limit_usd": apiKey.Group.DailyLimitUSD,
|
||||
"weekly_limit_usd": apiKey.Group.WeeklyLimitUSD,
|
||||
"monthly_limit_usd": apiKey.Group.MonthlyLimitUSD,
|
||||
"expires_at": subscription.ExpiresAt,
|
||||
"daily_usage_usd": subscription.DailyUsageUSD,
|
||||
"weekly_usage_usd": subscription.WeeklyUsageUSD,
|
||||
"monthly_usage_usd": subscription.MonthlyUsageUSD,
|
||||
"daily_limit_usd": apiKey.Group.DailyLimitUSD,
|
||||
"weekly_limit_usd": apiKey.Group.WeeklyLimitUSD,
|
||||
"monthly_limit_usd": apiKey.Group.MonthlyLimitUSD,
|
||||
"weekly_window_start": subscription.WeeklyWindowStart,
|
||||
"expires_at": subscription.ExpiresAt,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestUsageUnrestrictedIncludesWeeklyWindowStart(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/v1/usage", nil)
|
||||
|
||||
weeklyWindowStart := time.Date(2026, time.July, 13, 0, 30, 0, 0, time.FixedZone("UTC+8", 8*60*60))
|
||||
c.Set(string(middleware.ContextKeySubscription), &service.UserSubscription{
|
||||
WeeklyWindowStart: &weeklyWindowStart,
|
||||
})
|
||||
|
||||
handler := &GatewayHandler{}
|
||||
handler.usageUnrestricted(
|
||||
c,
|
||||
context.Background(),
|
||||
&service.APIKey{Group: &service.Group{
|
||||
Name: "Weekly plan",
|
||||
SubscriptionType: service.SubscriptionTypeSubscription,
|
||||
}},
|
||||
middleware.AuthSubject{},
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
var response struct {
|
||||
Subscription struct {
|
||||
WeeklyWindowStart *time.Time `json:"weekly_window_start"`
|
||||
} `json:"subscription"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response))
|
||||
require.NotNil(t, response.Subscription.WeeklyWindowStart)
|
||||
require.True(t, weeklyWindowStart.Equal(*response.Subscription.WeeklyWindowStart))
|
||||
}
|
||||
@@ -718,7 +718,7 @@ func TestStreamingToolCallDoneWithoutDeltaEmitsArguments(t *testing.T) {
|
||||
assert.Equal(t, "content_block_stop", events[1].Type)
|
||||
}
|
||||
|
||||
func TestStreamingReadToolDropsEmptyPages(t *testing.T) {
|
||||
func TestStreamingReadToolStreamsDeltas(t *testing.T) {
|
||||
state := NewResponsesEventToAnthropicState()
|
||||
|
||||
ResponsesEventToAnthropicEvents(&ResponsesStreamEvent{
|
||||
@@ -739,18 +739,17 @@ func TestStreamingReadToolDropsEmptyPages(t *testing.T) {
|
||||
OutputIndex: 0,
|
||||
Delta: `{"file_path":"/tmp/demo.py","limit":2000,"offset":0,"pages":""}`,
|
||||
}, state)
|
||||
assert.Len(t, events, 0)
|
||||
require.Len(t, events, 1, "Read tool deltas must be streamed like any other tool")
|
||||
assert.Equal(t, "content_block_delta", events[0].Type)
|
||||
assert.Equal(t, "input_json_delta", events[0].Delta.Type)
|
||||
|
||||
events = ResponsesEventToAnthropicEvents(&ResponsesStreamEvent{
|
||||
Type: "response.function_call_arguments.done",
|
||||
OutputIndex: 0,
|
||||
Arguments: `{"file_path":"/tmp/demo.py","limit":2000,"offset":0,"pages":""}`,
|
||||
}, state)
|
||||
require.Len(t, events, 2)
|
||||
assert.Equal(t, "content_block_delta", events[0].Type)
|
||||
assert.Equal(t, "input_json_delta", events[0].Delta.Type)
|
||||
assert.JSONEq(t, `{"file_path":"/tmp/demo.py","limit":2000,"offset":0}`, events[0].Delta.PartialJSON)
|
||||
assert.Equal(t, "content_block_stop", events[1].Type)
|
||||
require.Len(t, events, 1, "after streaming deltas, .done should just close the block")
|
||||
assert.Equal(t, "content_block_stop", events[0].Type)
|
||||
}
|
||||
|
||||
func TestStreamingReasoning(t *testing.T) {
|
||||
|
||||
@@ -164,6 +164,8 @@ type AnthropicEventToResponsesState struct {
|
||||
OutputTokens int
|
||||
CacheReadInputTokens int
|
||||
CacheCreationInputTokens int
|
||||
|
||||
StopReason string
|
||||
}
|
||||
|
||||
// NewAnthropicEventToResponsesState returns an initialised stream state.
|
||||
@@ -405,7 +407,6 @@ func anthToResHandleContentBlockStop(evt *AnthropicStreamEvent, state *Anthropic
|
||||
}
|
||||
|
||||
func anthToResHandleMessageDelta(evt *AnthropicStreamEvent, state *AnthropicEventToResponsesState) []ResponsesStreamEvent {
|
||||
// Update usage
|
||||
if evt.Usage != nil {
|
||||
state.OutputTokens = evt.Usage.OutputTokens
|
||||
if evt.Usage.InputTokens > 0 {
|
||||
@@ -418,6 +419,9 @@ func anthToResHandleMessageDelta(evt *AnthropicStreamEvent, state *AnthropicEven
|
||||
state.CacheCreationInputTokens = evt.Usage.CacheCreationInputTokens
|
||||
}
|
||||
}
|
||||
if evt.Delta != nil && evt.Delta.StopReason != "" {
|
||||
state.StopReason = evt.Delta.StopReason
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -428,15 +432,15 @@ func anthToResHandleMessageStop(state *AnthropicEventToResponsesState) []Respons
|
||||
}
|
||||
|
||||
var events []ResponsesStreamEvent
|
||||
|
||||
// Close any open item
|
||||
events = append(events, closeCurrentResponsesItem(state)...)
|
||||
|
||||
// Determine status
|
||||
status := "completed"
|
||||
var incompleteDetails *ResponsesIncompleteDetails
|
||||
if state.StopReason == "max_tokens" {
|
||||
status = "incomplete"
|
||||
incompleteDetails = &ResponsesIncompleteDetails{Reason: "max_output_tokens"}
|
||||
}
|
||||
|
||||
// Emit response.completed
|
||||
events = append(events, makeResponsesCompletedEvent(state, status, incompleteDetails))
|
||||
state.CompletedSent = true
|
||||
return events
|
||||
@@ -509,15 +513,20 @@ func makeResponsesCompletedEvent(
|
||||
}
|
||||
}
|
||||
|
||||
eventType := "response.completed"
|
||||
if status == "incomplete" {
|
||||
eventType = "response.incomplete"
|
||||
}
|
||||
|
||||
return ResponsesStreamEvent{
|
||||
Type: "response.completed",
|
||||
Type: eventType,
|
||||
SequenceNumber: seq,
|
||||
Response: &ResponsesResponse{
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
Model: state.Model,
|
||||
Status: status,
|
||||
Output: []ResponsesOutput{}, // Simplified; full output tracking would add complexity
|
||||
Output: []ResponsesOutput{},
|
||||
Usage: usage,
|
||||
IncompleteDetails: incompleteDetails,
|
||||
},
|
||||
|
||||
@@ -35,8 +35,12 @@ func ResponsesToChatCompletionsRequest(req *ResponsesRequest) (*ChatCompletionsR
|
||||
if req.Reasoning != nil {
|
||||
out.ReasoningEffort = req.Reasoning.Effort
|
||||
}
|
||||
if len(req.Tools) > 0 {
|
||||
tools, err := responsesToolsToChatTools(req.Tools)
|
||||
effectiveTools, err := EffectiveResponsesTools(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(effectiveTools) > 0 {
|
||||
tools, err := responsesToolsToChatTools(effectiveTools)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -63,6 +67,44 @@ func ResponsesToChatCompletionsRequest(req *ResponsesRequest) (*ChatCompletionsR
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// EffectiveResponsesTools returns every client-executable tool declared by a
|
||||
// Responses request. Newer Codex clients place their runtime tools in an
|
||||
// input item shaped as {"type":"additional_tools","tools":[...]} instead of
|
||||
// the top-level tools field. Chat-only upstreams must receive both forms.
|
||||
func EffectiveResponsesTools(req *ResponsesRequest) ([]ResponsesTool, error) {
|
||||
if req == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
tools := append([]ResponsesTool(nil), req.Tools...)
|
||||
inputRaw := bytesTrimSpace(req.Input)
|
||||
if len(inputRaw) == 0 || string(inputRaw) == "null" || inputRaw[0] != '[' {
|
||||
return tools, nil
|
||||
}
|
||||
|
||||
var items []json.RawMessage
|
||||
if err := json.Unmarshal(inputRaw, &items); err != nil {
|
||||
return nil, fmt.Errorf("parse responses input for additional tools: %w", err)
|
||||
}
|
||||
for _, raw := range items {
|
||||
raw = bytesTrimSpace(raw)
|
||||
if len(raw) == 0 || raw[0] != '{' {
|
||||
continue
|
||||
}
|
||||
var item struct {
|
||||
Type string `json:"type"`
|
||||
Tools []ResponsesTool `json:"tools"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &item); err != nil {
|
||||
return nil, fmt.Errorf("parse responses additional tools item: %w", err)
|
||||
}
|
||||
if item.Type == "additional_tools" {
|
||||
tools = append(tools, item.Tools...)
|
||||
}
|
||||
}
|
||||
return tools, nil
|
||||
}
|
||||
|
||||
// CustomToolNames 收集 Responses 请求中 custom/freeform 工具的名字。chat 桥回程时
|
||||
// 需要据此把模型对这些工具的调用还原为 custom_tool_call 项(codex 只按该类型路由)。
|
||||
func CustomToolNames(tools []ResponsesTool) map[string]bool {
|
||||
|
||||
@@ -34,6 +34,51 @@ func TestResponsesToChatCompletionsRequest_CustomToolBecomesFunctionTool(t *test
|
||||
assert.Equal(t, "wait", out.Tools[1].Function.Name)
|
||||
}
|
||||
|
||||
func TestResponsesToChatCompletionsRequest_AdditionalToolsItem(t *testing.T) {
|
||||
req := &ResponsesRequest{
|
||||
Model: "gpt-test",
|
||||
Input: json.RawMessage(`[
|
||||
{"type":"additional_tools","role":"developer","tools":[
|
||||
{"type":"custom","name":"exec","description":"Run PowerShell","format":{"type":"text"}},
|
||||
{"type":"function","name":"wait","parameters":{"type":"object","properties":{}}},
|
||||
{"type":"namespace","name":"collaboration","tools":[
|
||||
{"type":"function","name":"send_message","parameters":{"type":"object","properties":{}}}
|
||||
]}
|
||||
]},
|
||||
{"type":"message","role":"user","content":[{"type":"input_text","text":"run Get-Location"}]}
|
||||
]`),
|
||||
ToolChoice: json.RawMessage(`"auto"`),
|
||||
}
|
||||
|
||||
effective, err := EffectiveResponsesTools(req)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, effective, 3)
|
||||
assert.True(t, CustomToolNames(effective)["exec"])
|
||||
assert.Equal(t, NamespacedToolName{Namespace: "collaboration", Name: "send_message"}, NamespaceToolNames(effective)["collaboration__send_message"])
|
||||
|
||||
out, err := ResponsesToChatCompletionsRequest(req)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, out.Tools, 3)
|
||||
assert.Equal(t, "exec", out.Tools[0].Function.Name)
|
||||
assert.Equal(t, "wait", out.Tools[1].Function.Name)
|
||||
assert.Equal(t, "collaboration__send_message", out.Tools[2].Function.Name)
|
||||
assert.JSONEq(t, `"auto"`, string(out.ToolChoice))
|
||||
|
||||
require.Len(t, out.Messages, 1, "additional_tools must not become a chat message")
|
||||
assert.Equal(t, "user", out.Messages[0].Role)
|
||||
}
|
||||
|
||||
func TestEffectiveResponsesTools_SkipsStringInputItems(t *testing.T) {
|
||||
req := &ResponsesRequest{
|
||||
Input: json.RawMessage(`["plain input",{"type":"additional_tools","tools":[{"type":"custom","name":"exec"}]}]`),
|
||||
}
|
||||
|
||||
tools, err := EffectiveResponsesTools(req)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, tools, 1)
|
||||
assert.Equal(t, "exec", tools[0].Name)
|
||||
}
|
||||
|
||||
func TestResponsesToChatCompletionsRequest_DropsToolChoiceWhenNoConvertibleTools(t *testing.T) {
|
||||
req := &ResponsesRequest{
|
||||
Model: "glm-5.2",
|
||||
|
||||
@@ -413,10 +413,6 @@ func resToAnthHandleFuncArgsDelta(evt *ResponsesStreamEvent, state *ResponsesEve
|
||||
return nil
|
||||
}
|
||||
|
||||
if state.CurrentBlockType == "tool_use" && state.CurrentToolName == "Read" {
|
||||
state.CurrentToolArgs += evt.Delta
|
||||
return nil
|
||||
}
|
||||
if state.CurrentBlockType == "tool_use" {
|
||||
state.CurrentToolHadDelta = true
|
||||
}
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
package apicompat
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestResToAnthFuncArgsDelta_ReadToolStreamsDeltas(t *testing.T) {
|
||||
state := NewResponsesEventToAnthropicState()
|
||||
state.MessageStartSent = true
|
||||
state.CurrentBlockType = "tool_use"
|
||||
state.CurrentToolName = "Read"
|
||||
state.OutputIndexToBlockIdx = map[int]int{0: 0}
|
||||
|
||||
evt := &ResponsesStreamEvent{
|
||||
Type: "response.function_call_arguments.delta",
|
||||
OutputIndex: 0,
|
||||
Delta: `{"file_path":"/tmp/test.go"}`,
|
||||
}
|
||||
|
||||
events := ResponsesEventToAnthropicEvents(evt, state)
|
||||
|
||||
require.Len(t, events, 1, "Read tool delta must produce content_block_delta")
|
||||
assert.Equal(t, "content_block_delta", events[0].Type)
|
||||
assert.Equal(t, "input_json_delta", events[0].Delta.Type)
|
||||
assert.Equal(t, `{"file_path":"/tmp/test.go"}`, events[0].Delta.PartialJSON)
|
||||
assert.True(t, state.CurrentToolHadDelta, "Read deltas should set CurrentToolHadDelta")
|
||||
}
|
||||
|
||||
func TestResToAnthFuncArgsDelta_ReadToolWithoutDone(t *testing.T) {
|
||||
state := NewResponsesEventToAnthropicState()
|
||||
state.MessageStartSent = true
|
||||
state.ContentBlockIndex = 0
|
||||
state.ContentBlockOpen = true
|
||||
state.CurrentBlockType = "tool_use"
|
||||
state.CurrentToolName = "Read"
|
||||
state.OutputIndexToBlockIdx = map[int]int{0: 0}
|
||||
|
||||
delta := &ResponsesStreamEvent{
|
||||
Type: "response.function_call_arguments.delta",
|
||||
OutputIndex: 0,
|
||||
Delta: `{"file_path":"/tmp/test.go"}`,
|
||||
}
|
||||
events := ResponsesEventToAnthropicEvents(delta, state)
|
||||
require.Len(t, events, 1, "delta should be streamed")
|
||||
|
||||
completed := &ResponsesStreamEvent{
|
||||
Type: "response.completed",
|
||||
Response: &ResponsesResponse{
|
||||
Status: "completed",
|
||||
},
|
||||
}
|
||||
events = ResponsesEventToAnthropicEvents(completed, state)
|
||||
|
||||
hasStop := false
|
||||
for _, e := range events {
|
||||
if e.Type == "content_block_stop" {
|
||||
hasStop = true
|
||||
}
|
||||
}
|
||||
assert.True(t, hasStop, "block should be closed even without .done event")
|
||||
}
|
||||
|
||||
func TestResToAnthFuncArgsDelta_NonReadToolUnchanged(t *testing.T) {
|
||||
state := NewResponsesEventToAnthropicState()
|
||||
state.MessageStartSent = true
|
||||
state.CurrentBlockType = "tool_use"
|
||||
state.CurrentToolName = "Write"
|
||||
state.OutputIndexToBlockIdx = map[int]int{0: 0}
|
||||
|
||||
evt := &ResponsesStreamEvent{
|
||||
Type: "response.function_call_arguments.delta",
|
||||
OutputIndex: 0,
|
||||
Delta: `{"file_path":"/tmp/out.txt","content":"hello"}`,
|
||||
}
|
||||
|
||||
events := ResponsesEventToAnthropicEvents(evt, state)
|
||||
|
||||
require.Len(t, events, 1)
|
||||
assert.Equal(t, "content_block_delta", events[0].Type)
|
||||
assert.True(t, state.CurrentToolHadDelta)
|
||||
}
|
||||
@@ -89,8 +89,13 @@ func ResponsesToChatCompletions(resp *ResponsesResponse, model string) *ChatComp
|
||||
func responsesStatusToChatFinishReason(status string, details *ResponsesIncompleteDetails, toolCalls []ChatToolCall) string {
|
||||
switch status {
|
||||
case "incomplete":
|
||||
if details != nil && details.Reason == "max_output_tokens" {
|
||||
return "length"
|
||||
if details != nil {
|
||||
switch details.Reason {
|
||||
case "max_output_tokens":
|
||||
return "length"
|
||||
case "content_filter":
|
||||
return "content_filter"
|
||||
}
|
||||
}
|
||||
return "stop"
|
||||
case "completed":
|
||||
@@ -299,8 +304,13 @@ func resToChatHandleCompleted(evt *ResponsesStreamEvent, state *ResponsesEventTo
|
||||
|
||||
switch evt.Response.Status {
|
||||
case "incomplete":
|
||||
if evt.Response.IncompleteDetails != nil && evt.Response.IncompleteDetails.Reason == "max_output_tokens" {
|
||||
finishReason = "length"
|
||||
if evt.Response.IncompleteDetails != nil {
|
||||
switch evt.Response.IncompleteDetails.Reason {
|
||||
case "max_output_tokens":
|
||||
finishReason = "length"
|
||||
case "content_filter":
|
||||
finishReason = "content_filter"
|
||||
}
|
||||
}
|
||||
case "completed":
|
||||
if state.SawToolCall {
|
||||
|
||||
@@ -0,0 +1,122 @@
|
||||
package apicompat
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAnthropicStreamingMaxTokens_MapsToIncomplete(t *testing.T) {
|
||||
state := NewAnthropicEventToResponsesState()
|
||||
|
||||
AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
|
||||
Type: "message_start",
|
||||
Message: &AnthropicResponse{ID: "msg_test", Model: "claude-opus-4-6", Role: "assistant"},
|
||||
}, state)
|
||||
|
||||
AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
|
||||
Type: "message_delta",
|
||||
Delta: &AnthropicDelta{
|
||||
StopReason: "max_tokens",
|
||||
},
|
||||
Usage: &AnthropicUsage{OutputTokens: 4096},
|
||||
}, state)
|
||||
|
||||
require.Equal(t, "max_tokens", state.StopReason)
|
||||
|
||||
events := AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
|
||||
Type: "message_stop",
|
||||
}, state)
|
||||
|
||||
var completed *ResponsesStreamEvent
|
||||
for i := range events {
|
||||
if events[i].Type == "response.completed" || events[i].Type == "response.incomplete" {
|
||||
completed = &events[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotNil(t, completed, "should have terminal event")
|
||||
assert.Equal(t, "response.incomplete", completed.Type)
|
||||
require.NotNil(t, completed.Response)
|
||||
assert.Equal(t, "incomplete", completed.Response.Status)
|
||||
require.NotNil(t, completed.Response.IncompleteDetails)
|
||||
assert.Equal(t, "max_output_tokens", completed.Response.IncompleteDetails.Reason)
|
||||
}
|
||||
|
||||
func TestAnthropicStreamingEndTurn_MapsToCompleted(t *testing.T) {
|
||||
state := NewAnthropicEventToResponsesState()
|
||||
|
||||
AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
|
||||
Type: "message_start",
|
||||
Message: &AnthropicResponse{ID: "msg_test", Model: "claude-opus-4-6", Role: "assistant"},
|
||||
}, state)
|
||||
|
||||
AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
|
||||
Type: "message_delta",
|
||||
Delta: &AnthropicDelta{StopReason: "end_turn"},
|
||||
Usage: &AnthropicUsage{OutputTokens: 100},
|
||||
}, state)
|
||||
|
||||
events := AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
|
||||
Type: "message_stop",
|
||||
}, state)
|
||||
|
||||
var completed *ResponsesStreamEvent
|
||||
for i := range events {
|
||||
if events[i].Type == "response.completed" {
|
||||
completed = &events[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotNil(t, completed)
|
||||
assert.Equal(t, "completed", completed.Response.Status)
|
||||
assert.Nil(t, completed.Response.IncompleteDetails)
|
||||
}
|
||||
|
||||
func TestResponsesToChatCompletions_ContentFilter(t *testing.T) {
|
||||
resp := &ResponsesResponse{
|
||||
ID: "resp_cf",
|
||||
Status: "incomplete",
|
||||
IncompleteDetails: &ResponsesIncompleteDetails{
|
||||
Reason: "content_filter",
|
||||
},
|
||||
Output: []ResponsesOutput{{
|
||||
Type: "message",
|
||||
Content: []ResponsesContentPart{{Type: "output_text", Text: "partial"}},
|
||||
}},
|
||||
Usage: &ResponsesUsage{InputTokens: 10, OutputTokens: 5},
|
||||
}
|
||||
|
||||
cc := ResponsesToChatCompletions(resp, "gpt-5.5")
|
||||
require.Len(t, cc.Choices, 1)
|
||||
assert.Equal(t, "content_filter", cc.Choices[0].FinishReason)
|
||||
}
|
||||
|
||||
func TestResponsesToChatCompletionsStreaming_ContentFilter(t *testing.T) {
|
||||
state := NewResponsesEventToChatState()
|
||||
state.ID = "resp_cf"
|
||||
state.Model = "gpt-5.5"
|
||||
state.SentRole = true
|
||||
|
||||
events := ResponsesEventToChatChunks(&ResponsesStreamEvent{
|
||||
Type: "response.completed",
|
||||
Response: &ResponsesResponse{
|
||||
ID: "resp_cf",
|
||||
Status: "incomplete",
|
||||
IncompleteDetails: &ResponsesIncompleteDetails{
|
||||
Reason: "content_filter",
|
||||
},
|
||||
},
|
||||
}, state)
|
||||
|
||||
hasContentFilter := false
|
||||
for _, chunk := range events {
|
||||
for _, choice := range chunk.Choices {
|
||||
if choice.FinishReason != nil && *choice.FinishReason == "content_filter" {
|
||||
hasContentFilter = true
|
||||
}
|
||||
}
|
||||
}
|
||||
assert.True(t, hasContentFilter, "streaming content_filter should map to finish_reason content_filter")
|
||||
}
|
||||
@@ -38,7 +38,7 @@ func (p PaginationParams) Offset() int {
|
||||
if p.Page < 1 {
|
||||
p.Page = 1
|
||||
}
|
||||
return (p.Page - 1) * p.PageSize
|
||||
return (p.Page - 1) * p.Limit()
|
||||
}
|
||||
|
||||
// Limit 获取限制数
|
||||
|
||||
@@ -69,3 +69,30 @@ func TestPaginationParamsLimit(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPaginationParamsOffsetUsesNormalizedLimit(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
page int
|
||||
pageSize int
|
||||
want int
|
||||
}{
|
||||
{name: "invalid page uses first page", page: 0, pageSize: 50, want: 0},
|
||||
{name: "zero page size uses default", page: 2, pageSize: 0, want: 20},
|
||||
{name: "negative page size uses default", page: 2, pageSize: -1, want: 20},
|
||||
{name: "normal values", page: 3, pageSize: 50, want: 100},
|
||||
{name: "page size beyond max is clamped", page: 2, pageSize: 1500, want: 1000},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
params := PaginationParams{Page: tt.page, PageSize: tt.pageSize}
|
||||
if got := params.Offset(); got != tt.want {
|
||||
t.Fatalf("Offset() for Page=%d, PageSize=%d = %d, want %d", tt.page, tt.pageSize, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -525,17 +525,19 @@ func (r *apiKeyRepository) latestUsageLogIPs(ctx context.Context, apiKeyIDs []in
|
||||
|
||||
func latestUsageLogIPsQuery(apiKeyIDs []int64, dialectName string) (string, []any) {
|
||||
if dialectName == dialect.Postgres {
|
||||
// Keep each key lookup bounded to one ordered index probe instead of ranking its full history.
|
||||
return `
|
||||
SELECT api_key_id, ip_address
|
||||
FROM (
|
||||
SELECT api_key_id, ip_address,
|
||||
ROW_NUMBER() OVER (PARTITION BY api_key_id ORDER BY created_at DESC, id DESC) AS rn
|
||||
FROM usage_logs
|
||||
WHERE api_key_id = ANY($1::bigint[])
|
||||
AND ip_address IS NOT NULL
|
||||
AND ip_address <> ''
|
||||
) ranked
|
||||
WHERE rn = 1`, []any{pq.Array(apiKeyIDs)}
|
||||
SELECT requested.api_key_id, latest.ip_address
|
||||
FROM unnest($1::bigint[]) AS requested(api_key_id)
|
||||
CROSS JOIN LATERAL (
|
||||
SELECT ul.ip_address
|
||||
FROM usage_logs AS ul
|
||||
WHERE ul.api_key_id = requested.api_key_id
|
||||
AND ul.ip_address IS NOT NULL
|
||||
AND ul.ip_address <> ''
|
||||
ORDER BY ul.created_at DESC, ul.id DESC
|
||||
LIMIT 1
|
||||
) AS latest`, []any{pq.Array(apiKeyIDs)}
|
||||
}
|
||||
|
||||
placeholders := make([]string, len(apiKeyIDs))
|
||||
|
||||
@@ -3,6 +3,7 @@ package repository
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -125,6 +126,20 @@ func TestAPIKeyRepositoryListByUserIDAttachesLastUsedIP(t *testing.T) {
|
||||
require.Nil(t, byID[noLogs.ID].LastUsedIP)
|
||||
}
|
||||
|
||||
func TestLatestUsageLogIPsQueryPostgresUsesPerKeyLateralLookup(t *testing.T) {
|
||||
query, args := latestUsageLogIPsQuery([]int64{11, 22}, dialect.Postgres)
|
||||
normalizedQuery := strings.Join(strings.Fields(query), " ")
|
||||
|
||||
require.Contains(t, normalizedQuery, "FROM unnest($1::bigint[]) AS requested(api_key_id)")
|
||||
require.Contains(t, normalizedQuery, "CROSS JOIN LATERAL")
|
||||
require.Contains(t, normalizedQuery, "WHERE ul.api_key_id = requested.api_key_id")
|
||||
require.Contains(t, normalizedQuery, "AND ul.ip_address IS NOT NULL")
|
||||
require.Contains(t, normalizedQuery, "AND ul.ip_address <> ''")
|
||||
require.Contains(t, normalizedQuery, "ORDER BY ul.created_at DESC, ul.id DESC LIMIT 1")
|
||||
require.NotContains(t, normalizedQuery, "ROW_NUMBER")
|
||||
require.Len(t, args, 1)
|
||||
}
|
||||
|
||||
func TestAPIKeyRepository_CreateWithLastUsedAt(t *testing.T) {
|
||||
repo, client := newAPIKeyRepoSQLite(t)
|
||||
ctx := context.Background()
|
||||
|
||||
@@ -55,6 +55,8 @@ const paymentOrdersOutTradeNoUniqueMigration = "120_enforce_payment_orders_out_t
|
||||
const paymentOrdersOutTradeNoUniqueIndex = "paymentorder_out_trade_no_unique"
|
||||
const schedulerOutboxPendingDedupKeyMigration = "153_scheduler_outbox_pending_dedup_key_index_notx.sql"
|
||||
const schedulerOutboxPendingDedupKeyIndex = "idx_scheduler_outbox_pending_dedup_key"
|
||||
const latestAPIKeyIPIndexMigration = "174_add_usage_logs_api_key_latest_ip_index_notx.sql"
|
||||
const latestAPIKeyIPIndex = "idx_usage_logs_api_key_latest_ip"
|
||||
|
||||
type migrationChecksumCompatibilityRule struct {
|
||||
fileChecksum string
|
||||
@@ -264,6 +266,8 @@ func prepareNonTransactionalMigration(ctx context.Context, db *sql.DB, name stri
|
||||
return preparePaymentOrdersOutTradeNoUniqueMigration(ctx, db)
|
||||
case schedulerOutboxPendingDedupKeyMigration:
|
||||
return dropInvalidIndexIfPresent(ctx, db, schedulerOutboxPendingDedupKeyIndex)
|
||||
case latestAPIKeyIPIndexMigration:
|
||||
return dropInvalidIndexIfPresent(ctx, db, latestAPIKeyIPIndex)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -116,6 +116,45 @@ CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_t_b ON t(b);
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
|
||||
func TestApplyMigrationsFS_NonTransactionalMigration_LatestAPIKeyIPIndexDropsInvalidIndexBeforeRetry(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = db.Close() }()
|
||||
|
||||
prepareMigrationsBootstrapExpectations(mock)
|
||||
mock.ExpectQuery("SELECT checksum FROM schema_migrations WHERE filename = \\$1").
|
||||
WithArgs(latestAPIKeyIPIndexMigration).
|
||||
WillReturnError(sql.ErrNoRows)
|
||||
mock.ExpectQuery("SELECT EXISTS \\(").
|
||||
WithArgs(latestAPIKeyIPIndex).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true))
|
||||
mock.ExpectExec("DROP INDEX CONCURRENTLY IF EXISTS idx_usage_logs_api_key_latest_ip").
|
||||
WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
mock.ExpectExec("CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_api_key_latest_ip").
|
||||
WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
mock.ExpectExec("INSERT INTO schema_migrations \\(filename, checksum\\) VALUES \\(\\$1, \\$2\\)").
|
||||
WithArgs(latestAPIKeyIPIndexMigration, sqlmock.AnyArg()).
|
||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
||||
mock.ExpectExec("SELECT pg_advisory_unlock\\(\\$1\\)").
|
||||
WithArgs(migrationsAdvisoryLockID).
|
||||
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
|
||||
fsys := fstest.MapFS{
|
||||
latestAPIKeyIPIndexMigration: &fstest.MapFile{
|
||||
Data: []byte(`
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_api_key_latest_ip
|
||||
ON usage_logs (api_key_id, created_at DESC, id DESC)
|
||||
INCLUDE (ip_address)
|
||||
WHERE ip_address IS NOT NULL AND ip_address <> '';
|
||||
`),
|
||||
},
|
||||
}
|
||||
|
||||
err = applyMigrationsFS(context.Background(), db, fsys)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
|
||||
func TestApplyMigrationsFS_PaymentOrdersOutTradeNoUniqueMigration_FailsFastOnDuplicatePrecheck(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
@@ -163,14 +164,15 @@ func (c *schedulerCache) SetSnapshot(ctx context.Context, bucket service.Schedul
|
||||
versionStr := strconv.FormatInt(version, 10)
|
||||
snapshotKey := schedulerSnapshotKey(bucket, versionStr)
|
||||
|
||||
if err := c.writeAccounts(ctx, accounts); err != nil {
|
||||
cacheableAccounts, err := c.writeAccounts(ctx, accounts)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if len(accounts) > 0 {
|
||||
if len(cacheableAccounts) > 0 {
|
||||
// 使用序号作为 score,保持数据库返回的排序语义。
|
||||
members := make([]redis.Z, 0, len(accounts))
|
||||
for idx, account := range accounts {
|
||||
members := make([]redis.Z, 0, len(cacheableAccounts))
|
||||
for idx, account := range cacheableAccounts {
|
||||
members = append(members, redis.Z{
|
||||
Score: float64(idx),
|
||||
Member: strconv.FormatInt(account.ID, 10),
|
||||
@@ -224,7 +226,14 @@ func (c *schedulerCache) SetAccount(ctx context.Context, account *service.Accoun
|
||||
if account == nil || account.ID <= 0 {
|
||||
return nil
|
||||
}
|
||||
return c.writeAccounts(ctx, []service.Account{*account})
|
||||
cacheableAccounts, err := c.writeAccounts(ctx, []service.Account{*account})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(cacheableAccounts) == 0 {
|
||||
return c.DeleteAccount(ctx, account.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *schedulerCache) DeleteAccount(ctx context.Context, accountID int64) error {
|
||||
@@ -262,13 +271,14 @@ func (c *schedulerCache) UpdateLastUsed(ctx context.Context, updates map[int64]t
|
||||
return err
|
||||
}
|
||||
account.LastUsedAt = ptrTime(updates[ids[i]])
|
||||
updated, err := json.Marshal(account)
|
||||
updated, metaPayload, err := marshalSchedulerCacheAccount(*account)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
metaPayload, err := json.Marshal(buildSchedulerMetadataAccount(*account))
|
||||
if err != nil {
|
||||
return err
|
||||
slog.Warn("scheduler cache removes account with unencodable payload",
|
||||
"account_id", ids[i],
|
||||
"error", err,
|
||||
)
|
||||
pipe.Del(ctx, keys[i], schedulerAccountMetaKey(strconv.FormatInt(ids[i], 10)))
|
||||
continue
|
||||
}
|
||||
pipe.Set(ctx, keys[i], updated, 0)
|
||||
pipe.Set(ctx, schedulerAccountMetaKey(strconv.FormatInt(ids[i], 10)), metaPayload, 0)
|
||||
@@ -359,12 +369,13 @@ func decodeCachedAccount(val any) (*service.Account, error) {
|
||||
return &account, nil
|
||||
}
|
||||
|
||||
func (c *schedulerCache) writeAccounts(ctx context.Context, accounts []service.Account) error {
|
||||
func (c *schedulerCache) writeAccounts(ctx context.Context, accounts []service.Account) ([]service.Account, error) {
|
||||
if len(accounts) == 0 {
|
||||
return nil
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
pipe := c.rdb.Pipeline()
|
||||
cacheableAccounts := make([]service.Account, 0, len(accounts))
|
||||
pending := 0
|
||||
flush := func() error {
|
||||
if pending == 0 {
|
||||
@@ -379,27 +390,43 @@ func (c *schedulerCache) writeAccounts(ctx context.Context, accounts []service.A
|
||||
}
|
||||
|
||||
for _, account := range accounts {
|
||||
fullPayload, err := json.Marshal(account)
|
||||
fullPayload, metaPayload, err := marshalSchedulerCacheAccount(account)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
metaPayload, err := json.Marshal(buildSchedulerMetadataAccount(account))
|
||||
if err != nil {
|
||||
return err
|
||||
slog.Warn("scheduler cache skips account with unencodable payload",
|
||||
"account_id", account.ID,
|
||||
"error", err,
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
id := strconv.FormatInt(account.ID, 10)
|
||||
pipe.Set(ctx, schedulerAccountKey(id), fullPayload, 0)
|
||||
pipe.Set(ctx, schedulerAccountMetaKey(id), metaPayload, 0)
|
||||
cacheableAccounts = append(cacheableAccounts, account)
|
||||
pending++
|
||||
if pending >= c.writeChunkSize {
|
||||
if err := flush(); err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return flush()
|
||||
if err := flush(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cacheableAccounts, nil
|
||||
}
|
||||
|
||||
func marshalSchedulerCacheAccount(account service.Account) ([]byte, []byte, error) {
|
||||
fullPayload, err := json.Marshal(account)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("marshal account: %w", err)
|
||||
}
|
||||
metaPayload, err := json.Marshal(buildSchedulerMetadataAccount(account))
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("marshal account metadata: %w", err)
|
||||
}
|
||||
return fullPayload, metaPayload, nil
|
||||
}
|
||||
|
||||
func (c *schedulerCache) mgetChunked(ctx context.Context, keys []string) ([]any, error) {
|
||||
|
||||
@@ -3,12 +3,78 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func newSchedulerCacheUnit(t *testing.T) *schedulerCache {
|
||||
t.Helper()
|
||||
mr := miniredis.RunT(t)
|
||||
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
|
||||
t.Cleanup(func() { _ = rdb.Close() })
|
||||
cache, ok := newSchedulerCacheWithChunkSizes(rdb, defaultSchedulerSnapshotMGetChunkSize, defaultSchedulerSnapshotWriteChunkSize).(*schedulerCache)
|
||||
require.True(t, ok)
|
||||
return cache
|
||||
}
|
||||
|
||||
func TestSchedulerCacheWriteAccountsSkipsUnencodableTimes(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
cache := newSchedulerCacheUnit(t)
|
||||
invalidTime := time.Date(10000, time.January, 1, 0, 0, 0, 0, time.UTC)
|
||||
|
||||
cacheable, err := cache.writeAccounts(ctx, []service.Account{
|
||||
{ID: 111, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey},
|
||||
{ID: 112, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey, ExpiresAt: &invalidTime},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, cacheable, 1)
|
||||
require.Equal(t, int64(111), cacheable[0].ID)
|
||||
|
||||
cached, err := cache.GetAccount(ctx, 111)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, cached)
|
||||
|
||||
invalid, err := cache.GetAccount(ctx, 112)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, invalid)
|
||||
}
|
||||
|
||||
func TestSchedulerCacheSetAccountClearsUnencodablePayload(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
cache := newSchedulerCacheUnit(t)
|
||||
|
||||
account := service.Account{ID: 113, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey}
|
||||
require.NoError(t, cache.SetAccount(ctx, &account))
|
||||
|
||||
invalidTime := time.Date(10000, time.January, 1, 0, 0, 0, 0, time.UTC)
|
||||
account.ExpiresAt = &invalidTime
|
||||
require.NoError(t, cache.SetAccount(ctx, &account))
|
||||
|
||||
cached, err := cache.GetAccount(ctx, account.ID)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, cached)
|
||||
}
|
||||
|
||||
func TestSchedulerCacheUpdateLastUsedClearsUnencodablePayload(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
cache := newSchedulerCacheUnit(t)
|
||||
account := service.Account{ID: 114, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey}
|
||||
require.NoError(t, cache.SetAccount(ctx, &account))
|
||||
|
||||
invalidTime := time.Date(10000, time.January, 1, 0, 0, 0, 0, time.UTC)
|
||||
require.NoError(t, cache.UpdateLastUsed(ctx, map[int64]time.Time{account.ID: invalidTime}))
|
||||
|
||||
cached, err := cache.GetAccount(ctx, account.ID)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, cached)
|
||||
}
|
||||
|
||||
func TestBuildSchedulerMetadataAccount_KeepsOpenAIWSFlags(t *testing.T) {
|
||||
account := service.Account{
|
||||
ID: 42,
|
||||
|
||||
@@ -43,9 +43,14 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
|
||||
// custom_tool_call 项,先记下名字集合;tool_search 工具同理,回程还原为
|
||||
// tool_search_call 项;namespace 子工具(如 MCP 工具)摊平转发,回程按映射还原
|
||||
// 为带 namespace 字段的 function_call 项。
|
||||
customTools := apicompat.CustomToolNames(responsesReq.Tools)
|
||||
toolSearch := apicompat.HasToolSearchTool(responsesReq.Tools)
|
||||
namespaceTools := apicompat.NamespaceToolNames(responsesReq.Tools)
|
||||
effectiveTools, err := apicompat.EffectiveResponsesTools(&responsesReq)
|
||||
if err != nil {
|
||||
writeOpenAIResponsesFallbackError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
||||
return nil, fmt.Errorf("resolve responses tools: %w", err)
|
||||
}
|
||||
customTools := apicompat.CustomToolNames(effectiveTools)
|
||||
toolSearch := apicompat.HasToolSearchTool(effectiveTools)
|
||||
namespaceTools := apicompat.NamespaceToolNames(effectiveTools)
|
||||
|
||||
chatReq, err := apicompat.ResponsesToChatCompletionsRequest(&responsesReq)
|
||||
if err != nil {
|
||||
|
||||
@@ -131,6 +131,8 @@ func (s *AccountTestService) buildUpstreamModelsRequest(ctx context.Context, acc
|
||||
switch {
|
||||
case account.Platform == PlatformAntigravity:
|
||||
return s.buildAntigravityAPIKeyModelsRequest(ctx, account)
|
||||
case account.IsGrok():
|
||||
return s.buildGrokUpstreamModelsRequest(ctx, account)
|
||||
case account.IsOpenAI():
|
||||
return s.buildOpenAIUpstreamModelsRequest(ctx, account)
|
||||
case account.IsGemini():
|
||||
@@ -144,6 +146,36 @@ func (s *AccountTestService) buildUpstreamModelsRequest(ctx context.Context, acc
|
||||
}
|
||||
}
|
||||
|
||||
func (s *AccountTestService) buildGrokUpstreamModelsRequest(ctx context.Context, account *Account) (*http.Request, error) {
|
||||
if account.Type != AccountTypeAPIKey {
|
||||
return nil, newUpstreamModelSyncUnsupportedError(
|
||||
fmt.Sprintf("Unsupported Grok account type for upstream model sync: %s", account.Type), nil,
|
||||
)
|
||||
}
|
||||
apiKey := strings.TrimSpace(account.GetCredential("api_key"))
|
||||
if apiKey == "" {
|
||||
return nil, newUpstreamModelSyncConfigError("No Grok API key is available", nil)
|
||||
}
|
||||
|
||||
baseURL := strings.TrimSpace(account.GetCredential("base_url"))
|
||||
if baseURL == "" {
|
||||
baseURL = "https://api.x.ai"
|
||||
}
|
||||
normalizedBaseURL, err := s.validateUpstreamBaseURL(baseURL)
|
||||
if err != nil {
|
||||
return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, buildOpenAIModelsURL(normalizedBaseURL), nil)
|
||||
if err != nil {
|
||||
return nil, newUpstreamModelSyncConfigError("Invalid Grok model list URL", err)
|
||||
}
|
||||
req.Header.Set("Accept", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
account.ApplyHeaderOverrides(req.Header)
|
||||
return req, nil
|
||||
}
|
||||
|
||||
func (s *AccountTestService) buildAnthropicUpstreamModelsRequest(ctx context.Context, account *Account) (*http.Request, error) {
|
||||
if account.IsBedrock() || account.Type == AccountTypeServiceAccount {
|
||||
return nil, newUpstreamModelSyncUnsupportedError(
|
||||
|
||||
@@ -177,6 +177,18 @@ func TestBuildUpstreamModelsRequestsForAPIKeyAccounts(t *testing.T) {
|
||||
require.Equal(t, "https://openai.example.com/v1/models", openAIReq.URL.String())
|
||||
require.Equal(t, "Bearer openai-key", openAIReq.Header.Get("Authorization"))
|
||||
|
||||
grokReq, err := svc.buildUpstreamModelsRequest(ctx, &Account{
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeAPIKey,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "xai-key",
|
||||
"base_url": "https://xai.example.com/v1",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https://xai.example.com/v1/models", grokReq.URL.String())
|
||||
require.Equal(t, "Bearer xai-key", grokReq.Header.Get("Authorization"))
|
||||
|
||||
geminiReq, err := svc.buildGeminiUpstreamModelsRequest(ctx, &Account{
|
||||
Platform: PlatformGemini,
|
||||
Type: AccountTypeAPIKey,
|
||||
@@ -202,6 +214,22 @@ func TestBuildUpstreamModelsRequestsForAPIKeyAccounts(t *testing.T) {
|
||||
require.Equal(t, "antigravity-key", antigravityReq.Header.Get("x-api-key"))
|
||||
}
|
||||
|
||||
func TestBuildUpstreamModelsRequestRejectsGrokOAuth(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
svc := &AccountTestService{cfg: upstreamModelSyncTestConfig()}
|
||||
_, err := svc.buildUpstreamModelsRequest(context.Background(), &Account{
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
})
|
||||
require.Error(t, err)
|
||||
|
||||
var syncErr *UpstreamModelSyncError
|
||||
require.True(t, errors.As(err, &syncErr))
|
||||
require.Equal(t, UpstreamModelSyncErrorUnsupported, syncErr.Kind)
|
||||
require.Contains(t, syncErr.SafeMessage(), "Unsupported Grok account type")
|
||||
}
|
||||
|
||||
func TestBuildAntigravityAPIKeyModelsRequestRejectsOfficialCloudCodeBase(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -265,6 +293,34 @@ func TestFetchUpstreamSupportedModelsParsesOpenAIResponse(t *testing.T) {
|
||||
require.Equal(t, "Bearer openai-key", upstream.lastReq.Header.Get("Authorization"))
|
||||
}
|
||||
|
||||
func TestFetchUpstreamSupportedModelsParsesGrokAPIKeyResponse(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"data":[{"id":"grok-4.5"},{"id":"grok-4.5"},{"id":"grok-imagine"}]}`)),
|
||||
}}
|
||||
svc := &AccountTestService{
|
||||
httpUpstream: upstream,
|
||||
cfg: upstreamModelSyncTestConfig(),
|
||||
}
|
||||
|
||||
models, err := svc.FetchUpstreamSupportedModels(context.Background(), &Account{
|
||||
ID: 9,
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeAPIKey,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "xai-key",
|
||||
"base_url": "https://xai.example.com/v1",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"grok-4.5", "grok-imagine"}, models)
|
||||
require.Equal(t, "https://xai.example.com/v1/models", upstream.lastReq.URL.String())
|
||||
require.Equal(t, "Bearer xai-key", upstream.lastReq.Header.Get("Authorization"))
|
||||
}
|
||||
|
||||
func TestFetchUpstreamSupportedModelsDoesNotExposeUpstreamBody(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -109,7 +109,8 @@ func (s *FrontendServer) Middleware() gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// Serve static files normally
|
||||
// Serve static files normally (hashed assets get long-lived cache headers)
|
||||
applyStaticAssetCacheHeaders(c.Writer.Header(), cleanPath)
|
||||
s.fileServer.ServeHTTP(c.Writer, c.Request)
|
||||
c.Abort()
|
||||
}
|
||||
@@ -135,6 +136,7 @@ func (s *FrontendServer) tryServeOverride(c *gin.Context, cleanPath string) bool
|
||||
if err != nil || info.IsDir() {
|
||||
return false
|
||||
}
|
||||
applyStaticAssetCacheHeaders(c.Writer.Header(), cleanPath)
|
||||
c.File(filePath)
|
||||
c.Abort()
|
||||
return true
|
||||
@@ -273,6 +275,7 @@ func ServeEmbeddedFrontend() gin.HandlerFunc {
|
||||
if tryServeOverrideFile(c, overrideDir, cleanPath) {
|
||||
return
|
||||
}
|
||||
applyStaticAssetCacheHeaders(c.Writer.Header(), cleanPath)
|
||||
fileServer.ServeHTTP(c.Writer, c.Request)
|
||||
c.Abort()
|
||||
return
|
||||
@@ -292,6 +295,7 @@ func tryServeOverrideFile(c *gin.Context, overrideDir, cleanPath string) bool {
|
||||
if err != nil || info.IsDir() {
|
||||
return false
|
||||
}
|
||||
applyStaticAssetCacheHeaders(c.Writer.Header(), cleanPath)
|
||||
c.File(filePath)
|
||||
c.Abort()
|
||||
return true
|
||||
@@ -308,6 +312,7 @@ func shouldBypassEmbeddedFrontend(path string) bool {
|
||||
trimmed == "/health" ||
|
||||
trimmed == "/responses" ||
|
||||
strings.HasPrefix(trimmed, "/responses/") ||
|
||||
trimmed == "/alpha/search" ||
|
||||
strings.HasPrefix(trimmed, "/images/") ||
|
||||
strings.HasPrefix(trimmed, "/videos/")
|
||||
}
|
||||
|
||||
@@ -507,6 +507,32 @@ func TestFrontendServer_Middleware(t *testing.T) {
|
||||
assert.JSONEq(t, `{"ok":true}`, w.Body.String())
|
||||
})
|
||||
|
||||
t.Run("skips_alpha_search_post_route", func(t *testing.T) {
|
||||
provider := &mockSettingsProvider{
|
||||
settings: map[string]string{"test": "value"},
|
||||
}
|
||||
|
||||
server, err := NewFrontendServer(provider)
|
||||
require.NoError(t, err)
|
||||
|
||||
router := gin.New()
|
||||
router.Use(server.Middleware())
|
||||
nextCalled := false
|
||||
router.POST("/alpha/search", func(c *gin.Context) {
|
||||
nextCalled = true
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
})
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodPost, "/alpha/search", strings.NewReader(`{"model":"gpt-5.6-sol"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.True(t, nextCalled, "next handler should be called for alpha search API route")
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
assert.JSONEq(t, `{"ok":true}`, w.Body.String())
|
||||
})
|
||||
|
||||
t.Run("serves_index_for_spa_routes", func(t *testing.T) {
|
||||
provider := &mockSettingsProvider{
|
||||
settings: map[string]string{"test": "value"},
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
//go:build embed || unit
|
||||
|
||||
package web
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// staticAssetsCacheControl matches deploy/Caddyfile for hashed frontend assets.
|
||||
// Vite emits content-hashed filenames under assets/, so long-lived immutable
|
||||
// caching is safe without relying on a reverse proxy.
|
||||
const staticAssetsCacheControl = "public, max-age=31536000, immutable"
|
||||
|
||||
// isLongCacheStaticPath reports whether a cleaned URL path (no leading slash)
|
||||
// should receive long-lived Cache-Control headers. Aligned with deploy/Caddyfile.
|
||||
func isLongCacheStaticPath(cleanPath string) bool {
|
||||
cleanPath = strings.TrimPrefix(cleanPath, "/")
|
||||
return strings.HasPrefix(cleanPath, "assets/") ||
|
||||
cleanPath == "logo.png" ||
|
||||
cleanPath == "favicon.ico"
|
||||
}
|
||||
|
||||
// applyStaticAssetCacheHeaders sets Cache-Control for long-cacheable static paths.
|
||||
// index.html / SPA routes must keep no-cache and are not handled here.
|
||||
func applyStaticAssetCacheHeaders(header http.Header, cleanPath string) {
|
||||
if header == nil || !isLongCacheStaticPath(cleanPath) {
|
||||
return
|
||||
}
|
||||
header.Set("Cache-Control", staticAssetsCacheControl)
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
//go:build unit
|
||||
|
||||
package web
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestIsLongCacheStaticPath(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
path string
|
||||
want bool
|
||||
}{
|
||||
{name: "hashed_js", path: "assets/index-abc123.js", want: true},
|
||||
{name: "hashed_css", path: "assets/app-def456.css", want: true},
|
||||
{name: "nested_asset", path: "assets/vendor/chunk.js", want: true},
|
||||
{name: "leading_slash_asset", path: "/assets/index.js", want: true},
|
||||
{name: "logo", path: "logo.png", want: true},
|
||||
{name: "favicon", path: "favicon.ico", want: true},
|
||||
{name: "index_html", path: "index.html", want: false},
|
||||
{name: "spa_route", path: "dashboard", want: false},
|
||||
{name: "assets_prefix_only", path: "assets", want: false},
|
||||
{name: "similar_name", path: "assets-backup/x.js", want: false},
|
||||
{name: "empty", path: "", want: false},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
assert.Equal(t, tc.want, isLongCacheStaticPath(tc.path))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyStaticAssetCacheHeaders(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("sets_immutable_cache_for_assets", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
header := make(http.Header)
|
||||
applyStaticAssetCacheHeaders(header, "assets/index-abc.js")
|
||||
assert.Equal(t, staticAssetsCacheControl, header.Get("Cache-Control"))
|
||||
})
|
||||
|
||||
t.Run("sets_immutable_cache_for_logo", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
header := make(http.Header)
|
||||
applyStaticAssetCacheHeaders(header, "logo.png")
|
||||
assert.Equal(t, staticAssetsCacheControl, header.Get("Cache-Control"))
|
||||
})
|
||||
|
||||
t.Run("skips_index_html", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
header := make(http.Header)
|
||||
applyStaticAssetCacheHeaders(header, "index.html")
|
||||
assert.Empty(t, header.Get("Cache-Control"))
|
||||
})
|
||||
|
||||
t.Run("nil_header_is_noop", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
assert.NotPanics(t, func() {
|
||||
applyStaticAssetCacheHeaders(nil, "assets/x.js")
|
||||
})
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
-- Support the per-key latest non-empty source IP lookup without scanning full key history.
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_api_key_latest_ip
|
||||
ON usage_logs (api_key_id, created_at DESC, id DESC)
|
||||
INCLUDE (ip_address)
|
||||
WHERE ip_address IS NOT NULL AND ip_address <> '';
|
||||
@@ -0,0 +1,19 @@
|
||||
package migrations
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestLatestAPIKeyIPIndexMigration(t *testing.T) {
|
||||
content, err := FS.ReadFile("174_add_usage_logs_api_key_latest_ip_index_notx.sql")
|
||||
require.NoError(t, err)
|
||||
|
||||
sql := strings.Join(strings.Fields(string(content)), " ")
|
||||
require.Contains(t, sql, "CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_api_key_latest_ip")
|
||||
require.Contains(t, sql, "ON usage_logs (api_key_id, created_at DESC, id DESC)")
|
||||
require.Contains(t, sql, "INCLUDE (ip_address)")
|
||||
require.Contains(t, sql, "WHERE ip_address IS NOT NULL AND ip_address <> ''")
|
||||
}
|
||||
Reference in New Issue
Block a user