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 返回该用户所有已批准且供应商/模型仍启用的模型。 // decided_at 兜底 updated_at:历史数据/直接落库的批准记录可能为空。 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,COALESCE(max(r.decided_at),max(r.updated_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,COALESCE(r.decided_at,r.updated_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 COALESCE(r.decided_at,r.updated_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() } // appendChatMessages 在会话上批量追加消息(user+assistant 一轮):单事务内 // 连续插入、序列号一次锁定一次递增,模型调用成功后才落库——要么整轮落库 // 要么整轮不落,客户端重试不会产生孤儿或重复消息。 func (s *Service) appendChatMessages(ctx context.Context, sessionID string, messages []ConversationMessage) ([]ConversationMessage, error) { if len(messages) == 0 { return nil, nil } tx, err := s.pool.Begin(ctx) if err != nil { return nil, 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 nil, err } if sequence+len(messages)-1 > 200 { return nil, 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 nil, err } } firstTitle := messages[0].Content out := make([]ConversationMessage, 0, len(messages)) for _, message := range messages { id, err := platformid.NewUUID() if err != nil { return nil, err } hash := messageDigest(previous, sequence, message.Role, message.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, message.Role, message.Content, previous, hash).Scan(&created); err != nil { return nil, err } out = append(out, ConversationMessage{Sequence: sequence, Role: message.Role, Content: message.Content, CreatedAt: created}) previous = hash sequence++ } _, err = tx.Exec(ctx, `UPDATE gateway.portal_chat_sessions SET next_sequence=$2,title=CASE WHEN next_sequence=1 THEN left($3,160) ELSE title END,updated_at=clock_timestamp() WHERE id=$1`, sessionID, sequence, firstTitle) if err != nil { return nil, err } if err = tx.Commit(ctx); err != nil { return nil, err } return out, 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 } // chatStreamWriter 把网关响应实时转发给客户端,同时累积完整响应体: // - 2xx 且 text/event-stream:透传模式,边写边 flush,供提取流式回答; // - 其他(错误 JSON、或上游忽略 stream 返回普通 JSON):缓冲模式,header // 不落盘,由调用方决定输出方式或转成统一错误。 type chatStreamWriter struct { w http.ResponseWriter buf bytes.Buffer mode int // 0 未知 / 1 透传 / 2 缓冲 code int // shown 记录是否已向客户端落 header(透传模式才落)。 shown bool } func (c *chatStreamWriter) Header() http.Header { return c.w.Header() } func (c *chatStreamWriter) WriteHeader(code int) { c.code = code if code >= 200 && code < 300 && strings.Contains(c.w.Header().Get("Content-Type"), "text/event-stream") { c.mode = 1 c.shown = true c.w.WriteHeader(code) return } // 非流式(错误 JSON 或普通 JSON 响应):缓冲模式,header 不落盘。 c.mode = 2 } func (c *chatStreamWriter) Write(p []byte) (int, error) { if c.mode == 0 { c.WriteHeader(http.StatusOK) } c.buf.Write(p) if c.mode != 1 { return len(p), nil } n, err := c.w.Write(p) c.flush() return n, err } func (c *chatStreamWriter) Flush() { c.flush() } func (c *chatStreamWriter) flush() { if c.mode != 1 || !c.shown { return } if f, ok := c.w.(http.Flusher); ok { f.Flush() } } // extractStreamAnswer 从 OpenAI 兼容 SSE 响应中拼接完整回答(delta/message 的 // content 逐段累积)。 func extractStreamAnswer(body []byte) string { var answer strings.Builder for _, line := range strings.Split(string(body), "\n") { trimmed := strings.TrimSpace(line) if !strings.HasPrefix(trimmed, "data:") { continue } payload := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) if payload == "" || payload == "[DONE]" { continue } var chunk struct { Choices []struct { Delta struct{ Content string `json:"content"` } `json:"delta"` Message *struct{ Content string `json:"content"` } `json:"message"` } `json:"choices"` } if json.Unmarshal([]byte(payload), &chunk) != nil { continue } for _, choice := range chunk.Choices { if choice.Delta.Content != "" { answer.WriteString(choice.Delta.Content) } else if choice.Message != nil { answer.WriteString(choice.Message.Content) } } } return answer.String() } // extractNonStreamAnswer 从普通 JSON 响应中提取回答(上游忽略 stream 参数时)。 func extractNonStreamAnswer(body []byte) string { var response struct { Choices []struct { Message *struct{ Content string `json:"content"` } `json:"message"` } `json:"choices"` } if json.Unmarshal(body, &response) != nil || len(response.Choices) == 0 || response.Choices[0].Message == nil { return "" } return response.Choices[0].Message.Content } // chatErrorFromBody 从网关错误 JSON 中提取可读错误信息。 func chatErrorFromBody(body []byte, code int) error { var response struct { Error struct{ Message string `json:"message"` } `json:"error"` } message := fmt.Sprintf("模型调用失败(HTTP %d)", code) if json.Unmarshal(body, &response) == nil && response.Error.Message != "" { message = response.Error.Message } return errors.New(message) } // AppendChatMessageStream 在会话上追加一轮对话,网关响应以 SSE 实时透传给 // 客户端,流结束后把 user+assistant 整轮落库(失败不落库)。上游忽略 stream // 返回普通 JSON 时自动转成 SSE 事件,前端统一走流式解析。 // 缓冲模式(未落 header)的错误直接返回,由 handler 走统一错误处理;透传模式 // 已向客户端输出 200+SSE 头,错误只能通过 SSE error 事件告知。 func (s *Service) AppendChatMessageStream(ctx context.Context, w http.ResponseWriter, account identity.Account, id, message string) error { message = strings.TrimSpace(message) if message == "" || len(message) > 100000 { return 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= 300 { return chatErrorFromBody(stream.buf.Bytes(), stream.code) } answer := extractNonStreamAnswer(stream.buf.Bytes()) if strings.TrimSpace(answer) == "" { return errors.New("模型未返回文本回答") } if _, err = s.appendChatMessages(context.WithoutCancel(ctx), id, []ConversationMessage{{Role: "user", Content: message}, {Role: "assistant", Content: answer}}); err != nil { return err } w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") w.WriteHeader(http.StatusOK) writeDone(map[string]any{"choices": []map[string]any{{"delta": map[string]any{"content": answer}}}}) return nil } // 透传模式:已向客户端输出 200 + SSE 头,错误只能通过事件告知。 answer := extractStreamAnswer(stream.buf.Bytes()) if strings.TrimSpace(answer) == "" { writeDone(map[string]any{"error": map[string]string{"message": "模型未返回文本回答"}}) return nil } if _, err = s.appendChatMessages(context.WithoutCancel(ctx), id, []ConversationMessage{{Role: "user", Content: message}, {Role: "assistant", Content: answer}}); err != nil { writeDone(map[string]any{"error": map[string]string{"message": err.Error()}}) return nil } writeDone(map[string]any{"done": true, "conversation_id": id}) return nil } // 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