定义并实现三个 AI 服务客户端,每个服务一个独立包 #36

Merged
huanghaosheng merged 10 commits from feature/phase3 into develop 2026-06-13 15:53:10 +08:00
12 changed files with 1374 additions and 6 deletions

View File

@@ -6,6 +6,7 @@ require (
github.com/gin-gonic/gin v1.10.0
github.com/google/uuid v1.6.0
github.com/gorilla/websocket v1.5.3
github.com/redis/go-redis/v9 v9.20.1
github.com/spf13/viper v1.21.0
go.uber.org/zap v1.28.0
)
@@ -31,7 +32,6 @@ require (
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
github.com/modern-go/reflect2 v1.0.2 // indirect
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
github.com/redis/go-redis/v9 v9.20.1 // indirect
github.com/sagikazarmark/locafero v0.11.0 // indirect
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 // indirect
github.com/spf13/afero v1.15.0 // indirect

View File

@@ -1,3 +1,7 @@
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
github.com/bytedance/sonic v1.11.6 h1:oUp34TzMlL+OY1OUWxHqsdkgC/Zfc85zGqw9siXjrc0=
github.com/bytedance/sonic v1.11.6/go.mod h1:LysEHSvpvDySVdC2f87zGWf6CIKJcAvqab1ZaiQtds4=
github.com/bytedance/sonic/loader v0.1.1 h1:c+e5Pt1k/cy5wMveRDyk2X4B9hF4g7an8N3zCYjJFNM=
@@ -43,8 +47,6 @@ github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/ad
github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM=
github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo=
github.com/klauspost/cpuid/v2 v2.0.9/go.mod h1:FInQzS24/EEf25PyTYn52gqo7WaD8xa0213Md/qVLRg=
github.com/klauspost/cpuid/v2 v2.2.7 h1:ZWSB3igEs+d0qvnxR/ZBzXVmxkgt8DdzP6m9pfuVLDM=
github.com/klauspost/cpuid/v2 v2.2.7/go.mod h1:Lcz8mBdAVJIBVzewtcLocK12l3Y+JytZYpaMropDUws=
github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE=
github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
github.com/knz/go-libedit v1.10.1/go.mod h1:MZTVkCWyz0oBc7JOWP3wNAzd002ZbM/5hgShxwh4x8M=
@@ -97,6 +99,8 @@ github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS
github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08=
github.com/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65EE=
github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg=
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
@@ -114,10 +118,7 @@ golang.org/x/crypto v0.23.0 h1:dIJU/v2J8Mdglj/8rJ6UUOM3Zc9zLZxVZwwxMooUSAI=
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
golang.org/x/net v0.25.0 h1:d/OCCoBEUq33pjydKrGQhw7IlUPI2Oylr+8qLx49kac=
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU=
golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/sys v0.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc=
golang.org/x/sys v0.30.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/text v0.28.0 h1:rhazDwis8INMIwQ4tpjLDzUhx6RlXqZNPEM0huQojng=

View File

@@ -0,0 +1,37 @@
package llm
import (
"context"
"github.com/hhs/camtalk/internal/models"
)
// Service 多模态大模型服务契约。
type Service interface {
// ChatStream 流式推理,返回增量文本的 channel。
// 调用方必须消费 channel 直到 Done=true否则需 cancel ctx 以释放连接。
ChatStream(ctx context.Context, req Request) (<-chan Chunk, error)
}
// Request 推理请求。
type Request struct {
Image []byte // JPEG 图片(已从 Base64 解码)
Text string // 用户语音识别后的文本
History []models.Message // 最近 N 轮对话历史
Language string // 语言,如 "zh-CN"
}
// Chunk 流式推理的一个增量片段。
type Chunk struct {
Delta string // 增量文本
Done bool // 是否结束
TokensUsed *TokenUsage // 仅 Done=true 时有值
Model string // 实际使用的模型名
}
// TokenUsage 用量统计。
type TokenUsage struct {
Prompt int
Completion int
Total int
}

View File

@@ -0,0 +1,237 @@
package llm
import (
"bufio"
"bytes"
"context"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
"go.uber.org/zap"
)
// OpenAIService 基于 OpenAI Chat Completions API 的 LLM 实现。
type OpenAIService struct {
apiKey string
model string
endpoint string
timeout time.Duration
logger *zap.SugaredLogger
client *http.Client
}
// NewOpenAIService 创建 OpenAI LLM 服务。
func NewOpenAIService(apiKey, model, endpoint string, timeoutSec int, logger *zap.SugaredLogger) *OpenAIService {
if model == "" {
model = "gpt-4o"
}
if endpoint == "" {
endpoint = "https://api.openai.com/v1"
}
timeout := time.Duration(timeoutSec) * time.Second
if timeout <= 0 {
timeout = 10 * time.Second
}
return &OpenAIService{
apiKey: apiKey,
model: model,
endpoint: endpoint,
timeout: timeout,
logger: logger,
client: &http.Client{Timeout: 60 * time.Second}, // HTTP client timeout > LLM timeout
}
}
// --- OpenAI API 请求/响应结构 ---
type chatRequest struct {
Model string `json:"model"`
Messages []chatMessage `json:"messages"`
Stream bool `json:"stream"`
}
type chatMessage struct {
Role string `json:"role"`
Content []contentPart `json:"content"`
}
type contentPart struct {
Type string `json:"type"`
Text string `json:"text,omitempty"`
ImageURL *imageURL `json:"image_url,omitempty"`
}
type imageURL struct {
URL string `json:"url"`
}
// streamDelta SSE 流式响应的单个 delta。
type streamDelta struct {
Choices []struct {
Delta struct {
Content string `json:"content"`
} `json:"delta"`
FinishReason *string `json:"finish_reason"`
} `json:"choices"`
Usage *struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
Model string `json:"model"`
}
// ChatStream 实现 llm.Service。调用 OpenAI Chat Completions API 流式推理。
func (o *OpenAIService) ChatStream(ctx context.Context, req Request) (<-chan Chunk, error) {
// 构建请求
messages := o.buildMessages(req)
body := chatRequest{
Model: o.model,
Messages: messages,
Stream: true,
}
payload, err := json.Marshal(body)
if err != nil {
return nil, fmt.Errorf("llm: marshal request: %w", err)
}
// 创建带超时的 context
ctx, cancel := context.WithTimeout(ctx, o.timeout)
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, o.endpoint+"/chat/completions", bytes.NewReader(payload))
if err != nil {
cancel()
return nil, fmt.Errorf("llm: create request: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
httpReq.Header.Set("Authorization", "Bearer "+o.apiKey)
resp, err := o.client.Do(httpReq)
if err != nil {
cancel()
return nil, fmt.Errorf("llm: send request: %w", err)
}
if resp.StatusCode != http.StatusOK {
cancel()
bodyBytes, _ := io.ReadAll(resp.Body)
resp.Body.Close()
return nil, fmt.Errorf("llm: api error (status %d): %s", resp.StatusCode, string(bodyBytes))
}
// 启动 goroutine 解析 SSE 流
ch := make(chan Chunk, 64)
go func() {
defer close(ch)
defer cancel()
defer resp.Body.Close()
o.parseSSEStream(resp.Body, ch)
}()
return ch, nil
}
// parseSSEStream 解析 SSE 流,将 delta 发送到 channel。
func (o *OpenAIService) parseSSEStream(body io.Reader, ch chan<- Chunk) {
scanner := bufio.NewScanner(body)
scanner.Buffer(make([]byte, 0, 64*1024), 256*1024)
var fullText strings.Builder
var lastModel string
for scanner.Scan() {
line := scanner.Text()
// SSE 格式data: {...}
if !strings.HasPrefix(line, "data: ") {
continue
}
data := strings.TrimPrefix(line, "data: ")
if data == "[DONE]" {
// 流结束,发送最终 chunk
ch <- Chunk{Delta: "", Done: true, Model: lastModel}
return
}
var delta streamDelta
if err := json.Unmarshal([]byte(data), &delta); err != nil {
o.logger.Warnw("llm: unmarshal delta failed", "error", err, "data", data)
continue
}
if delta.Model != "" {
lastModel = delta.Model
}
// 提取增量文本
if len(delta.Choices) > 0 {
content := delta.Choices[0].Delta.Content
if content != "" {
fullText.WriteString(content)
ch <- Chunk{Delta: content, Done: false, Model: lastModel}
}
// 某些模型在最后一个 choice 中携带 usage
if delta.Choices[0].FinishReason != nil && delta.Usage != nil {
ch <- Chunk{
Delta: "",
Done: true,
Model: lastModel,
TokensUsed: &TokenUsage{
Prompt: delta.Usage.PromptTokens,
Completion: delta.Usage.CompletionTokens,
Total: delta.Usage.TotalTokens,
},
}
return
}
}
}
// scanner 结束但没收到 [DONE]
if err := scanner.Err(); err != nil {
o.logger.Warnw("llm: scan error", "error", err)
}
ch <- Chunk{Delta: "", Done: true, Model: lastModel}
}
// buildMessages 构建 OpenAI Chat API 的 messages 数组。
func (o *OpenAIService) buildMessages(req Request) []chatMessage {
var messages []chatMessage
// System prompt
messages = append(messages, chatMessage{
Role: "system",
Content: []contentPart{{Type: "text", Text: BuildSystemPrompt(req.Language, "")}},
})
// 历史消息
for _, msg := range req.History {
messages = append(messages, chatMessage{
Role: msg.Role,
Content: []contentPart{{Type: "text", Text: msg.Content}},
})
}
// 当前用户消息(图像 + 文本)
var parts []contentPart
if len(req.Image) > 0 {
b64 := base64.StdEncoding.EncodeToString(req.Image)
parts = append(parts, contentPart{
Type: "image_url",
ImageURL: &imageURL{URL: "data:image/jpeg;base64," + b64},
})
}
parts = append(parts, contentPart{Type: "text", Text: req.Text})
messages = append(messages, chatMessage{Role: "user", Content: parts})
return messages
}

View File

@@ -0,0 +1,251 @@
package llm
import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"go.uber.org/zap"
"github.com/hhs/camtalk/internal/models"
)
// mockLLMServer 创建模拟 OpenAI SSE 流式响应的 HTTP 服务器。
func mockLLMServer(t *testing.T, handler http.HandlerFunc) *httptest.Server {
t.Helper()
return httptest.NewServer(handler)
}
func TestOpenAIService_ChatStream_Success(t *testing.T) {
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
// 验证请求
if r.Method != http.MethodPost {
t.Errorf("method = %s, want POST", r.Method)
}
if !strings.Contains(r.URL.Path, "/chat/completions") {
t.Errorf("path = %s, should contain /chat/completions", r.URL.Path)
}
auth := r.Header.Get("Authorization")
if auth != "Bearer test-key" {
t.Errorf("Authorization = %q, want %q", auth, "Bearer test-key")
}
w.Header().Set("Content-Type", "text/event-stream")
flusher, ok := w.(http.Flusher)
if !ok {
t.Fatal("ResponseWriter does not support Flusher")
}
// 发送几个 delta
deltas := []string{"你好", "世界", ""}
for _, d := range deltas {
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"%s\"}}],\"model\":\"gpt-4o\"}\n\n", d)
flusher.Flush()
}
// 发送 [DONE]
fmt.Fprintf(w, "data: [DONE]\n\n")
flusher.Flush()
})
defer srv.Close()
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 10, zap.NewNop().Sugar())
ch, err := svc.ChatStream(context.Background(), Request{
Text: "这是什么?",
Language: "zh-CN",
})
if err != nil {
t.Fatalf("ChatStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 应该有 3 个文本 chunk + 1 个 Done chunk
if len(chunks) != 4 {
t.Fatalf("got %d chunks, want 4", len(chunks))
}
// 验证文本内容
if chunks[0].Delta != "你好" {
t.Errorf("chunk[0].Delta = %q, want %q", chunks[0].Delta, "你好")
}
if chunks[1].Delta != "世界" {
t.Errorf("chunk[1].Delta = %q, want %q", chunks[1].Delta, "世界")
}
// 验证最后一个 chunk 是 Done
last := chunks[len(chunks)-1]
if !last.Done {
t.Error("last chunk should be Done")
}
if last.Model != "gpt-4o" {
t.Errorf("last chunk Model = %q, want %q", last.Model, "gpt-4o")
}
}
func TestOpenAIService_ChatStream_WithImage(t *testing.T) {
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}],\"model\":\"gpt-4o\"}\n\n")
fmt.Fprintf(w, "data: [DONE]\n\n")
})
defer srv.Close()
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 10, zap.NewNop().Sugar())
ch, err := svc.ChatStream(context.Background(), Request{
Image: []byte("fake-jpeg-data"),
Text: "描述图片",
Language: "zh-CN",
})
if err != nil {
t.Fatalf("ChatStream() error: %v", err)
}
// 消费 channel
for range ch {
}
}
func TestOpenAIService_ChatStream_WithHistory(t *testing.T) {
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}],\"model\":\"gpt-4o\"}\n\n")
fmt.Fprintf(w, "data: [DONE]\n\n")
})
defer srv.Close()
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 10, zap.NewNop().Sugar())
ch, err := svc.ChatStream(context.Background(), Request{
Text: "继续",
Language: "zh-CN",
History: []models.Message{
{Role: "user", Content: "你好"},
{Role: "assistant", Content: "你好!有什么可以帮助你的吗?"},
},
})
if err != nil {
t.Fatalf("ChatStream() error: %v", err)
}
for range ch {
}
}
func TestOpenAIService_ChatStream_APIError(t *testing.T) {
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusUnauthorized)
fmt.Fprintf(w, `{"error":{"message":"Invalid API key"}}`)
})
defer srv.Close()
svc := NewOpenAIService("bad-key", "gpt-4o", srv.URL, 10, zap.NewNop().Sugar())
_, err := svc.ChatStream(context.Background(), Request{
Text: "test",
})
if err == nil {
t.Fatal("ChatStream() should return error for 401")
}
if !strings.Contains(err.Error(), "401") {
t.Errorf("error should mention 401, got: %v", err)
}
}
func TestOpenAIService_ChatStream_Timeout(t *testing.T) {
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
// 模拟慢响应
time.Sleep(5 * time.Second)
w.Header().Set("Content-Type", "text/event-stream")
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"late\"}}]}\n\n")
fmt.Fprintf(w, "data: [DONE]\n\n")
})
defer srv.Close()
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 1, zap.NewNop().Sugar()) // 1s timeout
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
ch, err := svc.ChatStream(ctx, Request{Text: "test"})
if err != nil {
// 超时可能在建立连接时或读取时发生
return
}
// 如果连接成功,消费 channel 应该超时
var gotContent bool
for c := range ch {
if c.Delta != "" {
gotContent = true
}
}
if gotContent {
t.Error("should not receive content before timeout")
}
}
func TestOpenAIService_ChatStream_UsageInResponse(t *testing.T) {
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
// 带 usage 的最后一个 chunk
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"hi\"},\"finish_reason\":\"stop\"}],\"model\":\"gpt-4o\",\"usage\":{\"prompt_tokens\":10,\"completion_tokens\":5,\"total_tokens\":15}}\n\n")
fmt.Fprintf(w, "data: [DONE]\n\n")
})
defer srv.Close()
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 10, zap.NewNop().Sugar())
ch, err := svc.ChatStream(context.Background(), Request{Text: "test"})
if err != nil {
t.Fatalf("ChatStream() error: %v", err)
}
var last Chunk
for c := range ch {
last = c
}
if !last.Done {
t.Error("last chunk should be Done")
}
if last.TokensUsed == nil {
t.Fatal("last chunk should have TokensUsed")
}
if last.TokensUsed.Total != 15 {
t.Errorf("TokensUsed.Total = %d, want 15", last.TokensUsed.Total)
}
}
func TestBuildSystemPrompt(t *testing.T) {
tests := []struct {
name string
language string
detailLevel string
wantContain string
}{
{"chinese default", "zh-CN", "", "视觉助手"},
{"chinese high", "zh-CN", "high", "更详细"},
{"english default", "en", "", "visual assistant"},
{"english high", "en", "high", "detailed"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := BuildSystemPrompt(tt.language, tt.detailLevel)
if !strings.Contains(got, tt.wantContain) {
t.Errorf("BuildSystemPrompt(%q, %q) should contain %q", tt.language, tt.detailLevel, tt.wantContain)
}
})
}
}

View File

@@ -0,0 +1,25 @@
package llm
import "strings"
// BuildSystemPrompt 根据语言和细节级别构建系统提示词。
func BuildSystemPrompt(language, detailLevel string) string {
isChinese := strings.HasPrefix(language, "zh")
var prompt strings.Builder
if isChinese {
prompt.WriteString("你是一个视觉助手。用户通过摄像头看到一个场景,并用语音向你提问。请用简洁自然的中文回答。如果涉及视觉描述,先说\"我看到……\"。回答控制在3-5句话以内除非用户要求详细说明。")
} else {
prompt.WriteString("You are a visual assistant. The user sees a scene through their camera and asks questions by voice. Answer concisely and naturally. If describing visual content, start with 'I see...'. Keep answers to 3-5 sentences unless the user asks for detail.")
}
if detailLevel == "high" {
if isChinese {
prompt.WriteString("请提供更详细的视觉描述,包括颜色、位置、数量等细节。")
} else {
prompt.WriteString(" Provide detailed visual descriptions including colors, positions, quantities, and other details.")
}
}
return prompt.String()
}

View File

@@ -0,0 +1,136 @@
package stt
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/url"
"strings"
"time"
"github.com/gorilla/websocket"
"go.uber.org/zap"
)
// DeepgramService 基于 Deepgram WebSocket API 的语音识别实现。
type DeepgramService struct {
apiKey string
endpoint string
logger *zap.SugaredLogger
}
// NewDeepgramService 创建 Deepgram STT 服务。
func NewDeepgramService(apiKey, endpoint string, logger *zap.SugaredLogger) *DeepgramService {
if endpoint == "" {
endpoint = "wss://api.deepgram.com/v1/listen"
}
return &DeepgramService{
apiKey: apiKey,
endpoint: endpoint,
logger: logger,
}
}
// deepgramResponse Deepgram WebSocket 响应。
type deepgramResponse struct {
Channel struct {
Alternatives []struct {
Transcript string `json:"transcript"`
Confidence float64 `json:"confidence"`
} `json:"alternatives"`
} `json:"channel"`
IsFinal bool `json:"is_final"`
}
// Recognize 实现 stt.Service。通过 WebSocket 发送音频到 Deepgram返回最终识别文本。
func (d *DeepgramService) Recognize(ctx context.Context, audio []byte, opts Options) (string, error) {
if len(audio) == 0 {
return "", fmt.Errorf("stt: empty audio")
}
// 构建 WebSocket URL附带查询参数
wsURL := d.buildURL(opts)
// 5 秒总超时
ctx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
// 建立 WebSocket 连接
conn, _, err := websocket.DefaultDialer.DialContext(ctx, wsURL, http.Header{
"Authorization": []string{"Token " + d.apiKey},
})
if err != nil {
return "", fmt.Errorf("stt: connect deepgram: %w", err)
}
defer conn.Close()
// 发送音频数据(一次性)
if err := conn.WriteMessage(websocket.BinaryMessage, audio); err != nil {
return "", fmt.Errorf("stt: send audio: %w", err)
}
// 发送 Close 消息通知服务端音频已发送完毕
closeMsg := websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")
_ = conn.WriteMessage(websocket.CloseMessage, closeMsg)
// 读取识别结果
var transcript strings.Builder
for {
_, message, err := conn.ReadMessage()
if err != nil {
// Close 帧是正常的结束信号
if websocket.IsCloseError(err, websocket.CloseNormalClosure) {
break
}
if websocket.IsUnexpectedCloseError(err, websocket.CloseNormalClosure) {
break
}
return "", fmt.Errorf("stt: read response: %w", err)
}
var resp deepgramResponse
if err := json.Unmarshal(message, &resp); err != nil {
d.logger.Warnw("stt: unmarshal response failed", "error", err)
continue
}
// 只累积 final 结果,跳过中间结果
if resp.IsFinal && len(resp.Channel.Alternatives) > 0 {
text := strings.TrimSpace(resp.Channel.Alternatives[0].Transcript)
if text != "" {
transcript.WriteString(text)
}
}
}
return strings.TrimSpace(transcript.String()), nil
}
// buildURL 构建 Deepgram WebSocket URL包含音频格式参数。
func (d *DeepgramService) buildURL(opts Options) string {
u, _ := url.Parse(d.endpoint)
encoding := opts.Encoding
if encoding == "" {
encoding = "pcm_s16le"
}
sampleRate := opts.SampleRate
if sampleRate == 0 {
sampleRate = 16000
}
language := opts.Language
if language == "" {
language = "zh-CN"
}
q := u.Query()
q.Set("encoding", encoding)
q.Set("sample_rate", fmt.Sprintf("%d", sampleRate))
q.Set("language", language)
q.Set("model", "nova-2")
q.Set("punctuate", "true")
u.RawQuery = q.Encode()
return u.String()
}

View File

@@ -0,0 +1,191 @@
package stt
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/gorilla/websocket"
"go.uber.org/zap"
)
var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true },
}
// newMockDeepgram 创建模拟 Deepgram WebSocket 服务。
// 返回 httptest.Server 和对应的 ws:// URL。
func newMockDeepgram(t *testing.T, handler func(conn *websocket.Conn)) *httptest.Server {
t.Helper()
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
t.Logf("upgrade error: %v", err)
return
}
defer conn.Close()
handler(conn)
}))
return srv
}
// wsToWss 将 http:// 转换为 ws://。
func wsToWss(httpURL string) string {
return "ws" + strings.TrimPrefix(httpURL, "http")
}
func TestDeepgramService_Recognize_Success(t *testing.T) {
srv := newMockDeepgram(t, func(conn *websocket.Conn) {
// 读取音频数据
_, _, err := conn.ReadMessage()
if err != nil {
t.Errorf("read audio: %v", err)
return
}
// 发送中间结果(非 final
intermediate := deepgramResponse{
IsFinal: false,
}
intermediate.Channel.Alternatives = []struct {
Transcript string `json:"transcript"`
Confidence float64 `json:"confidence"`
}{{Transcript: "你好", Confidence: 0.9}}
data, _ := json.Marshal(intermediate)
_ = conn.WriteMessage(websocket.TextMessage, data)
// 发送最终结果
final := deepgramResponse{
IsFinal: true,
}
final.Channel.Alternatives = []struct {
Transcript string `json:"transcript"`
Confidence float64 `json:"confidence"`
}{{Transcript: "你好世界", Confidence: 0.95}}
data, _ = json.Marshal(final)
_ = conn.WriteMessage(websocket.TextMessage, data)
// 等待客户端关闭
_, _, _ = conn.ReadMessage()
})
defer srv.Close()
svc := NewDeepgramService("test-key", wsToWss(srv.URL)+"/v1/listen", zap.NewNop().Sugar())
text, err := svc.Recognize(context.Background(), []byte("fake-pcm-audio"), Options{
Encoding: "pcm_s16le",
SampleRate: 16000,
Language: "zh-CN",
})
if err != nil {
t.Fatalf("Recognize() error: %v", err)
}
if text != "你好世界" {
t.Errorf("Recognize() = %q, want %q", text, "你好世界")
}
}
func TestDeepgramService_Recognize_EmptyAudio(t *testing.T) {
svc := NewDeepgramService("test-key", "ws://localhost", zap.NewNop().Sugar())
_, err := svc.Recognize(context.Background(), nil, Options{})
if err == nil {
t.Fatal("Recognize() with empty audio should return error")
}
}
func TestDeepgramService_Recognize_ConnectError(t *testing.T) {
svc := NewDeepgramService("test-key", "ws://localhost:1", zap.NewNop().Sugar())
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
_, err := svc.Recognize(ctx, []byte("audio"), Options{})
if err == nil {
t.Fatal("Recognize() with bad endpoint should return error")
}
}
func TestDeepgramService_Recognize_Timeout(t *testing.T) {
// 模拟一个永不响应的服务端
srv := newMockDeepgram(t, func(conn *websocket.Conn) {
// 读取音频但不发送任何结果,让客户端超时
_, _, _ = conn.ReadMessage()
time.Sleep(10 * time.Second)
})
defer srv.Close()
svc := NewDeepgramService("test-key", wsToWss(srv.URL)+"/v1/listen", zap.NewNop().Sugar())
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Second)
defer cancel()
_, err := svc.Recognize(ctx, []byte("audio"), Options{})
if err == nil {
t.Fatal("Recognize() should timeout")
}
}
func TestDeepgramService_Recognize_MultipleFinals(t *testing.T) {
srv := newMockDeepgram(t, func(conn *websocket.Conn) {
_, _, _ = conn.ReadMessage()
// 发送多个 final 结果(多句话场景)
for _, text := range []string{"你好", "世界"} {
resp := deepgramResponse{IsFinal: true}
resp.Channel.Alternatives = []struct {
Transcript string `json:"transcript"`
Confidence float64 `json:"confidence"`
}{{Transcript: text, Confidence: 0.9}}
data, _ := json.Marshal(resp)
_ = conn.WriteMessage(websocket.TextMessage, data)
}
_, _, _ = conn.ReadMessage()
})
defer srv.Close()
svc := NewDeepgramService("test-key", wsToWss(srv.URL)+"/v1/listen", zap.NewNop().Sugar())
text, err := svc.Recognize(context.Background(), []byte("audio"), Options{})
if err != nil {
t.Fatalf("Recognize() error: %v", err)
}
if text != "你好世界" {
t.Errorf("Recognize() = %q, want %q", text, "你好世界")
}
}
func TestDeepgramService_buildURL(t *testing.T) {
svc := NewDeepgramService("key", "wss://api.deepgram.com/v1/listen", zap.NewNop().Sugar())
tests := []struct {
name string
opts Options
want []string // URL 中应包含的参数
}{
{
name: "defaults",
opts: Options{},
want: []string{"encoding=pcm_s16le", "sample_rate=16000", "language=zh-CN"},
},
{
name: "custom",
opts: Options{Encoding: "wav", SampleRate: 44100, Language: "en"},
want: []string{"encoding=wav", "sample_rate=44100", "language=en"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
u := svc.buildURL(tt.opts)
for _, param := range tt.want {
if !strings.Contains(u, param) {
t.Errorf("buildURL() = %q, should contain %q", u, param)
}
}
})
}
}

View File

@@ -0,0 +1,16 @@
package stt
import "context"
// Service 语音识别服务契约。
type Service interface {
// Recognize 识别一段完整音频,返回最终文本。
Recognize(ctx context.Context, audio []byte, opts Options) (string, error)
}
// Options 语音识别参数。
type Options struct {
Encoding string // 音频编码,如 "pcm_s16le"
SampleRate int // 采样率,如 16000
Language string // 语言,如 "zh-CN"
}

View File

@@ -0,0 +1,149 @@
package tts
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"time"
"go.uber.org/zap"
)
// OpenAIService 基于 OpenAI TTS API 的语音合成实现。
type OpenAIService struct {
apiKey string
voice string
speed float64
endpoint string
timeout time.Duration
logger *zap.SugaredLogger
client *http.Client
}
// NewOpenAIService 创建 OpenAI TTS 服务。
func NewOpenAIService(apiKey, voice, endpoint string, speed float64, timeoutSec int, logger *zap.SugaredLogger) *OpenAIService {
if voice == "" {
voice = "alloy"
}
if speed <= 0 {
speed = 1.0
}
if endpoint == "" {
endpoint = "https://api.openai.com/v1"
}
timeout := time.Duration(timeoutSec) * time.Second
if timeout <= 0 {
timeout = 5 * time.Second
}
return &OpenAIService{
apiKey: apiKey,
voice: voice,
speed: speed,
endpoint: endpoint,
timeout: timeout,
logger: logger,
client: &http.Client{Timeout: 30 * time.Second},
}
}
// ttsRequest OpenAI TTS API 请求。
type ttsRequest struct {
Model string `json:"model"`
Input string `json:"input"`
Voice string `json:"voice"`
ResponseFormat string `json:"response_format"`
Speed float64 `json:"speed"`
}
// SynthesizeStream 实现 tts.Service。从 textStream 读取句子,逐句调用 OpenAI TTS API。
func (o *OpenAIService) SynthesizeStream(ctx context.Context, textStream <-chan string, opts Options) (<-chan Chunk, error) {
voice := opts.Voice
if voice == "" {
voice = o.voice
}
speed := opts.Speed
if speed <= 0 {
speed = o.speed
}
ch := make(chan Chunk, 4)
go func() {
defer close(ch)
for text := range textStream {
if text == "" {
continue
}
audio, err := o.synthesize(ctx, text, voice, speed)
if err != nil {
o.logger.Warnw("tts: synthesize failed", "error", err, "text", text)
// 静默跳过,不中断整个流
continue
}
select {
case ch <- Chunk{Audio: audio, IsLast: false}:
case <-ctx.Done():
return
}
}
// textStream 关闭,发送 IsLast 标记
select {
case ch <- Chunk{Audio: nil, IsLast: true}:
case <-ctx.Done():
}
}()
return ch, nil
}
// synthesize 调用 OpenAI TTS API 合成单个句子。
func (o *OpenAIService) synthesize(ctx context.Context, text, voice string, speed float64) ([]byte, error) {
// 单句超时
ctx, cancel := context.WithTimeout(ctx, o.timeout)
defer cancel()
body := ttsRequest{
Model: "tts-1",
Input: text,
Voice: voice,
ResponseFormat: "mp3",
Speed: speed,
}
payload, err := json.Marshal(body)
if err != nil {
return nil, fmt.Errorf("tts: marshal request: %w", err)
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, o.endpoint+"/audio/speech", bytes.NewReader(payload))
if err != nil {
return nil, fmt.Errorf("tts: create request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+o.apiKey)
resp, err := o.client.Do(req)
if err != nil {
return nil, fmt.Errorf("tts: send request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
errBody, _ := io.ReadAll(resp.Body)
return nil, fmt.Errorf("tts: api error (status %d): %s", resp.StatusCode, string(errBody))
}
// 读取整个 MP3 响应
audio, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("tts: read response: %w", err)
}
return audio, nil
}

View File

@@ -0,0 +1,300 @@
package tts
import (
"context"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"go.uber.org/zap"
)
// mockTTSServer 创建模拟 OpenAI TTS API 的 HTTP 服务器。
func mockTTSServer(t *testing.T, handler http.HandlerFunc) *httptest.Server {
t.Helper()
return httptest.NewServer(handler)
}
// sendSentences 向 channel 发送句子并关闭。
func sendSentences(sentences ...string) <-chan string {
ch := make(chan string, len(sentences))
for _, s := range sentences {
ch <- s
}
close(ch)
return ch
}
func TestOpenAIService_SynthesizeStream_Success(t *testing.T) {
var callCount int32
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
atomic.AddInt32(&callCount, 1)
if r.Method != http.MethodPost {
t.Errorf("method = %s, want POST", r.Method)
}
if !strings.Contains(r.URL.Path, "/audio/speech") {
t.Errorf("path = %s, should contain /audio/speech", r.URL.Path)
}
auth := r.Header.Get("Authorization")
if auth != "Bearer test-key" {
t.Errorf("Authorization = %q, want %q", auth, "Bearer test-key")
}
// 验证请求体
body, _ := io.ReadAll(r.Body)
if !strings.Contains(string(body), "tts-1") {
t.Errorf("request body should contain model tts-1")
}
// 返回假 MP3 数据
w.Header().Set("Content-Type", "audio/mpeg")
fmt.Fprintf(w, "fake-mp3-data")
})
defer srv.Close()
svc := NewOpenAIService("test-key", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
textStream := sendSentences("你好", "世界", "")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{
Voice: "alloy", Speed: 1.0, OutputFmt: "mp3", SampleRate: 24000,
})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 应该有 3 个音频 chunk + 1 个 IsLast 标记
if len(chunks) != 4 {
t.Fatalf("got %d chunks, want 4", len(chunks))
}
// 验证前 3 个有音频数据
for i := 0; i < 3; i++ {
if string(chunks[i].Audio) != "fake-mp3-data" {
t.Errorf("chunk[%d].Audio = %q, want %q", i, string(chunks[i].Audio), "fake-mp3-data")
}
if chunks[i].IsLast {
t.Errorf("chunk[%d].IsLast should be false", i)
}
}
// 验证最后一个是 IsLast
if !chunks[3].IsLast {
t.Error("last chunk should be IsLast")
}
if chunks[3].Audio != nil {
t.Error("last chunk Audio should be nil")
}
// 验证调用了 3 次 API3 个句子)
if atomic.LoadInt32(&callCount) != 3 {
t.Errorf("API called %d times, want 3", callCount)
}
}
func TestOpenAIService_SynthesizeStream_APIError(t *testing.T) {
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
fmt.Fprintf(w, "internal error")
})
defer srv.Close()
svc := NewOpenAIService("test-key", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
textStream := sendSentences("你好")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
// 应该只有一个 IsLast chunk音频被跳过
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
if len(chunks) != 1 {
t.Fatalf("got %d chunks, want 1 (IsLast only)", len(chunks))
}
if !chunks[0].IsLast {
t.Error("chunk should be IsLast")
}
}
func TestOpenAIService_SynthesizeStream_Timeout(t *testing.T) {
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
time.Sleep(3 * time.Second)
w.Header().Set("Content-Type", "audio/mpeg")
fmt.Fprintf(w, "late-mp3")
})
defer srv.Close()
// 1 秒超时
svc := NewOpenAIService("test-key", "alloy", srv.URL, 1.0, 1, zap.NewNop().Sugar())
textStream := sendSentences("很长的句子")
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
ch, err := svc.SynthesizeStream(ctx, textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 超时后音频被跳过,只有 IsLast
if len(chunks) != 1 {
t.Fatalf("got %d chunks, want 1", len(chunks))
}
if !chunks[0].IsLast {
t.Error("chunk should be IsLast")
}
}
func TestOpenAIService_SynthesizeStream_EmptyText(t *testing.T) {
var callCount int32
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
atomic.AddInt32(&callCount, 1)
w.Header().Set("Content-Type", "audio/mpeg")
fmt.Fprintf(w, "mp3")
})
defer srv.Close()
svc := NewOpenAIService("test-key", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
// 空句子应该被跳过
textStream := sendSentences("", "你好", "")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 只有 "你好" 应该被合成
if atomic.LoadInt32(&callCount) != 1 {
t.Errorf("API called %d times, want 1", callCount)
}
// 1 个音频 + 1 个 IsLast
if len(chunks) != 2 {
t.Fatalf("got %d chunks, want 2", len(chunks))
}
}
func TestOpenAIService_SynthesizeStream_ContextCancelled(t *testing.T) {
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "audio/mpeg")
fmt.Fprintf(w, "mp3")
})
defer srv.Close()
svc := NewOpenAIService("test-key", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
// 发送多个句子,但在第一个后取消
textStream := make(chan string, 3)
textStream <- "第一句"
textStream <- "第二句"
textStream <- "第三句"
close(textStream)
ctx, cancel := context.WithCancel(context.Background())
// 立即取消
cancel()
ch, err := svc.SynthesizeStream(ctx, textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
// 消费 channel应该很快结束
var count int
for range ch {
count++
}
// 可能收到 0 个或 1 个 chunk取决于时序
t.Logf("received %d chunks after context cancel", count)
}
func TestOpenAIService_SynthesizeStream_PartialFailure(t *testing.T) {
var callCount int32
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
n := atomic.AddInt32(&callCount, 1)
if n == 2 {
// 第二个句子失败
w.WriteHeader(http.StatusInternalServerError)
fmt.Fprintf(w, "error")
return
}
w.Header().Set("Content-Type", "audio/mpeg")
fmt.Fprintf(w, "mp3-%d", n)
})
defer srv.Close()
svc := NewOpenAIService("test-key", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
textStream := sendSentences("第一句", "第二句", "第三句")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 2 个成功音频 + 1 个 IsLast第二句被跳过
if len(chunks) != 3 {
t.Fatalf("got %d chunks, want 3", len(chunks))
}
if !chunks[len(chunks)-1].IsLast {
t.Error("last chunk should be IsLast")
}
}
func TestOpenAIService_SynthesizeStream_CustomVoice(t *testing.T) {
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
if !strings.Contains(string(body), "nova") {
t.Errorf("request body should contain voice 'nova', got: %s", string(body))
}
w.Header().Set("Content-Type", "audio/mpeg")
fmt.Fprintf(w, "mp3")
})
defer srv.Close()
svc := NewOpenAIService("test-key", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
textStream := sendSentences("你好")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{Voice: "nova"})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
for range ch {
}
}

View File

@@ -0,0 +1,25 @@
package tts
import "context"
// Service 语音合成服务契约。
type Service interface {
// SynthesizeStream 流式合成。
// textStream 接收句子级文本(由 Orchestrator 的句子切分器产出),
// 返回的 channel 持续输出 MP3 音频 chunk。
SynthesizeStream(ctx context.Context, textStream <-chan string, opts Options) (<-chan Chunk, error)
}
// Options 合成参数。
type Options struct {
Voice string // "alloy" | "nova" | "shimmer" 等
Speed float64 // 1.0 为正常语速
OutputFmt string // "mp3" — 固定使用 MP3
SampleRate int // 24000
}
// Chunk 一个音频片段。
type Chunk struct {
Audio []byte // MP3 音频数据(未 Base64 编码)
IsLast bool // 是否为最后一片
}