133 lines
3.0 KiB
Go
133 lines
3.0 KiB
Go
package services
|
|||
|
|
|
||
|
|
import (
|
||
|
|
"bytes"
|
||
|
|
"encoding/json"
|
||
|
|
"fmt"
|
||
|
|
"io"
|
||
|
|
"net/http"
|
||
|
|
|
||
|
|
"d:\项目\dobaochet\backend\config"
|
||
|
|
)
|
||
|
|
|
||
|
|
// AIService AI服务结构体
|
||
|
|
type AIService struct {
|
||
|
|
cfg *config.Config
|
||
|
|
client *http.Client
|
||
|
|
}
|
||
|
|
|
||
|
|
// NewAIService 创建AI服务实例
|
||
|
|
func NewAIService(cfg *config.Config) *AIService {
|
||
|
|
return &AIService{
|
||
|
|
cfg: cfg,
|
||
|
|
client: &http.Client{},
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// AIChatRequest AI聊天请求结构体
|
||
|
|
type AIChatRequest struct {
|
||
|
|
Model string `json:"model"`
|
||
|
|
Messages []Message `json:"messages"`
|
||
|
|
MaxTokens int64 `json:"max_tokens"`
|
||
|
|
Temperature float64 `json:"temperature"`
|
||
|
|
}
|
||
|
|
|
||
|
|
// Message 消息结构体
|
||
|
|
type Message struct {
|
||
|
|
Role string `json:"role"`
|
||
|
|
Content string `json:"content"`
|
||
|
|
}
|
||
|
|
|
||
|
|
// AIChatResponse AI聊天响应结构体
|
||
|
|
type AIChatResponse struct {
|
||
|
|
ID string `json:"id"`
|
||
|
|
Object string `json:"object"`
|
||
|
|
Created int64 `json:"created"`
|
||
|
|
Model string `json:"model"`
|
||
|
|
Choices []Choice `json:"choices"`
|
||
|
|
Usage Usage `json:"usage"`
|
||
|
|
}
|
||
|
|
|
||
|
|
// Choice 响应选项结构体
|
||
|
|
type Choice struct {
|
||
|
|
Index int `json:"index"`
|
||
|
|
Message Message `json:"message"`
|
||
|
|
FinishReason string `json:"finish_reason"`
|
||
|
|
}
|
||
|
|
|
||
|
|
// Usage 用量结构体
|
||
|
|
type Usage struct {
|
||
|
|
PromptTokens int64 `json:"prompt_tokens"`
|
||
|
|
CompletionTokens int64 `json:"completion_tokens"`
|
||
|
|
TotalTokens int64 `json:"total_tokens"`
|
||
|
|
}
|
||
|
|
|
||
|
|
// Chat 调用AI模型聊天
|
||
|
|
func (s *AIService) Chat(messages []Message) (*AIChatResponse, error) {
|
||
|
|
// 构建请求
|
||
|
|
reqBody := AIChatRequest{
|
||
|
|
Model: s.cfg.AI.ModelName,
|
||
|
|
Messages: messages,
|
||
|
|
MaxTokens: s.cfg.AI.MaxTokens,
|
||
|
|
Temperature: s.cfg.AI.Temperature,
|
||
|
|
}
|
||
|
|
|
||
|
|
// 序列化请求体
|
||
|
|
reqBytes, err := json.Marshal(reqBody)
|
||
|
|
if err != nil {
|
||
|
|
return nil, fmt.Errorf("序列化请求失败: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// 创建HTTP请求
|
||
|
|
req, err := http.NewRequest("POST", s.cfg.AI.APIURL, bytes.NewBuffer(reqBytes))
|
||
|
|
if err != nil {
|
||
|
|
return nil, fmt.Errorf("创建请求失败: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// 设置请求头
|
||
|
|
req.Header.Set("Content-Type", "application/json")
|
||
|
|
if s.cfg.AI.APIKey != "" {
|
||
|
|
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", s.cfg.AI.APIKey))
|
||
|
|
}
|
||
|
|
|
||
|
|
// 发送请求
|
||
|
|
resp, err := s.client.Do(req)
|
||
|
|
if err != nil {
|
||
|
|
return nil, fmt.Errorf("发送请求失败: %w", err)
|
||
|
|
}
|
||
|
|
defer resp.Body.Close()
|
||
|
|
|
||
|
|
// 读取响应
|
||
|
|
respBody, err := io.ReadAll(resp.Body)
|
||
|
|
if err != nil {
|
||
|
|
return nil, fmt.Errorf("读取响应失败: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// 检查响应状态
|
||
|
|
if resp.StatusCode != http.StatusOK {
|
||
|
|
return nil, fmt.Errorf("AI API返回错误: %s, 响应: %s", resp.Status, string(respBody))
|
||
|
|
}
|
||
|
|
|
||
|
|
// 解析响应
|
||
|
|
var aiResp AIChatResponse
|
||
|
|
if err := json.Unmarshal(respBody, &aiResp); err != nil {
|
||
|
|
return nil, fmt.Errorf("解析响应失败: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
return &aiResp, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// CountTokens 估算令牌数量(简化实现,实际应使用更准确的方法)
|
||
|
|
func (s *AIService) CountTokens(text string) int64 {
|
||
|
|
// 简单估算:每个汉字算2个令牌,每个英文单词算1个令牌
|
||
|
|
var count int64
|
||
|
|
for _, r := range text {
|
||
|
|
if r > 127 {
|
||
|
|
count += 2
|
||
|
|
} else {
|
||
|
|
count++
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return count / 2 // 平均估算
|
||
|
|
}
|