增加相关功能
This commit is contained in:
@@ -0,0 +1,794 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"server/models"
|
||||
"server/pkg/jwtutil"
|
||||
|
||||
beego "github.com/beego/beego/v2/server/web"
|
||||
)
|
||||
|
||||
// BackendAiChatController AI聊天消息控制器
|
||||
type BackendAiChatController struct {
|
||||
beego.Controller
|
||||
}
|
||||
|
||||
func (c *BackendAiChatController) chatClaims() (*jwtutil.Claims, error) {
|
||||
auth := c.Ctx.Request.Header.Get("Authorization")
|
||||
if auth == "" {
|
||||
return nil, fmt.Errorf("未登录")
|
||||
}
|
||||
parts := strings.SplitN(auth, " ", 2)
|
||||
if len(parts) != 2 || parts[0] != "Bearer" {
|
||||
return nil, fmt.Errorf("认证信息格式错误")
|
||||
}
|
||||
claims, err := jwtutil.ParseToken(parts[1])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("无效的token")
|
||||
}
|
||||
if claims.UserType != "backend" {
|
||||
return nil, fmt.Errorf("无权访问")
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
func (c *BackendAiChatController) chatJsonErr(httpStatus, bizCode int, msg string) {
|
||||
c.Ctx.Output.SetStatus(httpStatus)
|
||||
c.Data["json"] = map[string]interface{}{"code": bizCode, "msg": msg}
|
||||
_ = c.ServeJSON()
|
||||
}
|
||||
|
||||
func (c *BackendAiChatController) chatOk(data interface{}) {
|
||||
c.Data["json"] = map[string]interface{}{"code": 200, "msg": "success", "data": data}
|
||||
_ = c.ServeJSON()
|
||||
}
|
||||
|
||||
type chatSendPayload struct {
|
||||
SessionID uint64 `json:"session_id"`
|
||||
Content string `json:"content"`
|
||||
ProviderID uint64 `json:"provider_id"`
|
||||
Model string `json:"model"`
|
||||
}
|
||||
|
||||
type openaiMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
type openaiRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []openaiMessage `json:"messages"`
|
||||
}
|
||||
|
||||
type openaiResponse struct {
|
||||
Choices []struct {
|
||||
Message struct {
|
||||
Content string `json:"content"`
|
||||
} `json:"message"`
|
||||
} `json:"choices"`
|
||||
Error *struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"error"`
|
||||
}
|
||||
|
||||
type anthropicRequest struct {
|
||||
Model string `json:"model"`
|
||||
MaxTokens int `json:"max_tokens"`
|
||||
System string `json:"system,omitempty"`
|
||||
Messages []openaiMessage `json:"messages"`
|
||||
}
|
||||
|
||||
type anthropicResponse struct {
|
||||
Content []struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"content"`
|
||||
Error *struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"error"`
|
||||
}
|
||||
|
||||
// MessageList GET /backend/ai/chat/message/list?session_id=xxx
|
||||
func (c *BackendAiChatController) MessageList() {
|
||||
claims, err := c.chatClaims()
|
||||
if err != nil {
|
||||
c.chatJsonErr(401, 401, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
sessionIDStr := strings.TrimSpace(c.GetString("session_id"))
|
||||
if sessionIDStr == "" {
|
||||
c.chatJsonErr(400, 400, "缺少会话ID")
|
||||
return
|
||||
}
|
||||
sessionID, err := strconv.ParseUint(sessionIDStr, 10, 64)
|
||||
if err != nil {
|
||||
c.chatJsonErr(400, 400, "会话ID格式错误")
|
||||
return
|
||||
}
|
||||
|
||||
// 验证会话归属
|
||||
session := models.BackendAiChatSession{ID: sessionID}
|
||||
if err := models.Orm.Read(&session); err != nil {
|
||||
c.chatJsonErr(404, 404, "会话不存在")
|
||||
return
|
||||
}
|
||||
if session.TenantID != fmt.Sprintf("%d", claims.TenantId) || session.UserID != uint64(claims.UserID) {
|
||||
c.chatJsonErr(403, 403, "无权访问")
|
||||
return
|
||||
}
|
||||
|
||||
var list []models.BackendAiChatMessage
|
||||
_, err = models.Orm.QueryTable(new(models.BackendAiChatMessage)).
|
||||
Filter("session_id", sessionID).
|
||||
OrderBy("id").
|
||||
All(&list)
|
||||
if err != nil {
|
||||
c.chatJsonErr(500, 500, "查询失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.chatOk(map[string]interface{}{"list": list, "title": session.Title})
|
||||
}
|
||||
|
||||
// DeleteMessage DELETE /backend/ai/chat/message/:id
|
||||
func (c *BackendAiChatController) DeleteMessage() {
|
||||
claims, err := c.chatClaims()
|
||||
if err != nil {
|
||||
c.chatJsonErr(401, 401, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
idStr := c.Ctx.Input.Param(":id")
|
||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||
if err != nil {
|
||||
c.chatJsonErr(400, 400, "消息ID格式错误")
|
||||
return
|
||||
}
|
||||
|
||||
msg := models.BackendAiChatMessage{ID: id}
|
||||
if err := models.Orm.Read(&msg); err != nil {
|
||||
c.chatJsonErr(404, 404, "消息不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 验证会话归属
|
||||
session := models.BackendAiChatSession{ID: msg.SessionID}
|
||||
if err := models.Orm.Read(&session); err != nil {
|
||||
c.chatJsonErr(404, 404, "会话不存在")
|
||||
return
|
||||
}
|
||||
if session.TenantID != fmt.Sprintf("%d", claims.TenantId) || session.UserID != uint64(claims.UserID) {
|
||||
c.chatJsonErr(403, 403, "无权操作")
|
||||
return
|
||||
}
|
||||
|
||||
if _, err := models.Orm.Delete(&msg); err != nil {
|
||||
c.chatJsonErr(500, 500, "删除失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.chatOk(nil)
|
||||
}
|
||||
|
||||
// Send POST /backend/ai/chat/send
|
||||
func (c *BackendAiChatController) Send() {
|
||||
claims, err := c.chatClaims()
|
||||
if err != nil {
|
||||
c.chatJsonErr(401, 401, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(c.Ctx.Request.Body)
|
||||
if err != nil {
|
||||
c.chatJsonErr(400, 400, "读取请求体失败")
|
||||
return
|
||||
}
|
||||
var p chatSendPayload
|
||||
if err := json.Unmarshal(body, &p); err != nil {
|
||||
c.chatJsonErr(400, 400, "参数格式错误")
|
||||
return
|
||||
}
|
||||
|
||||
if strings.TrimSpace(p.Content) == "" {
|
||||
c.chatJsonErr(400, 400, "消息内容不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
// 查找或创建会话
|
||||
var session models.BackendAiChatSession
|
||||
if p.SessionID > 0 {
|
||||
session = models.BackendAiChatSession{ID: p.SessionID}
|
||||
if err := models.Orm.Read(&session); err != nil {
|
||||
c.chatJsonErr(404, 404, "会话不存在")
|
||||
return
|
||||
}
|
||||
if session.TenantID != fmt.Sprintf("%d", claims.TenantId) || session.UserID != uint64(claims.UserID) {
|
||||
c.chatJsonErr(403, 403, "无权操作")
|
||||
return
|
||||
}
|
||||
} else {
|
||||
// 新建会话,标题取消息前20字
|
||||
title := p.Content
|
||||
if len([]rune(title)) > 20 {
|
||||
title = string([]rune(title)[:20]) + "..."
|
||||
}
|
||||
session = models.BackendAiChatSession{
|
||||
TenantID: fmt.Sprintf("%d", claims.TenantId),
|
||||
UserID: uint64(claims.UserID),
|
||||
ProviderID: p.ProviderID,
|
||||
Title: title,
|
||||
CreateTime: time.Now(),
|
||||
UpdateTime: time.Now(),
|
||||
}
|
||||
id, err := models.Orm.Insert(&session)
|
||||
if err != nil {
|
||||
c.chatJsonErr(500, 500, "创建会话失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
session.ID = uint64(id)
|
||||
}
|
||||
|
||||
// 查找AI接入配置
|
||||
providerID := p.ProviderID
|
||||
if providerID == 0 {
|
||||
providerID = session.ProviderID
|
||||
}
|
||||
|
||||
var provider models.BackendAiProvider
|
||||
if providerID > 0 {
|
||||
provider = models.BackendAiProvider{ID: providerID}
|
||||
if err := models.Orm.Read(&provider); err != nil {
|
||||
c.chatJsonErr(400, 400, "指定的AI接入配置不存在,请先在设置中配置")
|
||||
return
|
||||
}
|
||||
if provider.TenantID != fmt.Sprintf("%d", claims.TenantId) || provider.UserID != uint64(claims.UserID) {
|
||||
c.chatJsonErr(403, 403, "无权使用该配置")
|
||||
return
|
||||
}
|
||||
if provider.Status != 1 {
|
||||
c.chatJsonErr(400, 400, "该AI接入配置已被禁用")
|
||||
return
|
||||
}
|
||||
} else {
|
||||
// 找第一个启用的配置
|
||||
var providers []models.BackendAiProvider
|
||||
_, err := models.Orm.QueryTable(new(models.BackendAiProvider)).
|
||||
Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)).
|
||||
Filter("user_id", uint64(claims.UserID)).
|
||||
Filter("status", 1).
|
||||
Filter("delete_time__isnull", true).
|
||||
OrderBy("-id").
|
||||
Limit(1).
|
||||
All(&providers)
|
||||
if err != nil || len(providers) == 0 {
|
||||
c.chatJsonErr(400, 400, "尚未配置AI接入,请先在设置中配置OpenAI或Anthropic接入")
|
||||
return
|
||||
}
|
||||
provider = providers[0]
|
||||
}
|
||||
|
||||
// 加载历史消息(最近20条)
|
||||
var history []models.BackendAiChatMessage
|
||||
_, _ = models.Orm.QueryTable(new(models.BackendAiChatMessage)).
|
||||
Filter("session_id", session.ID).
|
||||
OrderBy("-id").
|
||||
Limit(20).
|
||||
All(&history)
|
||||
// 反转顺序
|
||||
for i, j := 0, len(history)-1; i < j; i, j = i+1, j-1 {
|
||||
history[i], history[j] = history[j], history[i]
|
||||
}
|
||||
|
||||
// 构建消息列表
|
||||
messages := make([]openaiMessage, 0, len(history)+1)
|
||||
for _, m := range history {
|
||||
messages = append(messages, openaiMessage{Role: m.Role, Content: m.Content})
|
||||
}
|
||||
messages = append(messages, openaiMessage{Role: "user", Content: p.Content})
|
||||
|
||||
// 保存用户消息
|
||||
userMsg := models.BackendAiChatMessage{
|
||||
TenantID: fmt.Sprintf("%d", claims.TenantId),
|
||||
SessionID: session.ID,
|
||||
Role: "user",
|
||||
Content: p.Content,
|
||||
CreateTime: time.Now(),
|
||||
}
|
||||
_, _ = models.Orm.Insert(&userMsg)
|
||||
|
||||
// 确定使用的模型:用户指定 > provider第一个模型
|
||||
useModel := strings.TrimSpace(p.Model)
|
||||
if useModel == "" {
|
||||
providerModels := parseProviderModels(provider.Models)
|
||||
if len(providerModels) > 0 {
|
||||
useModel = providerModels[0]
|
||||
}
|
||||
}
|
||||
if useModel == "" {
|
||||
c.chatJsonErr(400, 400, "未指定模型且该接入配置无可用模型")
|
||||
return
|
||||
}
|
||||
|
||||
// 查询用户默认角色预设
|
||||
var systemPrompt string
|
||||
var presets []models.BackendAiChatPreset
|
||||
_, _ = models.Orm.QueryTable(new(models.BackendAiChatPreset)).
|
||||
Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)).
|
||||
Filter("user_id", uint64(claims.UserID)).
|
||||
Filter("is_default", 1).
|
||||
Filter("delete_time__isnull", true).
|
||||
Limit(1).
|
||||
All(&presets)
|
||||
if len(presets) > 0 {
|
||||
systemPrompt = presets[0].Content
|
||||
}
|
||||
|
||||
// 调用AI
|
||||
reply, err := callAI(provider, useModel, systemPrompt, messages)
|
||||
if err != nil {
|
||||
c.chatJsonErr(500, 500, "AI调用失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 保存AI回复
|
||||
assistantMsg := models.BackendAiChatMessage{
|
||||
TenantID: fmt.Sprintf("%d", claims.TenantId),
|
||||
SessionID: session.ID,
|
||||
Role: "assistant",
|
||||
Content: reply,
|
||||
CreateTime: time.Now(),
|
||||
}
|
||||
_, _ = models.Orm.Insert(&assistantMsg)
|
||||
|
||||
// 更新会话时间
|
||||
session.UpdateTime = time.Now()
|
||||
_, _ = models.Orm.Update(&session, "update_time")
|
||||
|
||||
c.chatOk(map[string]interface{}{
|
||||
"session_id": session.ID,
|
||||
"reply": reply,
|
||||
"message_id": assistantMsg.ID,
|
||||
})
|
||||
}
|
||||
|
||||
func callAI(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage) (string, error) {
|
||||
client := &http.Client{Timeout: 120 * time.Second}
|
||||
|
||||
if provider.ProviderType == "openai" {
|
||||
// OpenAI 兼容格式
|
||||
reqMessages := messages
|
||||
if systemPrompt != "" {
|
||||
reqMessages = append([]openaiMessage{{Role: "system", Content: systemPrompt}}, messages...)
|
||||
}
|
||||
reqBody := openaiRequest{
|
||||
Model: model,
|
||||
Messages: reqMessages,
|
||||
}
|
||||
jsonData, _ := json.Marshal(reqBody)
|
||||
|
||||
// 处理URL:如果已经包含/chat/completions,直接使用;否则拼接
|
||||
url := strings.TrimRight(provider.ApiBase, "/")
|
||||
if !strings.HasSuffix(url, "/chat/completions") {
|
||||
url = url + "/chat/completions"
|
||||
}
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+provider.ApiKey)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result openaiResponse
|
||||
if err := json.Unmarshal(body, &result); err != nil {
|
||||
return "", fmt.Errorf("解析响应失败: %s", string(body))
|
||||
}
|
||||
if result.Error != nil {
|
||||
return "", fmt.Errorf(result.Error.Message)
|
||||
}
|
||||
if len(result.Choices) == 0 {
|
||||
return "", fmt.Errorf("AI未返回内容")
|
||||
}
|
||||
return result.Choices[0].Message.Content, nil
|
||||
|
||||
} else {
|
||||
// Anthropic 格式
|
||||
reqBody := anthropicRequest{
|
||||
Model: model,
|
||||
MaxTokens: 4096,
|
||||
System: systemPrompt,
|
||||
Messages: messages,
|
||||
}
|
||||
jsonData, _ := json.Marshal(reqBody)
|
||||
|
||||
// 处理URL:如果已经包含/v1/messages,直接使用;否则拼接
|
||||
url := strings.TrimRight(provider.ApiBase, "/")
|
||||
if !strings.HasSuffix(url, "/v1/messages") {
|
||||
url = url + "/v1/messages"
|
||||
}
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("x-api-key", provider.ApiKey)
|
||||
req.Header.Set("anthropic-version", "2023-06-01")
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result anthropicResponse
|
||||
if err := json.Unmarshal(body, &result); err != nil {
|
||||
return "", fmt.Errorf("解析响应失败: %s", string(body))
|
||||
}
|
||||
if result.Error != nil {
|
||||
return "", fmt.Errorf(result.Error.Message)
|
||||
}
|
||||
if len(result.Content) == 0 {
|
||||
return "", fmt.Errorf("AI未返回内容")
|
||||
}
|
||||
return result.Content[0].Text, nil
|
||||
}
|
||||
}
|
||||
|
||||
// SendStream POST /backend/ai/chat/send-stream
|
||||
// 流式响应(SSE),逐块返回AI回复
|
||||
func (c *BackendAiChatController) SendStream() {
|
||||
claims, err := c.chatClaims()
|
||||
if err != nil {
|
||||
c.chatJsonErr(401, 401, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(c.Ctx.Request.Body)
|
||||
if err != nil {
|
||||
c.chatJsonErr(400, 400, "读取请求体失败")
|
||||
return
|
||||
}
|
||||
var p chatSendPayload
|
||||
if err := json.Unmarshal(body, &p); err != nil {
|
||||
c.chatJsonErr(400, 400, "参数格式错误")
|
||||
return
|
||||
}
|
||||
|
||||
if strings.TrimSpace(p.Content) == "" {
|
||||
c.chatJsonErr(400, 400, "消息内容不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
// 查找或创建会话
|
||||
var session models.BackendAiChatSession
|
||||
if p.SessionID > 0 {
|
||||
session = models.BackendAiChatSession{ID: p.SessionID}
|
||||
if err := models.Orm.Read(&session); err != nil {
|
||||
c.chatJsonErr(404, 404, "会话不存在")
|
||||
return
|
||||
}
|
||||
if session.TenantID != fmt.Sprintf("%d", claims.TenantId) || session.UserID != uint64(claims.UserID) {
|
||||
c.chatJsonErr(403, 403, "无权操作")
|
||||
return
|
||||
}
|
||||
} else {
|
||||
title := p.Content
|
||||
if len([]rune(title)) > 20 {
|
||||
title = string([]rune(title)[:20]) + "..."
|
||||
}
|
||||
session = models.BackendAiChatSession{
|
||||
TenantID: fmt.Sprintf("%d", claims.TenantId),
|
||||
UserID: uint64(claims.UserID),
|
||||
ProviderID: p.ProviderID,
|
||||
Title: title,
|
||||
CreateTime: time.Now(),
|
||||
UpdateTime: time.Now(),
|
||||
}
|
||||
id, err := models.Orm.Insert(&session)
|
||||
if err != nil {
|
||||
c.chatJsonErr(500, 500, "创建会话失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
session.ID = uint64(id)
|
||||
}
|
||||
|
||||
// 查找AI接入配置
|
||||
providerID := p.ProviderID
|
||||
if providerID == 0 {
|
||||
providerID = session.ProviderID
|
||||
}
|
||||
|
||||
var provider models.BackendAiProvider
|
||||
if providerID > 0 {
|
||||
provider = models.BackendAiProvider{ID: providerID}
|
||||
if err := models.Orm.Read(&provider); err != nil {
|
||||
c.chatJsonErr(400, 400, "指定的AI接入配置不存在")
|
||||
return
|
||||
}
|
||||
if provider.TenantID != fmt.Sprintf("%d", claims.TenantId) || provider.UserID != uint64(claims.UserID) {
|
||||
c.chatJsonErr(403, 403, "无权使用该配置")
|
||||
return
|
||||
}
|
||||
if provider.Status != 1 {
|
||||
c.chatJsonErr(400, 400, "该AI接入配置已被禁用")
|
||||
return
|
||||
}
|
||||
} else {
|
||||
var providers []models.BackendAiProvider
|
||||
_, err := models.Orm.QueryTable(new(models.BackendAiProvider)).
|
||||
Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)).
|
||||
Filter("user_id", uint64(claims.UserID)).
|
||||
Filter("status", 1).
|
||||
Filter("delete_time__isnull", true).
|
||||
OrderBy("-id").
|
||||
Limit(1).
|
||||
All(&providers)
|
||||
if err != nil || len(providers) == 0 {
|
||||
c.chatJsonErr(400, 400, "尚未配置AI接入,请先在设置中配置")
|
||||
return
|
||||
}
|
||||
provider = providers[0]
|
||||
}
|
||||
|
||||
// 确定模型
|
||||
useModel := strings.TrimSpace(p.Model)
|
||||
if useModel == "" {
|
||||
providerModels := parseProviderModels(provider.Models)
|
||||
if len(providerModels) > 0 {
|
||||
useModel = providerModels[0]
|
||||
}
|
||||
}
|
||||
if useModel == "" {
|
||||
c.chatJsonErr(400, 400, "未指定模型且该接入配置无可用模型")
|
||||
return
|
||||
}
|
||||
|
||||
// 查询默认预设
|
||||
var systemPrompt string
|
||||
var presets []models.BackendAiChatPreset
|
||||
_, _ = models.Orm.QueryTable(new(models.BackendAiChatPreset)).
|
||||
Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)).
|
||||
Filter("user_id", uint64(claims.UserID)).
|
||||
Filter("is_default", 1).
|
||||
Filter("delete_time__isnull", true).
|
||||
Limit(1).
|
||||
All(&presets)
|
||||
if len(presets) > 0 {
|
||||
systemPrompt = presets[0].Content
|
||||
}
|
||||
|
||||
// 加载历史消息
|
||||
var history []models.BackendAiChatMessage
|
||||
_, _ = models.Orm.QueryTable(new(models.BackendAiChatMessage)).
|
||||
Filter("session_id", session.ID).
|
||||
OrderBy("-id").
|
||||
Limit(20).
|
||||
All(&history)
|
||||
for i, j := 0, len(history)-1; i < j; i, j = i+1, j-1 {
|
||||
history[i], history[j] = history[j], history[i]
|
||||
}
|
||||
|
||||
messages := make([]openaiMessage, 0, len(history)+1)
|
||||
for _, m := range history {
|
||||
messages = append(messages, openaiMessage{Role: m.Role, Content: m.Content})
|
||||
}
|
||||
messages = append(messages, openaiMessage{Role: "user", Content: p.Content})
|
||||
|
||||
// 保存用户消息
|
||||
userMsg := models.BackendAiChatMessage{
|
||||
TenantID: fmt.Sprintf("%d", claims.TenantId),
|
||||
SessionID: session.ID,
|
||||
Role: "user",
|
||||
Content: p.Content,
|
||||
CreateTime: time.Now(),
|
||||
}
|
||||
_, _ = models.Orm.Insert(&userMsg)
|
||||
|
||||
// 设置SSE响应头 - 直接操作底层ResponseWriter,绕过beego包装层
|
||||
rw := c.Ctx.ResponseWriter.ResponseWriter
|
||||
rw.Header().Set("Content-Type", "text/event-stream")
|
||||
rw.Header().Set("Cache-Control", "no-cache")
|
||||
rw.Header().Set("Connection", "keep-alive")
|
||||
rw.Header().Set("X-Accel-Buffering", "no")
|
||||
rw.WriteHeader(200)
|
||||
|
||||
var flusher http.Flusher
|
||||
if f, ok := rw.(http.Flusher); ok {
|
||||
flusher = f
|
||||
}
|
||||
|
||||
writeSSE := func(event string, data string) {
|
||||
rw.Write([]byte(fmt.Sprintf("event: %s\ndata: %s\n\n", event, data)))
|
||||
if flusher != nil {
|
||||
flusher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
// 发送session_id事件
|
||||
writeSSE("session", fmt.Sprintf("{\"session_id\":%d}", session.ID))
|
||||
|
||||
// 流式调用AI
|
||||
fullReply := ""
|
||||
streamErr := callAIStream(provider, useModel, systemPrompt, messages, func(chunk string) {
|
||||
fullReply += chunk
|
||||
chunkJSON, _ := json.Marshal(map[string]string{"content": chunk})
|
||||
writeSSE("content", string(chunkJSON))
|
||||
})
|
||||
|
||||
if streamErr != nil {
|
||||
errJSON, _ := json.Marshal(map[string]string{"error": streamErr.Error()})
|
||||
writeSSE("error", string(errJSON))
|
||||
return
|
||||
}
|
||||
|
||||
// 保存AI回复
|
||||
assistantMsg := models.BackendAiChatMessage{
|
||||
TenantID: fmt.Sprintf("%d", claims.TenantId),
|
||||
SessionID: session.ID,
|
||||
Role: "assistant",
|
||||
Content: fullReply,
|
||||
CreateTime: time.Now(),
|
||||
}
|
||||
_, _ = models.Orm.Insert(&assistantMsg)
|
||||
|
||||
// 更新会话时间
|
||||
session.UpdateTime = time.Now()
|
||||
_, _ = models.Orm.Update(&session, "update_time")
|
||||
|
||||
// 发送done事件
|
||||
doneJSON, _ := json.Marshal(map[string]interface{}{
|
||||
"session_id": session.ID,
|
||||
"message_id": assistantMsg.ID,
|
||||
})
|
||||
writeSSE("done", string(doneJSON))
|
||||
}
|
||||
|
||||
// callAIStream 流式调用AI,onChunk回调每收到一个内容块就调用
|
||||
func callAIStream(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, onChunk func(string)) error {
|
||||
client := &http.Client{Timeout: 120 * time.Second}
|
||||
|
||||
if provider.ProviderType == "openai" {
|
||||
// OpenAI 兼容格式 - 流式
|
||||
reqMessages := messages
|
||||
if systemPrompt != "" {
|
||||
reqMessages = append([]openaiMessage{{Role: "system", Content: systemPrompt}}, messages...)
|
||||
}
|
||||
reqBody := map[string]interface{}{
|
||||
"model": model,
|
||||
"messages": reqMessages,
|
||||
"stream": true,
|
||||
}
|
||||
jsonData, _ := json.Marshal(reqBody)
|
||||
|
||||
// 处理URL:如果已经包含/chat/completions,直接使用;否则拼接
|
||||
url := strings.TrimRight(provider.ApiBase, "/")
|
||||
if !strings.HasSuffix(url, "/chat/completions") {
|
||||
url = url + "/chat/completions"
|
||||
}
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+provider.ApiKey)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != 200 {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
scanner.Buffer(make([]byte, 1024*1024), 1024*1024)
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
if !strings.HasPrefix(line, "data: ") {
|
||||
continue
|
||||
}
|
||||
data := strings.TrimPrefix(line, "data: ")
|
||||
if data == "[DONE]" {
|
||||
break
|
||||
}
|
||||
var chunk struct {
|
||||
Choices []struct {
|
||||
Delta struct {
|
||||
Content string `json:"content"`
|
||||
} `json:"delta"`
|
||||
} `json:"choices"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(data), &chunk); err != nil {
|
||||
continue
|
||||
}
|
||||
if len(chunk.Choices) > 0 && chunk.Choices[0].Delta.Content != "" {
|
||||
onChunk(chunk.Choices[0].Delta.Content)
|
||||
}
|
||||
}
|
||||
return scanner.Err()
|
||||
|
||||
} else {
|
||||
// Anthropic 格式 - 流式
|
||||
reqBody := map[string]interface{}{
|
||||
"model": model,
|
||||
"max_tokens": 4096,
|
||||
"messages": messages,
|
||||
"stream": true,
|
||||
}
|
||||
if systemPrompt != "" {
|
||||
reqBody["system"] = systemPrompt
|
||||
}
|
||||
jsonData, _ := json.Marshal(reqBody)
|
||||
|
||||
// 处理URL:如果已经包含/v1/messages,直接使用;否则拼接
|
||||
url := strings.TrimRight(provider.ApiBase, "/")
|
||||
if !strings.HasSuffix(url, "/v1/messages") {
|
||||
url = url + "/v1/messages"
|
||||
}
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("x-api-key", provider.ApiKey)
|
||||
req.Header.Set("anthropic-version", "2023-06-01")
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != 200 {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
scanner.Buffer(make([]byte, 1024*1024), 1024*1024)
|
||||
currentEvent := ""
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
if strings.HasPrefix(line, "event: ") {
|
||||
currentEvent = strings.TrimPrefix(line, "event: ")
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(line, "data: ") {
|
||||
data := strings.TrimPrefix(line, "data: ")
|
||||
if currentEvent == "content_block_delta" {
|
||||
var chunk struct {
|
||||
Delta struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"delta"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(data), &chunk); err == nil && chunk.Delta.Text != "" {
|
||||
onChunk(chunk.Delta.Text)
|
||||
}
|
||||
}
|
||||
if currentEvent == "message_stop" {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return scanner.Err()
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user