Files
ai-gateway-go/internal/portal/chat.go
T
LLMGuardX Dev 87c2b04174 0.11.3: 旗舰版第四轮完善(统一审批中心/工具治理/平台环境变量/数字员工入口/个人渠道/报表多维/租户配额)
- 统一审批中心:模型/资源/渠道/工具四类申请聚合审批,通过自动开通
  (marketplace 安装/渠道授权),outbox 双向站内信;门户可发起/撤回。
- 工具治理:rate_limit_rpm(固定窗口原子 upsert,多实例共享)+ approval_required
  (首次调用自动发起审批,批准前一律拒绝)。
- 平台环境变量:平台级注入 skill/MCP 运行时,个人可覆盖;系统管理员可写。
- 数字员工会话入口:门户列表/对话/调用记录,复用用户运行时凭据。
- 个人渠道:webhook 入站令牌 SHA-256 摘要 + constant-time 校验,绑定已批准
  模型,用量归属用户 Key。
- 报表多维:工具调用/审批授权/安全事件三组统计端点与页面。
- 租户配额:部门 Key/月 Token 上限,运行时凭据开通强制校验,概览展示用量。
- 迁移 000042-000045;修复渠道空 API Key NOT NULL 违约与 inet 扫描;
  25 包测试通过,前后端构建通过,端到端验证完成。
2026-08-13 13:41:22 +08:00

330 lines
14 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package portal
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"time"
"aigateway.local/core/internal/gateway"
"aigateway.local/core/internal/identity"
platformid "aigateway.local/core/internal/platform/id"
"github.com/jackc/pgx/v5"
)
// ChatModel 是门户用户经审批可用的模型。
type ChatModel struct {
ProviderCode string `json:"provider_code"`
Model string `json:"model"`
ApprovedAt time.Time `json:"approved_at"`
}
// ChatModels 返回该用户所有已批准且供应商/模型仍启用的模型。
func (s *Service) ChatModels(ctx context.Context, account identity.Account) ([]ChatModel, error) {
rows, err := s.pool.Query(ctx, `SELECT DISTINCT r.provider_code,r.model,max(r.decided_at)
FROM gateway.model_access_requests r
JOIN gateway.providers p ON p.code=r.provider_code AND p.enabled
JOIN gateway.provider_models m ON m.provider_id=p.id AND m.provider_model_id=r.model AND m.enabled
WHERE r.portal_user_id=$1 AND r.status='approved'
GROUP BY r.provider_code,r.model ORDER BY r.provider_code,r.model`, account.ID)
if err != nil {
return nil, err
}
defer rows.Close()
items := []ChatModel{}
for rows.Next() {
var item ChatModel
if err := rows.Scan(&item.ProviderCode, &item.Model, &item.ApprovedAt); err != nil {
return nil, err
}
items = append(items, item)
}
return items, rows.Err()
}
// approvedModel 校验模型是否在该用户的已批准清单内。
func (s *Service) approvedModel(ctx context.Context, account identity.Account, providerCode, model string) (ChatModel, error) {
providerCode = strings.ToLower(strings.TrimSpace(providerCode))
model = strings.TrimSpace(model)
var item ChatModel
err := s.pool.QueryRow(ctx, `SELECT r.provider_code,r.model,r.decided_at
FROM gateway.model_access_requests r
JOIN gateway.providers p ON p.code=r.provider_code AND p.enabled
JOIN gateway.provider_models m ON m.provider_id=p.id AND m.provider_model_id=r.model AND m.enabled
WHERE r.portal_user_id=$1 AND r.provider_code=$2 AND r.model=$3 AND r.status='approved'
ORDER BY r.decided_at DESC LIMIT 1`, account.ID, providerCode, model).Scan(&item.ProviderCode, &item.Model, &item.ApprovedAt)
if errors.Is(err, pgx.ErrNoRows) {
return ChatModel{}, ErrNotFound
}
return item, err
}
// ensureChatCredential 为用户开通/复用聊天运行时凭据,限额取已批准申请的最大值。
func (s *Service) ensureChatCredential(ctx context.Context, account identity.Account) (string, error) {
if s.credentials == nil || s.runtime == nil {
return "", errors.New("聊天服务未配置")
}
var rpm int
var monthlyTokens int64
if err := s.pool.QueryRow(ctx, `SELECT COALESCE(max(requested_rpm),0),COALESCE(max(requested_monthly_tokens),0) FROM gateway.model_access_requests WHERE portal_user_id=$1 AND status='approved'`, account.ID).Scan(&rpm, &monthlyTokens); err != nil {
return "", err
}
secret, _, err := s.credentials.EnsureUser(ctx, account.ID, account.DepartmentID, rpm, monthlyTokens)
return secret, err
}
// ChatSession 是一条通用聊天会话。
type ChatSession struct {
ID string `json:"id"`
Title string `json:"title"`
ProviderCode string `json:"provider_code"`
Model string `json:"model"`
Status string `json:"status"`
Messages []ConversationMessage `json:"messages,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
const chatSessionSelect = `SELECT id::text,title,provider_code,model,status,created_at,updated_at FROM gateway.portal_chat_sessions`
func (s *Service) ListChatSessions(ctx context.Context, account identity.Account, limit int) ([]ChatSession, error) {
if limit < 1 || limit > 200 {
limit = 50
}
rows, err := s.pool.Query(ctx, chatSessionSelect+` WHERE portal_user_id=$1 AND status='active' ORDER BY updated_at DESC LIMIT $2`, account.ID, limit)
if err != nil {
return nil, err
}
defer rows.Close()
items := []ChatSession{}
for rows.Next() {
var item ChatSession
if err := rows.Scan(&item.ID, &item.Title, &item.ProviderCode, &item.Model, &item.Status, &item.CreatedAt, &item.UpdatedAt); err != nil {
return nil, err
}
items = append(items, item)
}
return items, rows.Err()
}
func (s *Service) CreateChatSession(ctx context.Context, account identity.Account, providerCode, model string) (ChatSession, error) {
if _, err := s.approvedModel(ctx, account, providerCode, model); err != nil {
return ChatSession{}, ErrNotFound
}
if _, err := s.ensureChatCredential(ctx, account); err != nil {
return ChatSession{}, err
}
id, err := platformid.NewUUID()
if err != nil {
return ChatSession{}, err
}
var item ChatSession
err = s.pool.QueryRow(ctx, `INSERT INTO gateway.portal_chat_sessions(id,portal_user_id,provider_code,model) VALUES($1,$2,$3,$4) RETURNING id::text,'',provider_code,model,status,created_at,updated_at`, id, account.ID, providerCode, model).Scan(&item.ID, &item.Title, &item.ProviderCode, &item.Model, &item.Status, &item.CreatedAt, &item.UpdatedAt)
item.Messages = []ConversationMessage{}
return item, err
}
func (s *Service) RenameChatSession(ctx context.Context, account identity.Account, id, title string) (ChatSession, error) {
title = strings.TrimSpace(title)
if title == "" || len(title) > 128 {
return ChatSession{}, errors.New("会话标题必须为 1-128 个字符")
}
var item ChatSession
err := s.pool.QueryRow(ctx, `UPDATE gateway.portal_chat_sessions SET title=$3,updated_at=clock_timestamp() WHERE id=$1 AND portal_user_id=$2 AND status='active' RETURNING id::text,title,provider_code,model,status,created_at,updated_at`, id, account.ID, title).Scan(&item.ID, &item.Title, &item.ProviderCode, &item.Model, &item.Status, &item.CreatedAt, &item.UpdatedAt)
if errors.Is(err, pgx.ErrNoRows) {
return ChatSession{}, ErrNotFound
}
return item, err
}
func (s *Service) DeleteChatSession(ctx context.Context, account identity.Account, id string) error {
tag, err := s.pool.Exec(ctx, `UPDATE gateway.portal_chat_sessions SET status='archived' WHERE id=$1 AND portal_user_id=$2`, id, account.ID)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrNotFound
}
return nil
}
// ChatSession 返回会话与全部消息(哈希链完整性校验)。
func (s *Service) ChatSession(ctx context.Context, account identity.Account, id string) (ChatSession, error) {
var item ChatSession
err := s.pool.QueryRow(ctx, chatSessionSelect+` WHERE id=$1 AND portal_user_id=$2`, id, account.ID).Scan(&item.ID, &item.Title, &item.ProviderCode, &item.Model, &item.Status, &item.CreatedAt, &item.UpdatedAt)
if errors.Is(err, pgx.ErrNoRows) {
return ChatSession{}, ErrNotFound
}
if err != nil {
return ChatSession{}, err
}
rows, err := s.pool.Query(ctx, `SELECT sequence,role,content,previous_hash,message_hash,created_at FROM gateway.portal_chat_messages WHERE session_id=$1 ORDER BY sequence`, id)
if err != nil {
return ChatSession{}, err
}
defer rows.Close()
previous := strings.Repeat("0", 64)
item.Messages = []ConversationMessage{}
for rows.Next() {
var message ConversationMessage
var storedPrevious, storedHash string
if err = rows.Scan(&message.Sequence, &message.Role, &message.Content, &storedPrevious, &storedHash, &message.CreatedAt); err != nil {
return ChatSession{}, err
}
if storedPrevious != previous || storedHash != messageDigest(previous, message.Sequence, message.Role, message.Content) {
return ChatSession{}, errors.New("会话历史完整性校验失败")
}
previous = storedHash
item.Messages = append(item.Messages, message)
}
return item, rows.Err()
}
// appendChatMessage 在会话上追加一条消息(哈希链 + 序号,事务内完成)。
func (s *Service) appendChatMessage(ctx context.Context, sessionID, role, content string) (ConversationMessage, error) {
tx, err := s.pool.Begin(ctx)
if err != nil {
return ConversationMessage{}, err
}
defer func() { _ = tx.Rollback(ctx) }()
var sequence int
if err = tx.QueryRow(ctx, `SELECT next_sequence FROM gateway.portal_chat_sessions WHERE id=$1 FOR UPDATE`, sessionID).Scan(&sequence); err != nil {
return ConversationMessage{}, err
}
if sequence > 200 {
return ConversationMessage{}, errors.New("本会话已达到 200 条消息上限")
}
previous := strings.Repeat("0", 64)
if sequence > 1 {
if err = tx.QueryRow(ctx, `SELECT message_hash FROM gateway.portal_chat_messages WHERE session_id=$1 AND sequence=$2`, sessionID, sequence-1).Scan(&previous); err != nil {
return ConversationMessage{}, err
}
}
id, err := platformid.NewUUID()
if err != nil {
return ConversationMessage{}, err
}
hash := messageDigest(previous, sequence, role, content)
var created time.Time
if err = tx.QueryRow(ctx, `INSERT INTO gateway.portal_chat_messages(id,session_id,sequence,role,content,previous_hash,message_hash) VALUES($1,$2,$3,$4,$5,$6,$7) RETURNING created_at`, id, sessionID, sequence, role, content, previous, hash).Scan(&created); err != nil {
return ConversationMessage{}, err
}
_, err = tx.Exec(ctx, `UPDATE gateway.portal_chat_sessions SET next_sequence=next_sequence+1,title=CASE WHEN next_sequence=1 THEN left($2,160) ELSE title END,updated_at=clock_timestamp() WHERE id=$1`, sessionID, content)
if err != nil {
return ConversationMessage{}, err
}
if err = tx.Commit(ctx); err != nil {
return ConversationMessage{}, err
}
return ConversationMessage{Sequence: sequence, Role: role, Content: content, CreatedAt: created}, nil
}
// callChat 用用户的运行时凭据直接调用受管网关 /v1/chat/completions。
// 认证、限流、配额、审计与路由都由网关统一执行,与外部 API Key 调用完全同权。
func (s *Service) callChat(ctx context.Context, secret, providerCode, model string, messages []ConversationMessage) (map[string]any, string, error) {
payloadMessages := make([]map[string]any, 0, len(messages))
for _, m := range messages {
payloadMessages = append(payloadMessages, map[string]any{"role": m.Role, "content": m.Content})
}
payload, _ := json.Marshal(map[string]any{"model": model, "messages": payloadMessages, "stream": false})
request := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(payload)).WithContext(gateway.WithRequestID(ctx, "portal-chat-"+time.Now().UTC().Format("20060102150405.000000000")))
request.Header.Set("Authorization", "Bearer "+secret)
request.Header.Set("Content-Type", "application/json")
recorder := httptest.NewRecorder()
s.gateway.ServeHTTP(recorder, request)
var response map[string]any
if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil {
return nil, "", errors.New("模型响应无法解析")
}
if recorder.Code < 200 || recorder.Code >= 300 {
message := fmt.Sprintf("模型调用失败(HTTP %d", recorder.Code)
if value, ok := response["error"].(map[string]any); ok {
if text, ok := value["message"].(string); ok {
message = text
}
}
return response, "", errors.New(message)
}
choices, _ := response["choices"].([]any)
if len(choices) == 0 {
return response, "", errors.New("模型未返回回答")
}
choice, _ := choices[0].(map[string]any)
message, _ := choice["message"].(map[string]any)
answer, _ := message["content"].(string)
if strings.TrimSpace(answer) == "" {
return response, "", errors.New("模型未返回文本回答")
}
return response, answer, nil
}
// ChatOnce 一次性对话(不落库):模型须已批准,凭据自动开通。
func (s *Service) ChatOnce(ctx context.Context, account identity.Account, providerCode, model, message string) (map[string]any, error) {
if _, err := s.approvedModel(ctx, account, providerCode, model); err != nil {
return nil, ErrNotFound
}
message = strings.TrimSpace(message)
if message == "" || len(message) > 100000 {
return nil, errors.New("消息为空或过长")
}
secret, err := s.ensureChatCredential(ctx, account)
if err != nil {
return nil, err
}
response, _, err := s.callChat(ctx, secret, providerCode, model, []ConversationMessage{{Role: "user", Content: message}})
return response, err
}
// AppendChatMessage 在会话上追加一轮对话:busy 租约防并发交错,消息在模型
// 调用成功后才落库,失败重试不会产生孤儿消息或重复消息。
func (s *Service) AppendChatMessage(ctx context.Context, account identity.Account, id, message string) (map[string]any, error) {
message = strings.TrimSpace(message)
if message == "" || len(message) > 100000 {
return nil, errors.New("消息为空或过长")
}
lease, _ := platformid.NewUUID()
tag, err := s.pool.Exec(ctx, `UPDATE gateway.portal_chat_sessions SET busy=true,busy_token=$3,busy_since=clock_timestamp() WHERE id=$1 AND portal_user_id=$2 AND status='active' AND (NOT busy OR busy_since<clock_timestamp()-interval '10 minutes')`, id, account.ID, lease)
if err != nil {
return nil, err
}
if tag.RowsAffected() == 0 {
return nil, errors.New("会话不存在、已归档或上一条消息仍在处理")
}
defer func() {
_, _ = s.pool.Exec(context.WithoutCancel(ctx), `UPDATE gateway.portal_chat_sessions SET busy=false,busy_token=NULL,busy_since=NULL WHERE id=$1 AND busy_token=$2`, id, lease)
}()
var providerCode, model string
if err = s.pool.QueryRow(ctx, `SELECT provider_code,model FROM gateway.portal_chat_sessions WHERE id=$1`, id).Scan(&providerCode, &model); err != nil {
return nil, err
}
conversation, err := s.ChatSession(ctx, account, id)
if err != nil {
return nil, err
}
secret, err := s.ensureChatCredential(ctx, account)
if err != nil {
return nil, err
}
history := make([]ConversationMessage, len(conversation.Messages)+1)
copy(history, conversation.Messages)
history[len(conversation.Messages)] = ConversationMessage{Role: "user", Content: message}
response, answer, err := s.callChat(ctx, secret, providerCode, model, history)
if err != nil {
return response, err
}
if _, err = s.appendChatMessage(ctx, id, "user", message); err != nil {
return nil, err
}
if _, err = s.appendChatMessage(ctx, id, "assistant", answer); err != nil {
return nil, err
}
response["conversation_id"] = id
response["history_integrity"] = "verified"
return response, nil
}