mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-29 12:09:31 +08:00
fix(go): implement Anthropic streaming (ChatStreamlyWithSender) (#17380)
### Summary The Go Anthropic driver's `ChatStreamlyWithSender` (internal/entity/models/anthropic.go) was a stub that always returned `"no such method"`, so any caller requesting a streamed response from a Claude model via the Go path failed outright — diverging from the Python `AnthropicCV` driver, which already supports streaming. This implements the method by opening the Messages API with `stream=true` and parsing the SSE response via the shared `ParseSSEStream` helper, forwarding `text_delta`/`thinking_delta` content through the `sender` callback and treating `message_stop` as the terminal event — consistent with the other Go drivers in this package (e.g. Cohere). Fixes #17333 --------- Co-authored-by: Abhay Yadav <abhayyadav@Abhays-MacBook-Air.local>
This commit is contained in:
@@ -109,7 +109,7 @@ func (a *AnthropicModel) ChatWithMessages(ctx context.Context, modelName string,
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
setAnthropicHeaders(req, apiKey)
|
||||
setAnthropicHeaders(req, apiKey, false)
|
||||
|
||||
resp, err := a.baseModel.httpClient.Do(req)
|
||||
if err != nil {
|
||||
@@ -153,9 +153,13 @@ func applyAnthropicChatConfig(reqBody map[string]interface{}, chatModelConfig *C
|
||||
}
|
||||
}
|
||||
|
||||
func setAnthropicHeaders(req *http.Request, apiKey string) {
|
||||
func setAnthropicHeaders(req *http.Request, apiKey string, streaming bool) {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
if streaming {
|
||||
req.Header.Set("Accept", "text/event-stream")
|
||||
} else {
|
||||
req.Header.Set("Accept", "application/json")
|
||||
}
|
||||
req.Header.Set("x-api-key", apiKey)
|
||||
req.Header.Set("anthropic-version", anthropicVersion)
|
||||
}
|
||||
@@ -393,7 +397,7 @@ func (a *AnthropicModel) ListModels(ctx context.Context, apiConfig *APIConfig) (
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
setAnthropicHeaders(req, apiKey)
|
||||
setAnthropicHeaders(req, apiKey, false)
|
||||
|
||||
resp, err := a.baseModel.httpClient.Do(req)
|
||||
if err != nil {
|
||||
@@ -434,7 +438,137 @@ func (a *AnthropicModel) CheckConnection(ctx context.Context, apiConfig *APIConf
|
||||
}
|
||||
|
||||
func (a *AnthropicModel) ChatStreamlyWithSender(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, modelConfig *ChatConfig, modelUsage *common.ModelUsage, sender func(*string, *string) error) error {
|
||||
return fmt.Errorf("%s, no such method", a.Name())
|
||||
if err := a.baseModel.APIConfigCheck(apiConfig); err != nil {
|
||||
return err
|
||||
}
|
||||
apiKey := strings.TrimSpace(*apiConfig.ApiKey)
|
||||
if len(messages) == 0 {
|
||||
return fmt.Errorf("messages is empty")
|
||||
}
|
||||
|
||||
apiMessages, systemPrompt, err := anthropicMessages(messages)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
baseURLRegion := a.region(apiConfig)
|
||||
baseURLConfig := &APIConfig{Region: &baseURLRegion}
|
||||
if apiConfig != nil {
|
||||
baseURLConfig.BaseURL = apiConfig.BaseURL
|
||||
}
|
||||
baseURL, err := a.baseModel.GetBaseURL(baseURLConfig)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
baseURL = strings.TrimSpace(strings.TrimSuffix(baseURL, "/"))
|
||||
url := fmt.Sprintf("%s/%s", baseURL, strings.TrimLeft(a.baseModel.URLSuffix.Chat, "/"))
|
||||
|
||||
reqBody := map[string]interface{}{
|
||||
"model": modelName,
|
||||
"messages": apiMessages,
|
||||
"max_tokens": 1024,
|
||||
"stream": true,
|
||||
}
|
||||
if systemPrompt != "" {
|
||||
reqBody["system"] = systemPrompt
|
||||
}
|
||||
applyAnthropicChatConfig(reqBody, modelConfig)
|
||||
|
||||
jsonData, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, streamCallTimeout)
|
||||
defer cancel()
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
setAnthropicHeaders(req, apiKey, true)
|
||||
|
||||
resp, err := a.baseModel.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to send request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return fmt.Errorf("Anthropic messages API error: %s, body: %s", resp.Status, string(body))
|
||||
}
|
||||
|
||||
sawTerminal := false
|
||||
var streamUsage TokenUsage
|
||||
sawUsage := false
|
||||
done, err := ParseSSEStream[map[string]interface{}](resp.Body, func(event map[string]interface{}) error {
|
||||
eventType, _ := event["type"].(string)
|
||||
switch eventType {
|
||||
case "content_block_delta":
|
||||
delta, ok := event["delta"].(map[string]interface{})
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
deltaType, _ := delta["type"].(string)
|
||||
switch deltaType {
|
||||
case "text_delta":
|
||||
if text, ok := delta["text"].(string); ok && text != "" {
|
||||
if err := sender(&text, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
case "thinking_delta":
|
||||
if thinking, ok := delta["thinking"].(string); ok && thinking != "" {
|
||||
if err := sender(nil, &thinking); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
case "message_start":
|
||||
message, ok := event["message"].(map[string]interface{})
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
if usage, ok := message["usage"].(map[string]interface{}); ok {
|
||||
if inputTokens, ok := usage["input_tokens"].(float64); ok {
|
||||
streamUsage.PromptTokens = int(inputTokens)
|
||||
sawUsage = true
|
||||
}
|
||||
}
|
||||
case "message_delta":
|
||||
// message_delta carries the running total of output tokens
|
||||
// generated so far; the last event before message_stop is
|
||||
// authoritative.
|
||||
if usage, ok := event["usage"].(map[string]interface{}); ok {
|
||||
if outputTokens, ok := usage["output_tokens"].(float64); ok {
|
||||
streamUsage.CompletionTokens = int(outputTokens)
|
||||
sawUsage = true
|
||||
}
|
||||
}
|
||||
case "message_stop":
|
||||
sawTerminal = true
|
||||
case "error":
|
||||
errInfo, _ := event["error"].(map[string]interface{})
|
||||
message, _ := errInfo["message"].(string)
|
||||
return fmt.Errorf("Anthropic stream error: %s", message)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to scan response body: %w", err)
|
||||
}
|
||||
if !done && !sawTerminal {
|
||||
return fmt.Errorf("anthropic: stream ended before message_stop")
|
||||
}
|
||||
|
||||
if sawUsage {
|
||||
streamUsage.TotalTokens = streamUsage.PromptTokens + streamUsage.CompletionTokens
|
||||
applyStreamUsage(modelConfig, modelUsage, &streamUsage)
|
||||
}
|
||||
|
||||
endOfStream := "[DONE]"
|
||||
return sender(&endOfStream, nil)
|
||||
}
|
||||
|
||||
func (a *AnthropicModel) Embed(ctx context.Context, modelName *string, texts []string, apiConfig *APIConfig, embeddingConfig *EmbeddingConfig, modelUsage *common.ModelUsage) ([]EmbeddingData, error) {
|
||||
|
||||
@@ -291,6 +291,124 @@ func TestAnthropicChatRejectsMalformedResponse(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnthropicChatStreamlyWithSenderHappyPath(t *testing.T) {
|
||||
srv := newAnthropicServer(t, "/v1/messages", func(t *testing.T, body map[string]interface{}, w http.ResponseWriter) {
|
||||
if body["stream"] != true {
|
||||
t.Errorf("stream=%v, want true", body["stream"])
|
||||
}
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
events := []string{
|
||||
`{"type":"message_start","message":{"id":"msg_1","usage":{"input_tokens":5,"output_tokens":0}}}`,
|
||||
`{"type":"content_block_start","index":0,"content_block":{"type":"thinking","thinking":""}}`,
|
||||
`{"type":"content_block_delta","index":0,"delta":{"type":"thinking_delta","thinking":"pondering"}}`,
|
||||
`{"type":"content_block_stop","index":0}`,
|
||||
`{"type":"content_block_start","index":1,"content_block":{"type":"text","text":""}}`,
|
||||
`{"type":"content_block_delta","index":1,"delta":{"type":"text_delta","text":"pong"}}`,
|
||||
`{"type":"content_block_stop","index":1}`,
|
||||
`{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":2}}`,
|
||||
`{"type":"message_stop"}`,
|
||||
}
|
||||
for _, event := range events {
|
||||
_, _ = w.Write([]byte("data: " + event + "\n\n"))
|
||||
}
|
||||
})
|
||||
defer srv.Close()
|
||||
ctx := t.Context()
|
||||
|
||||
apiKey := "test-key"
|
||||
var answer, reasoning strings.Builder
|
||||
sawDone := false
|
||||
modelUsage := &common.ModelUsage{}
|
||||
err := newAnthropicForTest(srv.URL).ChatStreamlyWithSender(
|
||||
ctx,
|
||||
"claude-sonnet-4-5-20250929",
|
||||
[]Message{{Role: "user", Content: "ping"}},
|
||||
&APIConfig{ApiKey: &apiKey},
|
||||
&ChatConfig{},
|
||||
modelUsage,
|
||||
func(text, reason *string) error {
|
||||
if text != nil {
|
||||
if *text == "[DONE]" {
|
||||
sawDone = true
|
||||
} else {
|
||||
answer.WriteString(*text)
|
||||
}
|
||||
}
|
||||
if reason != nil {
|
||||
reasoning.WriteString(*reason)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("ChatStreamlyWithSender: %v", err)
|
||||
}
|
||||
if answer.String() != "pong" {
|
||||
t.Errorf("answer=%q, want pong", answer.String())
|
||||
}
|
||||
if reasoning.String() != "pondering" {
|
||||
t.Errorf("reasoning=%q, want pondering", reasoning.String())
|
||||
}
|
||||
if !sawDone {
|
||||
t.Error("expected terminal [DONE] sender call")
|
||||
}
|
||||
if modelUsage.InputTokens != 5 {
|
||||
t.Errorf("InputTokens=%d, want 5", modelUsage.InputTokens)
|
||||
}
|
||||
if modelUsage.OutputTokens != 2 {
|
||||
t.Errorf("OutputTokens=%d, want 2", modelUsage.OutputTokens)
|
||||
}
|
||||
if modelUsage.TotalTokens != 7 {
|
||||
t.Errorf("TotalTokens=%d, want 7", modelUsage.TotalTokens)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnthropicChatStreamlyWithSenderRejectsHTTPError(t *testing.T) {
|
||||
srv := newAnthropicServer(t, "/v1/messages", func(t *testing.T, _ map[string]interface{}, w http.ResponseWriter) {
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
_, _ = w.Write([]byte(`{"error":{"message":"bad key"}}`))
|
||||
})
|
||||
defer srv.Close()
|
||||
ctx := t.Context()
|
||||
|
||||
apiKey := "test-key"
|
||||
err := newAnthropicForTest(srv.URL).ChatStreamlyWithSender(
|
||||
ctx,
|
||||
"claude",
|
||||
[]Message{{Role: "user", Content: "x"}},
|
||||
&APIConfig{ApiKey: &apiKey},
|
||||
nil,
|
||||
nil,
|
||||
func(*string, *string) error { return nil },
|
||||
)
|
||||
if err == nil || !strings.Contains(err.Error(), "401") || !strings.Contains(err.Error(), "bad key") {
|
||||
t.Errorf("expected provider error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnthropicChatStreamlyWithSenderRejectsStreamError(t *testing.T) {
|
||||
srv := newAnthropicServer(t, "/v1/messages", func(t *testing.T, _ map[string]interface{}, w http.ResponseWriter) {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = w.Write([]byte(`data: {"type":"error","error":{"message":"overloaded"}}` + "\n\n"))
|
||||
})
|
||||
defer srv.Close()
|
||||
ctx := t.Context()
|
||||
|
||||
apiKey := "test-key"
|
||||
err := newAnthropicForTest(srv.URL).ChatStreamlyWithSender(
|
||||
ctx,
|
||||
"claude",
|
||||
[]Message{{Role: "user", Content: "x"}},
|
||||
&APIConfig{ApiKey: &apiKey},
|
||||
nil,
|
||||
nil,
|
||||
func(*string, *string) error { return nil },
|
||||
)
|
||||
if err == nil || !strings.Contains(err.Error(), "overloaded") {
|
||||
t.Errorf("expected stream error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnthropicListModelsAndCheckConnection(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
var calls int
|
||||
@@ -352,17 +470,6 @@ func TestAnthropicUnsupportedMethods(t *testing.T) {
|
||||
m := newAnthropicForTest("http://unused")
|
||||
apiKey := "test-key"
|
||||
modelName := "claude"
|
||||
checks := []struct {
|
||||
name string
|
||||
err error
|
||||
}{
|
||||
{"stream", m.ChatStreamlyWithSender(ctx, modelName, []Message{{Role: "user", Content: "x"}}, &APIConfig{ApiKey: &apiKey}, nil, nil, func(*string, *string) error { return nil })},
|
||||
}
|
||||
for _, check := range checks {
|
||||
if check.err == nil || !strings.Contains(check.err.Error(), "no such method") {
|
||||
t.Errorf("%s: want no such method, got %v", check.name, check.err)
|
||||
}
|
||||
}
|
||||
if _, err := m.Embed(ctx, &modelName, []string{"x"}, &APIConfig{ApiKey: &apiKey}, nil, nil); err == nil || !strings.Contains(err.Error(), "no such method") {
|
||||
t.Errorf("Embed: got %v", err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user