Reorganize repo: backend/, deploy/, docs/ layout + AGENTS.md

Restructure the monorepo into clear top-level directories:
- backend/: Go module root (cmd, internal, pkg, configs, migrations,
  docs/swagger, scripts, tests, go.mod, Makefile, .air.toml)
- deploy/: Docker (Dockerfile, docker-compose*), quickstart, fluentd
- docs/: project documentation + reports/ (moved from repo root)
- AGENTS.md: new AI coding-agent guide at repo root

Update all references to the new layout:
- Dockerfile: COPY backend/go.mod, COPY backend/ (context = repo root)
- docker-compose files: context ../.., dockerfile deploy/docker/Dockerfile,
  env_file ../../.env, volume mounts ../../backend:/app
- deploy/quickstart/compose.yaml: dockerfile deploy/docker/Dockerfile
- CI: working-directory: backend for go commands, file deploy/docker/Dockerfile,
  coverage path backend/coverage.out, health_check backend/scripts/
- backend/Makefile: docker target uses -f ../deploy/docker/Dockerfile ../
- README: architecture tree, quickstart, config paths updated

Move root stray scripts (rename_models.*, run_m11_tests.sh, verify_build.sh,
gorm_bool_main.go) to backend/scripts/legacy/. All moves via git mv to
preserve history. Build, vet, SQLite tests, and docker compose config verified.
This commit is contained in:
2026-07-07 14:44:12 +08:00
parent d4ef996f49
commit aeddedf2a3
1348 changed files with 176 additions and 57 deletions
+484
View File
@@ -0,0 +1,484 @@
package llm
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// --- Interface compliance test ---
func TestOpenAIProvider_ImplementsProvider(t *testing.T) {
var _ Provider = (*OpenAIProvider)(nil)
}
// --- Constructor tests ---
func TestNewOpenAIProvider_Defaults(t *testing.T) {
p := NewOpenAIProvider(OpenAIProviderConfig{
APIKey: "test-key",
})
assert.Equal(t, "test-key", p.apiKey)
assert.Equal(t, "https://api.openai.com/v1", p.baseURL)
assert.Equal(t, "gpt-4", p.model)
assert.Equal(t, "text-embedding-3-small", p.embedModel)
assert.Equal(t, 3, p.maxRetries)
}
func TestNewOpenAIProvider_CustomConfig(t *testing.T) {
p := NewOpenAIProvider(OpenAIProviderConfig{
APIKey: "custom-key",
BaseURL: "https://ark.cn-beijing.volces.com/api/v3",
Model: "doubao-pro-4k",
EmbedModel: "doubao-embedding",
MaxRetries: 5,
Timeout: 120,
})
assert.Equal(t, "custom-key", p.apiKey)
assert.Equal(t, "https://ark.cn-beijing.volces.com/api/v3", p.baseURL)
assert.Equal(t, "doubao-pro-4k", p.model)
assert.Equal(t, "doubao-embedding", p.embedModel)
assert.Equal(t, 5, p.maxRetries)
}
func TestNewOpenAIProvider_TrailingSlashTrimmed(t *testing.T) {
p := NewOpenAIProvider(OpenAIProviderConfig{
APIKey: "key",
BaseURL: "https://api.openai.com/v1/",
})
assert.Equal(t, "https://api.openai.com/v1", p.baseURL)
}
// --- Request serialization tests ---
func TestChatRequest_MarshalJSON(t *testing.T) {
req := ChatRequest{
Model: "gpt-4",
Messages: []ChatMessage{
{Role: "system", Content: "You are a helpful assistant."},
{Role: "user", Content: "Hello!"},
},
Temperature: 0.7,
MaxTokens: 100,
}
body, err := json.Marshal(req)
require.NoError(t, err)
// Verify key fields are present
assert.Contains(t, string(body), `"model":"gpt-4"`)
assert.Contains(t, string(body), `"messages"`)
assert.Contains(t, string(body), `"role":"system"`)
assert.Contains(t, string(body), `"role":"user"`)
assert.Contains(t, string(body), `"temperature":0.7`)
assert.Contains(t, string(body), `"max_tokens":100`)
}
func TestChatRequest_MarshalJSON_WithTools(t *testing.T) {
req := ChatRequest{
Model: "gpt-4",
Messages: []ChatMessage{{Role: "user", Content: "Search for docs"}},
Tools: []ToolDefinition{
{
Type: "function",
Function: ToolFunction{
Name: "search_documents",
Description: "Search knowledge base",
Parameters: map[string]interface{}{
"type": "object",
"properties": map[string]interface{}{
"query": map[string]interface{}{
"type": "string",
},
},
},
},
},
},
}
body, err := json.Marshal(req)
require.NoError(t, err)
assert.Contains(t, string(body), `"tools"`)
assert.Contains(t, string(body), `"search_documents"`)
}
func TestEmbeddingRequest_MarshalJSON(t *testing.T) {
req := EmbeddingRequest{
Model: "text-embedding-3-small",
Input: []string{"Hello world", "Test embedding"},
}
body, err := json.Marshal(req)
require.NoError(t, err)
assert.Contains(t, string(body), `"model":"text-embedding-3-small"`)
assert.Contains(t, string(body), `"Hello world"`)
assert.Contains(t, string(body), `"Test embedding"`)
}
func TestChatRequest_OmitEmptyFields(t *testing.T) {
req := ChatRequest{
Model: "gpt-4",
Messages: []ChatMessage{{Role: "user", Content: "Hi"}},
}
body, err := json.Marshal(req)
require.NoError(t, err)
// temperature=0 should NOT be omitted since 0 is the zero value for float64
// but "omitempty" will omit it — this is expected behavior
decoded := ChatRequest{}
err = json.Unmarshal(body, &decoded)
require.NoError(t, err)
assert.Equal(t, "gpt-4", decoded.Model)
assert.Equal(t, 1, len(decoded.Messages))
}
// --- Response deserialization tests ---
func TestChatResponse_UnmarshalJSON(t *testing.T) {
raw := `{
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "gpt-4",
"choices": [{
"index": 0,
"message": {"role": "assistant", "content": "Hello! How can I help you?"},
"finish_reason": "stop"
}],
"usage": {
"prompt_tokens": 9,
"completion_tokens": 12,
"total_tokens": 21
}
}`
var resp ChatResponse
err := json.Unmarshal([]byte(raw), &resp)
require.NoError(t, err)
assert.Equal(t, "chatcmpl-123", resp.ID)
assert.Equal(t, 1, len(resp.Choices))
assert.Equal(t, "assistant", resp.Choices[0].Message.Role)
assert.Equal(t, "Hello! How can I help you?", resp.Choices[0].Message.Content)
assert.Equal(t, "stop", resp.Choices[0].FinishReason)
assert.Equal(t, 21, resp.Usage.TotalTokens)
}
func TestEmbeddingResponse_UnmarshalJSON(t *testing.T) {
raw := `{
"object": "list",
"data": [{
"object": "embedding",
"index": 0,
"embedding": [0.0023064255, -0.009327292, 0.015871]
}],
"model": "text-embedding-3-small",
"usage": {
"prompt_tokens": 5,
"total_tokens": 5
}
}`
var resp EmbeddingResponse
err := json.Unmarshal([]byte(raw), &resp)
require.NoError(t, err)
assert.Equal(t, 1, len(resp.Data))
assert.Equal(t, 0, resp.Data[0].Index)
assert.Equal(t, 3, len(resp.Data[0].Embedding))
assert.InDelta(t, 0.0023064255, resp.Data[0].Embedding[0], 0.0001)
}
// --- Mock server tests ---
func TestOpenAIProvider_ChatCompletion_MockServer(t *testing.T) {
mockResp := ChatResponse{
ID: "chatcmpl-test",
Object: "chat.completion",
Model: "gpt-4",
Choices: []ChatChoice{
{
Index: 0,
Message: ChatMessage{Role: "assistant", Content: "Mocked response"},
FinishReason: "stop",
},
},
Usage: TokenUsage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15},
}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, "POST", r.Method)
assert.Equal(t, "/chat/completions", r.URL.Path)
assert.Equal(t, "Bearer test-key", r.Header.Get("Authorization"))
assert.Equal(t, "application/json", r.Header.Get("Content-Type"))
body, _ := json.Marshal(mockResp)
w.WriteHeader(http.StatusOK)
w.Write(body)
}))
defer server.Close()
p := NewOpenAIProvider(OpenAIProviderConfig{
APIKey: "test-key",
BaseURL: server.URL,
Model: "gpt-4",
MaxRetries: 0,
})
req := ChatRequest{
Messages: []ChatMessage{{Role: "user", Content: "Hello"}},
}
resp, err := p.ChatCompletion(context.Background(), req)
require.NoError(t, err)
assert.Equal(t, "chatcmpl-test", resp.ID)
assert.Equal(t, "Mocked response", resp.Choices[0].Message.Content)
}
func TestOpenAIProvider_CreateEmbedding_MockServer(t *testing.T) {
mockResp := EmbeddingResponse{
Object: "list",
Data: []EmbeddingData{
{
Object: "embedding",
Index: 0,
Embedding: []float64{0.1, 0.2, 0.3},
},
},
Model: "text-embedding-3-small",
Usage: TokenUsage{PromptTokens: 3, TotalTokens: 3},
}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, "POST", r.Method)
assert.Equal(t, "/embeddings", r.URL.Path)
assert.Equal(t, "Bearer test-key", r.Header.Get("Authorization"))
body, _ := json.Marshal(mockResp)
w.WriteHeader(http.StatusOK)
w.Write(body)
}))
defer server.Close()
p := NewOpenAIProvider(OpenAIProviderConfig{
APIKey: "test-key",
BaseURL: server.URL,
MaxRetries: 0,
})
req := EmbeddingRequest{
Input: []string{"Hello world"},
}
resp, err := p.CreateEmbedding(context.Background(), req)
require.NoError(t, err)
assert.Equal(t, 1, len(resp.Data))
assert.InDelta(t, 0.1, resp.Data[0].Embedding[0], 0.001)
}
func TestOpenAIProvider_ChatCompletion_APIError(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusUnauthorized)
w.Write([]byte(`{"error":{"message":"Invalid API key","type":"invalid_request_error","code":"invalid_api_key"}}`))
}))
defer server.Close()
p := NewOpenAIProvider(OpenAIProviderConfig{
APIKey: "bad-key",
BaseURL: server.URL,
MaxRetries: 0,
})
req := ChatRequest{
Messages: []ChatMessage{{Role: "user", Content: "Hello"}},
}
resp, err := p.ChatCompletion(context.Background(), req)
assert.Nil(t, resp)
require.Error(t, err)
var apiErr *APIError
require.True(t, errors.As(err, &apiErr), "error should wrap APIError")
assert.Equal(t, http.StatusUnauthorized, apiErr.StatusCode)
assert.Equal(t, "Invalid API key", apiErr.Message)
assert.Equal(t, "invalid_request_error", apiErr.Type)
}
func TestOpenAIProvider_ChatCompletion_RetryOn5xx(t *testing.T) {
callCount := 0
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
callCount++
if callCount < 3 {
w.WriteHeader(http.StatusInternalServerError)
w.Write([]byte(`{"error":{"message":"Internal server error"}}`))
return
}
mockResp := ChatResponse{
ID: "chatcmpl-retry",
Choices: []ChatChoice{{Message: ChatMessage{Role: "assistant", Content: "Success after retry"}}},
}
body, _ := json.Marshal(mockResp)
w.WriteHeader(http.StatusOK)
w.Write(body)
}))
defer server.Close()
p := NewOpenAIProvider(OpenAIProviderConfig{
APIKey: "test-key",
BaseURL: server.URL,
MaxRetries: 3,
Timeout: 5,
})
req := ChatRequest{
Messages: []ChatMessage{{Role: "user", Content: "Hello"}},
}
resp, err := p.ChatCompletion(context.Background(), req)
require.NoError(t, err)
assert.Equal(t, "Success after retry", resp.Choices[0].Message.Content)
assert.Equal(t, 3, callCount)
}
func TestOpenAIProvider_ChatCompletion_NoRetryOn4xx(t *testing.T) {
callCount := 0
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
callCount++
w.WriteHeader(http.StatusBadRequest)
w.Write([]byte(`{"error":{"message":"Bad request"}}`))
}))
defer server.Close()
p := NewOpenAIProvider(OpenAIProviderConfig{
APIKey: "test-key",
BaseURL: server.URL,
MaxRetries: 3,
Timeout: 5,
})
req := ChatRequest{
Messages: []ChatMessage{{Role: "user", Content: "Hello"}},
}
resp, err := p.ChatCompletion(context.Background(), req)
assert.Nil(t, resp)
require.Error(t, err)
assert.Equal(t, 1, callCount, "should not retry on 400")
}
func TestOpenAIProvider_DefaultModelApplied(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var reqBody ChatRequest
json.NewDecoder(r.Body).Decode(&reqBody)
assert.Equal(t, "gpt-4", reqBody.Model, "default model should be applied when not specified")
mockResp := ChatResponse{ID: "test"}
body, _ := json.Marshal(mockResp)
w.WriteHeader(http.StatusOK)
w.Write(body)
}))
defer server.Close()
p := NewOpenAIProvider(OpenAIProviderConfig{
APIKey: "key",
BaseURL: server.URL,
Model: "gpt-4",
})
req := ChatRequest{
Messages: []ChatMessage{{Role: "user", Content: "Hello"}},
// Model intentionally left empty
}
_, err := p.ChatCompletion(context.Background(), req)
require.NoError(t, err)
}
// --- SSE stream parsing test ---
func TestOpenAIProvider_ParseSSEStream(t *testing.T) {
sseData := `data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":""}]}
data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":""}]}
data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":" world"},"finish_reason":""}]}
data: [DONE]
`
p := NewOpenAIProvider(OpenAIProviderConfig{APIKey: "test"})
var chunks []StreamChunk
err := p.parseSSEStream(strings.NewReader(sseData), func(chunk StreamChunk) error {
chunks = append(chunks, chunk)
return nil
})
require.NoError(t, err)
assert.Equal(t, 3, len(chunks))
assert.Equal(t, "assistant", chunks[0].Choices[0].Delta.Role)
assert.Equal(t, "Hello", chunks[1].Choices[0].Delta.Content)
assert.Equal(t, " world", chunks[2].Choices[0].Delta.Content)
}
// --- Utility tests ---
func TestParseFloatEmbedding(t *testing.T) {
raw := []interface{}{float64(0.1), float32(0.2), int(3), int64(4), "5.5"}
result := ParseFloatEmbedding(raw)
assert.Equal(t, 5, len(result))
assert.InDelta(t, 0.1, result[0], 0.001)
assert.InDelta(t, 0.2, result[1], 0.001)
assert.InDelta(t, 3.0, result[2], 0.001)
assert.InDelta(t, 4.0, result[3], 0.001)
assert.InDelta(t, 5.5, result[4], 0.001)
}
// --- APIError tests ---
func TestAPIError_Error(t *testing.T) {
err := &APIError{
StatusCode: 401,
Message: "Invalid API key",
Type: "invalid_request_error",
Code: "invalid_api_key",
}
assert.Equal(t, "API error (status 401): Invalid API key", err.Error())
}
func TestIsNonRetriableError(t *testing.T) {
// 4xx (except 429) should not retry
assert.True(t, isNonRetriableError(&APIError{StatusCode: 400}))
assert.True(t, isNonRetriableError(&APIError{StatusCode: 401}))
assert.True(t, isNonRetriableError(&APIError{StatusCode: 403}))
assert.True(t, isNonRetriableError(&APIError{StatusCode: 404}))
// 429 (rate limit) should retry
assert.False(t, isNonRetriableError(&APIError{StatusCode: 429}))
// 5xx should retry
assert.False(t, isNonRetriableError(&APIError{StatusCode: 500}))
assert.False(t, isNonRetriableError(&APIError{StatusCode: 502}))
assert.False(t, isNonRetriableError(&APIError{StatusCode: 503}))
// Non-APIError should not be treated as non-retriable
assert.False(t, isNonRetriableError(fmt.Errorf("some error")))
}
+469
View File
@@ -0,0 +1,469 @@
package llm
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"time"
applogger "github.com/gochat/gochat/pkg/logger"
)
// OpenAIProvider implements the Provider interface using OpenAI-compatible APIs.
// Supports custom base_url for domestic providers like Volcengine/Doubao.
type OpenAIProvider struct {
apiKey string
baseURL string
model string
embedModel string
httpClient *http.Client
maxRetries int
}
// OpenAIProviderConfig holds configuration for creating an OpenAIProvider.
type OpenAIProviderConfig struct {
APIKey string
BaseURL string // defaults to "https://api.openai.com/v1"
Model string // defaults to "gpt-4"
EmbedModel string // defaults to "text-embedding-3-small"
MaxRetries int // defaults to 3
Timeout int // HTTP timeout in seconds, defaults to 60
}
// NewOpenAIProvider creates a new OpenAIProvider with the given configuration.
func NewOpenAIProvider(cfg OpenAIProviderConfig) *OpenAIProvider {
if cfg.BaseURL == "" {
cfg.BaseURL = "https://api.openai.com/v1"
}
// Ensure baseURL ends without trailing slash
cfg.BaseURL = strings.TrimRight(cfg.BaseURL, "/")
if cfg.Model == "" {
cfg.Model = "gpt-4"
}
if cfg.EmbedModel == "" {
cfg.EmbedModel = "text-embedding-3-small"
}
if cfg.MaxRetries == 0 {
cfg.MaxRetries = 3
}
if cfg.Timeout == 0 {
cfg.Timeout = 60
}
return &OpenAIProvider{
apiKey: cfg.APIKey,
baseURL: cfg.BaseURL,
model: cfg.Model,
embedModel: cfg.EmbedModel,
httpClient: &http.Client{
Timeout: time.Duration(cfg.Timeout) * time.Second,
},
maxRetries: cfg.MaxRetries,
}
}
// ChatCompletion sends a chat completion request to the OpenAI-compatible API.
func (p *OpenAIProvider) ChatCompletion(ctx context.Context, req ChatRequest) (*ChatResponse, error) {
// Set default model if not specified
if req.Model == "" {
req.Model = p.model
}
body, err := json.Marshal(req)
if err != nil {
applogger.L().Errorf("ChatCompletion: failed to marshal request: %v", err)
return nil, fmt.Errorf("marshal chat request: %w", err)
}
respBody, err := p.doRequestWithRetry(ctx, "/chat/completions", body)
if err != nil {
return nil, fmt.Errorf("chat completion request: %w", err)
}
var resp ChatResponse
if err := json.Unmarshal(respBody, &resp); err != nil {
applogger.L().Errorf("ChatCompletion: failed to unmarshal response: %v", err)
return nil, fmt.Errorf("unmarshal chat response: %w", err)
}
return &resp, nil
}
// ChatCompletionStream sends a streaming chat completion request and returns chunks via callback.
func (p *OpenAIProvider) ChatCompletionStream(ctx context.Context, req ChatRequest, onChunk func(StreamChunk) error) error {
if req.Model == "" {
req.Model = p.model
}
req.Stream = true
body, err := json.Marshal(req)
if err != nil {
applogger.L().Errorf("ChatCompletionStream: failed to marshal request: %v", err)
return fmt.Errorf("marshal chat request: %w", err)
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, p.baseURL+"/chat/completions", bytes.NewReader(body))
if err != nil {
return fmt.Errorf("create stream request: %w", err)
}
p.setHeaders(httpReq)
httpResp, err := p.httpClient.Do(httpReq)
if err != nil {
return fmt.Errorf("stream request: %w", err)
}
defer httpResp.Body.Close()
if httpResp.StatusCode != http.StatusOK {
respBody, _ := io.ReadAll(httpResp.Body)
applogger.L().Errorf("ChatCompletionStream: unexpected status %d: %s", httpResp.StatusCode, string(respBody))
return fmt.Errorf("stream request status %d: %s", httpResp.StatusCode, string(respBody))
}
return p.parseSSEStream(httpResp.Body, onChunk)
}
// CreateEmbedding sends an embedding request to the OpenAI-compatible API.
func (p *OpenAIProvider) CreateEmbedding(ctx context.Context, req EmbeddingRequest) (*EmbeddingResponse, error) {
if req.Model == "" {
req.Model = p.embedModel
}
body, err := json.Marshal(req)
if err != nil {
applogger.L().Errorf("CreateEmbedding: failed to marshal request: %v", err)
return nil, fmt.Errorf("marshal embedding request: %w", err)
}
respBody, err := p.doRequestWithRetry(ctx, "/embeddings", body)
if err != nil {
return nil, fmt.Errorf("embedding request: %w", err)
}
var resp EmbeddingResponse
if err := json.Unmarshal(respBody, &resp); err != nil {
applogger.L().Errorf("CreateEmbedding: failed to unmarshal response: %v", err)
return nil, fmt.Errorf("unmarshal embedding response: %w", err)
}
return &resp, nil
}
// doRequestWithRetry performs an HTTP request with retry logic.
func (p *OpenAIProvider) doRequestWithRetry(ctx context.Context, path string, body []byte) ([]byte, error) {
var lastErr error
for attempt := 0; attempt <= p.maxRetries; attempt++ {
if attempt > 0 {
// Exponential backoff: 1s, 2s, 4s
backoff := time.Duration(1<<uint(attempt-1)) * time.Second
applogger.L().Infof("Retrying request (attempt %d/%d) after %v: %v", attempt, p.maxRetries, backoff, lastErr)
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-time.After(backoff):
}
}
respBody, err := p.doRequest(ctx, path, body)
if err == nil {
return respBody, nil
}
// Don't retry on client errors (4xx except 429)
if isNonRetriableError(err) {
return nil, err
}
lastErr = err
}
return nil, fmt.Errorf("max retries (%d) exceeded: %w", p.maxRetries, lastErr)
}
// doRequest performs a single HTTP request to the API.
func (p *OpenAIProvider) doRequest(ctx context.Context, path string, body []byte) ([]byte, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.baseURL+path, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("create request: %w", err)
}
p.setHeaders(req)
resp, err := p.httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf("execute request: %w", err)
}
defer resp.Body.Close()
respBody, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("read response body: %w", err)
}
if resp.StatusCode >= 400 {
apiErr := parseAPIError(resp.StatusCode, respBody)
applogger.L().Errorf("API error (status %d): %v", resp.StatusCode, apiErr)
return nil, apiErr
}
return respBody, nil
}
// setHeaders sets common headers for API requests.
func (p *OpenAIProvider) setHeaders(req *http.Request) {
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+p.apiKey)
}
// parseSSEStream parses Server-Sent Events from a streaming response body.
func (p *OpenAIProvider) parseSSEStream(body io.Reader, onChunk func(StreamChunk) error) error {
reader := newSSEReader(body)
for {
event, err := reader.Next()
if err != nil {
return fmt.Errorf("read SSE event: %w", err)
}
if event == nil {
// Stream ended
return nil
}
// Skip non-data events
if event.Type != "message" || event.Data == "" {
continue
}
// OpenAI sends "[DONE]" to signal stream end
if event.Data == "[DONE]" {
return nil
}
var chunk StreamChunk
if err := json.Unmarshal([]byte(event.Data), &chunk); err != nil {
applogger.L().Errorf("parseSSEStream: failed to unmarshal chunk: %v (data: %s)", err, event.Data)
continue
}
if err := onChunk(chunk); err != nil {
return fmt.Errorf("chunk callback: %w", err)
}
}
}
// --- Error types ---
// APIError represents an error returned by the OpenAI-compatible API.
type APIError struct {
StatusCode int
Message string
Type string
Code string
}
func (e *APIError) Error() string {
return fmt.Sprintf("API error (status %d): %s", e.StatusCode, e.Message)
}
// parseAPIError creates an APIError from HTTP status and response body.
func parseAPIError(statusCode int, body []byte) *APIError {
apiErr := &APIError{
StatusCode: statusCode,
Message: string(body),
}
// Try to parse OpenAI error structure
var errResp struct {
Error struct {
Message string `json:"message"`
Type string `json:"type"`
Code string `json:"code"`
} `json:"error"`
}
if err := json.Unmarshal(body, &errResp); err == nil && errResp.Error.Message != "" {
apiErr.Message = errResp.Error.Message
apiErr.Type = errResp.Error.Type
apiErr.Code = errResp.Error.Code
}
return apiErr
}
// isNonRetriableError returns true for errors that should not be retried.
func isNonRetriableError(err error) bool {
if apiErr, ok := err.(*APIError); ok {
// Retry on rate limit (429) and server errors (5xx)
// Don't retry on client errors (400, 401, 403, 404, etc.)
return apiErr.StatusCode >= 400 && apiErr.StatusCode < 500 && apiErr.StatusCode != 429
}
return false
}
// --- SSE Reader ---
// sseEvent represents a parsed SSE event.
type sseEvent struct {
Type string // event type (default "message" if not specified)
Data string // data payload
ID string // event ID
}
// sseReader reads Server-Sent Events from a stream.
type sseReader struct {
scanner *sseLineScanner
}
func newSSEReader(body io.Reader) *sseReader {
return &sseReader{
scanner: newSSELineScanner(body),
}
}
// Next reads the next SSE event from the stream.
// Returns nil when the stream is complete.
func (r *sseReader) Next() (*sseEvent, error) {
var event *sseEvent
for {
line, err := r.scanner.Next()
if err != nil {
return nil, err
}
if line == nil {
// Stream ended
return event, nil
}
text := *line
if text == "" {
// Empty line = event boundary, dispatch current event
if event != nil {
return event, nil
}
continue
}
if strings.HasPrefix(text, ":") {
// Comment, skip
continue
}
field, value := parseSSEField(text)
switch field {
case "event":
if event == nil {
event = &sseEvent{}
}
event.Type = value
case "data":
if event == nil {
event = &sseEvent{Type: "message"}
}
if event.Data != "" {
event.Data += "\n"
}
event.Data += value
case "id":
if event == nil {
event = &sseEvent{}
}
event.ID = value
}
}
}
func parseSSEField(line string) (field, value string) {
idx := strings.Index(line, ":")
if idx == -1 {
return line, ""
}
field = line[:idx]
value = strings.TrimLeft(line[idx+1:], " ")
return field, value
}
// sseLineScanner reads lines from an SSE stream efficiently.
type sseLineScanner struct {
reader io.Reader
buffer []byte
hasData bool
}
func newSSELineScanner(reader io.Reader) *sseLineScanner {
return &sseLineScanner{
reader: reader,
buffer: make([]byte, 0, 4096),
}
}
// Next returns the next line from the stream.
// Returns nil when the stream is complete.
func (s *sseLineScanner) Next() (*string, error) {
for {
// Check if we have a complete line in the buffer
idx := bytes.IndexByte(s.buffer, '\n')
if idx != -1 {
line := string(s.buffer[:idx])
s.buffer = s.buffer[idx+1:]
// Strip \r if present (CRLF)
line = strings.TrimRight(line, "\r")
return &line, nil
}
// Read more data
tmp := make([]byte, 4096)
n, err := s.reader.Read(tmp)
if n > 0 {
s.buffer = append(s.buffer, tmp[:n]...)
s.hasData = true
}
if err != nil {
if err == io.EOF {
if len(s.buffer) > 0 {
line := string(s.buffer)
s.buffer = s.buffer[:0]
line = strings.TrimRight(line, "\r")
return &line, nil
}
return nil, nil
}
return nil, err
}
}
}
// --- Utility functions ---
// ParseFloatEmbedding converts a slice of any (JSON numbers) to []float64.
// Useful when embedding responses contain mixed numeric types.
func ParseFloatEmbedding(raw []interface{}) []float64 {
result := make([]float64, len(raw))
for i, v := range raw {
switch n := v.(type) {
case float64:
result[i] = n
case float32:
result[i] = float64(n)
case int:
result[i] = float64(n)
case int64:
result[i] = float64(n)
case string:
f, err := strconv.ParseFloat(n, 64)
if err == nil {
result[i] = f
}
}
}
return result
}
+127
View File
@@ -0,0 +1,127 @@
package llm
import "context"
// Provider defines the interface for LLM operations.
// Inspired by Chatwoot's CaptainAI which uses OpenAI API for:
// 1. Chat completion (copilot message suggestions)
// 2. Embedding generation (knowledge base document vectorization)
// 3. RAG Q&A (retrieve relevant documents via embedding then generate answers)
// 4. Streaming chat completion (SSE streaming for captain task endpoints)
type Provider interface {
// ChatCompletion sends a chat completion request and returns the response.
ChatCompletion(ctx context.Context, req ChatRequest) (*ChatResponse, error)
// CreateEmbedding generates embeddings for the given input texts.
CreateEmbedding(ctx context.Context, req EmbeddingRequest) (*EmbeddingResponse, error)
// ChatCompletionStream sends a streaming chat completion request.
// Chunks are delivered via the onChunk callback. Return nil when the
// stream completes naturally, or an error on failure. The caller is
// responsible for flushing/closing the downstream SSE connection.
ChatCompletionStream(ctx context.Context, req ChatRequest, onChunk func(StreamChunk) error) error
}
// StreamingProvider is a narrowed interface for contexts that only need
// streaming capability (e.g. SSE handler wiring in bootstrap). Providers
// that implement Provider automatically satisfy StreamingProvider since
// ChatCompletionStream is on the base interface.
type StreamingProvider interface {
ChatCompletionStream(ctx context.Context, req ChatRequest, onChunk func(StreamChunk) error) error
}
// ChatRequest represents a request to the chat completion API.
type ChatRequest struct {
Model string `json:"model"`
Messages []ChatMessage `json:"messages"`
Temperature float64 `json:"temperature,omitempty"`
MaxTokens int `json:"max_tokens,omitempty"`
Tools []ToolDefinition `json:"tools,omitempty"`
Stream bool `json:"stream,omitempty"`
}
// ChatMessage represents a single message in a chat conversation.
type ChatMessage struct {
Role string `json:"role"`
Content string `json:"content"`
}
// ToolDefinition represents a tool that the model can call.
type ToolDefinition struct {
Type string `json:"type"`
Function ToolFunction `json:"function"`
}
// ToolFunction describes a function tool definition.
type ToolFunction struct {
Name string `json:"name"`
Description string `json:"description"`
Parameters map[string]interface{} `json:"parameters"`
}
// ChatResponse represents the response from a chat completion API.
type ChatResponse struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
Model string `json:"model"`
Choices []ChatChoice `json:"choices"`
Usage TokenUsage `json:"usage"`
}
// ChatChoice represents a single choice in a chat completion response.
type ChatChoice struct {
Index int `json:"index"`
Message ChatMessage `json:"message"`
FinishReason string `json:"finish_reason"`
}
// TokenUsage represents token usage statistics.
type TokenUsage struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
}
// EmbeddingRequest represents a request to the embeddings API.
type EmbeddingRequest struct {
Model string `json:"model"`
Input []string `json:"input"`
}
// EmbeddingResponse represents the response from an embeddings API.
type EmbeddingResponse struct {
Object string `json:"object"`
Data []EmbeddingData `json:"data"`
Model string `json:"model"`
Usage TokenUsage `json:"usage"`
}
// EmbeddingData represents a single embedding result.
type EmbeddingData struct {
Object string `json:"object"`
Index int `json:"index"`
Embedding []float64 `json:"embedding"`
}
// StreamChunk represents a single chunk in a streaming response.
type StreamChunk struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
Model string `json:"model"`
Choices []StreamChoice `json:"choices"`
}
// StreamChoice represents a single choice in a streaming chunk.
type StreamChoice struct {
Index int `json:"index"`
Delta StreamDelta `json:"delta"`
FinishReason string `json:"finish_reason"`
}
// StreamDelta represents the delta content in a streaming chunk.
type StreamDelta struct {
Role string `json:"role,omitempty"`
Content string `json:"content,omitempty"`
}