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:
@@ -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")))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
Reference in New Issue
Block a user