Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions internal/aistudio/accounts.go
Original file line number Diff line number Diff line change
Expand Up @@ -214,6 +214,7 @@ type AccountSelection struct {
ResourceID string
AllowedAccountIDs []string
PlaygroundOnly bool
Channel Channel
}

const preferredBootstrapModelID = "gemini-flash-latest"
Expand Down
2 changes: 1 addition & 1 deletion internal/aistudio/build.go
Original file line number Diff line number Diff line change
Expand Up @@ -779,7 +779,7 @@ func buildTrailerError(raw json.RawMessage) error {

// sendBuild 编码并发送 Build 代理请求,返回响应与含 finishReason 校验的解码
func (c *Client) sendBuild(ctx context.Context, request GenerateRequest, entry modelEntry) (*RPCResponse, func(io.Reader, func(Event) error) error, error) {
unary := buildUsesUnary(entry.model)
unary := request.Unary || buildUsesUnary(entry.model)
path, body, err := EncodeBuildGenerateRequest(request, entry.defaults, request.ImageRoute, unary)
if err != nil {
return nil, nil, fmt.Errorf("%w: %v", ErrInvalidArgument, err)
Expand Down
3 changes: 3 additions & 0 deletions internal/aistudio/channel.go
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,9 @@ func generationChannelSelection(selection AccountSelection) bool {

// selectionChannelsLocked 返回选择可使用的通道顺序;非生成请求只使用 Playground RPC
func (p *AccountPool) selectionChannelsLocked(selection AccountSelection) []Channel {
if selection.Channel != "" && p.channelEnabledLocked(selection.Channel) {
return []Channel{selection.Channel}
}
if !generationChannelSelection(selection) {
return []Channel{ChannelPlayground}
}
Expand Down
5 changes: 5 additions & 0 deletions internal/aistudio/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -550,11 +550,16 @@ func (s *PooledService) Generate(ctx context.Context, request GenerateRequest) (
if err != nil {
return nil, err
}
channel := request.Channel
if request.Unary && channel == "" {
channel = ChannelBuild
}
selection := AccountSelection{
ModelID: modelID,
Method: "generateContent",
AccountID: strings.TrimSpace(request.AccountID),
ResourceID: resourceID,
Channel: channel,
}
pinned := selection.AccountID != "" || selection.ResourceID != ""
if _, ok := AccountLeaseFromContext(ctx); ok {
Expand Down
4 changes: 4 additions & 0 deletions internal/aistudio/types.go
Original file line number Diff line number Diff line change
Expand Up @@ -235,6 +235,10 @@ type GenerateRequest struct {
AccountID string `json:"account_id,omitempty"`
// ImageRoute 内部标记:图像生成模型(由模型目录能力推导)
ImageRoute bool `json:"-"`
// Unary 标记单次非流式请求(在 Build 代理中使用 ProxyUnaryCall)
Unary bool `json:"-"`
// Channel 指定请求优先或强制使用的上游通道(为空时按 AccountPool 规则调度)
Channel Channel `json:"-"`
}

// TokenCountRequest 表示计数请求
Expand Down
1 change: 1 addition & 0 deletions internal/api/anthropic.go
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,7 @@ func (s *server) handleAnthropicMessages(w http.ResponseWriter, r *http.Request)
writeAnthropicError(w, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
generateRequest.Unary = !request.Stream
if request.Stream {
if err := streamHeaders(w); err != nil {
return
Expand Down
67 changes: 56 additions & 11 deletions internal/api/gemini.go
Original file line number Diff line number Diff line change
Expand Up @@ -780,6 +780,7 @@ func (s *server) handleGeminiCountTokens(w http.ResponseWriter, r *http.Request,
}

func (s *server) handleGeminiGenerate(w http.ResponseWriter, r *http.Request, request aistudio.GenerateRequest, stream bool) {
request.Unary = !stream
events, err := s.service.Generate(r.Context(), request)
if err != nil {
if shouldWriteRequestError(r, err) {
Expand Down Expand Up @@ -828,46 +829,90 @@ func buildGeminiResponse(request aistudio.GenerateRequest, result generationResu

func geminiOutputParts(result generationResult) []map[string]any {
parts := make([]map[string]any, 0)
var pendingSignature string

attachSignature := func(part map[string]any, sig string) map[string]any {
if sig != "" {
part["thoughtSignature"] = sig
} else if pendingSignature != "" {
part["thoughtSignature"] = pendingSignature
pendingSignature = ""
}
return part
}

for _, event := range result.events {
switch event.Kind {
case aistudio.EventText:
parts = append(parts, geminiSignedPart(geminiTextPart(event), event.ThoughtSignature))
if len(parts) > 0 {
last := parts[len(parts)-1]
if lastText, ok := last["text"].(string); ok && last["thought"] != true && last["transcriptionMetadata"] == nil && event.Transcript == nil {
last["text"] = lastText + event.Text
if event.ThoughtSignature != "" {
last["thoughtSignature"] = event.ThoughtSignature
}
continue
}
}
parts = append(parts, attachSignature(geminiTextPart(event), event.ThoughtSignature))
case aistudio.EventReasoning:
parts = append(parts, geminiSignedPart(map[string]any{"text": event.Text, "thought": true}, event.ThoughtSignature))
if len(parts) > 0 {
last := parts[len(parts)-1]
if lastText, ok := last["text"].(string); ok && last["thought"] == true {
last["text"] = lastText + event.Text
if event.ThoughtSignature != "" {
last["thoughtSignature"] = event.ThoughtSignature
}
continue
}
}
parts = append(parts, attachSignature(map[string]any{"text": event.Text, "thought": true}, event.ThoughtSignature))
case aistudio.EventToolCall:
if event.ToolCall != nil {
parts = append(parts, geminiSignedPart(geminiFunctionCallPart(*event.ToolCall), event.ThoughtSignature))
parts = append(parts, attachSignature(geminiFunctionCallPart(*event.ToolCall), event.ThoughtSignature))
}
case aistudio.EventExecutableCode:
if event.ExecutableCode != nil {
parts = append(parts, geminiSignedPart(map[string]any{"executableCode": map[string]any{
parts = append(parts, attachSignature(map[string]any{"executableCode": map[string]any{
"language": event.ExecutableCode.Language, "code": event.ExecutableCode.Code,
}}, event.ThoughtSignature))
}
case aistudio.EventCodeExecutionResult:
if event.CodeExecutionResult != nil {
parts = append(parts, geminiSignedPart(map[string]any{
parts = append(parts, attachSignature(map[string]any{
"codeExecutionResult": geminiCodeExecutionResult(*event.CodeExecutionResult),
}, event.ThoughtSignature))
}
case aistudio.EventMedia:
if event.Media != nil {
var part map[string]any
if len(event.Media.Data) > 0 {
parts = append(parts, geminiSignedPart(map[string]any{"inlineData": map[string]any{
part = map[string]any{"inlineData": map[string]any{
"mimeType": event.Media.MIME, "data": base64.StdEncoding.EncodeToString(event.Media.Data),
}}, event.ThoughtSignature))
}}
} else if event.Media.URL != "" {
parts = append(parts, geminiSignedPart(map[string]any{"fileData": map[string]any{
part = map[string]any{"fileData": map[string]any{
"mimeType": event.Media.MIME, "fileUri": event.Media.URL, "displayName": event.Media.Name,
}}, event.ThoughtSignature))
}}
}
if part != nil {
parts = append(parts, attachSignature(part, event.ThoughtSignature))
}
}
case aistudio.EventThoughtSignature:
if event.ThoughtSignature != "" {
parts = append(parts, map[string]any{"thoughtSignature": event.ThoughtSignature})
if event.ThoughtSignature == "" {
continue
}
if len(parts) > 0 {
parts[len(parts)-1]["thoughtSignature"] = event.ThoughtSignature
} else {
pendingSignature = event.ThoughtSignature
}
}
}
if len(parts) == 0 && pendingSignature != "" {
parts = append(parts, map[string]any{"thoughtSignature": pendingSignature})
}
return parts
}

Expand Down
138 changes: 138 additions & 0 deletions internal/api/gemini_output_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,138 @@
package api

import (
"encoding/json"
"reflect"
"testing"

"github.com/Mag1cFall/AIStudio2API/internal/aistudio"
)

func TestGeminiOutputParts_MergeTextAndReasoning(t *testing.T) {
result := generationResult{
events: []aistudio.Event{
{Kind: aistudio.EventReasoning, Text: "I think "},
{Kind: aistudio.EventReasoning, Text: "therefore "},
{Kind: aistudio.EventReasoning, Text: "I am."},
{Kind: aistudio.EventThoughtSignature, ThoughtSignature: "sig_abc123"},
{Kind: aistudio.EventText, Text: "Hello "},
{Kind: aistudio.EventText, Text: "world!"},
},
}

parts := geminiOutputParts(result)
if len(parts) != 2 {
t.Fatalf("expected 2 parts, got %d: %+v", len(parts), parts)
}

// First part: merged thought
expectedThought := map[string]any{
"thought": true,
"text": "I think therefore I am.",
"thoughtSignature": "sig_abc123",
}
if !reflect.DeepEqual(parts[0], expectedThought) {
t.Errorf("parts[0] = %+v, want %+v", parts[0], expectedThought)
}

// Second part: merged text
expectedText := map[string]any{
"text": "Hello world!",
}
if !reflect.DeepEqual(parts[1], expectedText) {
t.Errorf("parts[1] = %+v, want %+v", parts[1], expectedText)
}
}

func TestGeminiOutputParts_TextOnly(t *testing.T) {
result := generationResult{
events: []aistudio.Event{
{Kind: aistudio.EventText, Text: "Chunk 1 "},
{Kind: aistudio.EventText, Text: "Chunk 2 "},
{Kind: aistudio.EventText, Text: "Chunk 3"},
},
}

parts := geminiOutputParts(result)
if len(parts) != 1 {
t.Fatalf("expected 1 part, got %d: %+v", len(parts), parts)
}

expectedText := map[string]any{
"text": "Chunk 1 Chunk 2 Chunk 3",
}
if !reflect.DeepEqual(parts[0], expectedText) {
t.Errorf("parts[0] = %+v, want %+v", parts[0], expectedText)
}
}

func TestGeminiOutputParts_LeadingThoughtSignature(t *testing.T) {
result := generationResult{
events: []aistudio.Event{
{Kind: aistudio.EventThoughtSignature, ThoughtSignature: "sig_pre"},
{Kind: aistudio.EventReasoning, Text: "Thinking step."},
{Kind: aistudio.EventText, Text: "Final text."},
},
}

parts := geminiOutputParts(result)
if len(parts) != 2 {
t.Fatalf("expected 2 parts, got %d: %+v", len(parts), parts)
}

if parts[0]["thoughtSignature"] != "sig_pre" {
t.Errorf("expected thoughtSignature 'sig_pre' on parts[0], got %v", parts[0]["thoughtSignature"])
}
if parts[0]["text"] != "Thinking step." {
t.Errorf("expected text 'Thinking step.', got %v", parts[0]["text"])
}
if parts[1]["text"] != "Final text." {
t.Errorf("expected text 'Final text.', got %v", parts[1]["text"])
}
}

func TestGeminiOutputParts_StandaloneThoughtSignatureFallback(t *testing.T) {
result := generationResult{
events: []aistudio.Event{
{Kind: aistudio.EventThoughtSignature, ThoughtSignature: "sig_only"},
},
}

parts := geminiOutputParts(result)
if len(parts) != 1 {
t.Fatalf("expected 1 part, got %d: %+v", len(parts), parts)
}
if parts[0]["thoughtSignature"] != "sig_only" {
t.Errorf("expected thoughtSignature 'sig_only', got %v", parts[0]["thoughtSignature"])
}
}

func TestGeminiOutputParts_WithToolCall(t *testing.T) {
call := aistudio.FunctionCall{
ID: "call_1",
Name: "get_weather",
Arguments: json.RawMessage(`{"location":"Tokyo"}`),
}
result := generationResult{
events: []aistudio.Event{
{Kind: aistudio.EventReasoning, Text: "Need weather for Tokyo."},
{Kind: aistudio.EventThoughtSignature, ThoughtSignature: "sig_call"},
{Kind: aistudio.EventToolCall, ToolCall: &call},
},
}

parts := geminiOutputParts(result)
if len(parts) != 2 {
t.Fatalf("expected 2 parts, got %d: %+v", len(parts), parts)
}

if parts[0]["thought"] != true || parts[0]["text"] != "Need weather for Tokyo." {
t.Errorf("unexpected parts[0]: %+v", parts[0])
}
if parts[0]["thoughtSignature"] != "sig_call" {
t.Errorf("expected sig_call on parts[0], got %v", parts[0]["thoughtSignature"])
}
if parts[1]["functionCall"] == nil {
t.Errorf("expected functionCall on parts[1], got %+v", parts[1])
}
}
1 change: 1 addition & 0 deletions internal/api/openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -135,6 +135,7 @@ func (s *server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
writeOpenAIError(w, http.StatusBadRequest, "invalid_request", err.Error())
return
}
generateRequest.Unary = !request.Stream
s.thoughtSignatures.Restore(generateRequest.Contents)
events, err := s.service.Generate(r.Context(), generateRequest)
if err != nil {
Expand Down
1 change: 1 addition & 0 deletions internal/api/responses.go
Original file line number Diff line number Diff line change
Expand Up @@ -164,6 +164,7 @@ func (s *server) handleResponses(w http.ResponseWriter, r *http.Request) {
instructions = append(instructions, inlineInstructions...)
generateRequest.System = strings.Join(instructions, "\n")
}
generateRequest.Unary = !request.Stream
events, err := s.service.Generate(r.Context(), generateRequest)
if err != nil {
if shouldWriteRequestError(r, err) {
Expand Down