diff --git a/im/backend/docs/ai-agent.md b/im/backend/docs/ai-agent.md new file mode 100644 index 0000000..a07a8e8 --- /dev/null +++ b/im/backend/docs/ai-agent.md @@ -0,0 +1,140 @@ +# 测试账号 AI 托管 + +测试账号(`users.is_test=1`)在收到真实用户的私信后,由后台配置的大模型生成回复并以该账号身份发出。目标是让新用户在冷启动阶段能得到真实的对话反馈,同时把模型接入做成后台可配置、可多协议共存的能力。 + +本文只描述设计与实施边界,未实现的部分标注为阶段目标。 + +## 协议抽象 + +模型厂商众多但线上协议只有少数几种,因此按**协议**而不是按厂商建适配器,与现有短信(`sms_providers.go`)和对象存储(`storage_providers.go`)的做法一致: + +| 协议 | 覆盖范围 | 关键差异 | +| --- | --- | --- | +| `openai` | OpenAI、DeepSeek、通义千问兼容模式、智谱、Moonshot、火山方舟、SiliconFlow、Ollama、vLLM、one-api/new-api 网关 | `POST {base_url}/chat/completions`,`Authorization: Bearer`,system 作为首条 message | +| `anthropic` | Claude | `POST {base_url}/v1/messages`,`x-api-key` + `anthropic-version`,system 是顶层字段 | +| `gemini` | Google Gemini | `POST {base_url}/v1beta/models/{model}:generateContent`,`contents[].parts[].text`,assistant 角色写作 `model`;密钥用 `x-goog-api-key` 头而不是查询参数,避免写进代理日志 | +| `webhook` | Coze、Dify、n8n、自建编排 | 后端 POST 规范化 JSON,期望返回 `{"text": "..."}`,不为每个编排平台写代码 | +| `debug` | 开发与测试 | 不出网,返回确定性文案,供本地联调和单元测试使用 | + +统一接口(`internal/app/ai_providers.go`): + +```go +type aiChatMessage struct{ Role, Text string } // role: user | assistant +type aiChatRequest struct { + System string + Messages []aiChatMessage + MaxTokens int + Temperature float64 +} +type aiChatResult struct { + Text string + PromptTokens, CompletionTokens int + FinishReason, RequestID string +} +type aiProvider interface { + Chat(ctx context.Context, req aiChatRequest) (aiChatResult, error) +} +``` + +新增一家厂商时,若走 OpenAI 兼容格式则只需在后台新增一条模型记录,不改代码。 + +## 数据模型 + +`migrations/030_ai_models.sql`(P0,已落地): + +- `ai_models`:`id / name / protocol / base_url / model_name / api_key_cipher / temperature / max_tokens / timeout_ms / status / is_default / extra_json`。密钥复用 `integration.go` 的 `encryptSecret`/`decryptSecret`(AES-GCM),与 `system_configs` 的 secret 字段同一套密钥。允许同时存在多条不同协议的记录,`is_default` 在写入时互斥。 +- `ai_call_logs`:`model_id / scene / agent_user_id / conversation_id / latency_ms / prompt_tokens / completion_tokens / status / error / prompt_digest / reply_text`,用于成本核算和事后审核,按 `ai.log_retention_days` 定期清理。后台探活也会写入(`scene='probe'`)。 + +P1 迁移新增: + +- `ai_agents`:`user_id`(唯一,必须 `is_test=1`)/ `model_id` / `persona`(系统提示词)/ `enabled` / `daily_reply_limit` / `reply_delay_min_ms` / `reply_delay_max_ms`。未单独配置 persona 时由 `user_profiles` 的昵称、年龄、城市、bio、标签拼装,保证人设和资料一致。 +- `ai_reply_jobs`:`id / conversation_id / agent_user_id / trigger_message_id / status(0待处理 1处理中 2成功 3失败) / attempts / run_after / claim_token / last_error`。同一会话同时只允许一个待回复任务:`pending_key` 是 `IF(status=0, conversation_id, NULL)` 的 STORED 生成列,加唯一键后,待处理状态天然互斥,离开该状态即为 NULL 不再约束。 + +全局开关放进 `integrationSpecs` 新增的 `ai` 组(`system_configs` 表),后台现成的通用面板即可渲染:`ai.enabled`、`ai.default_model_id`、`ai.max_concurrency`、`ai.daily_call_limit`、`ai.reply_delay_min/max_ms`、`ai.allow_batches`(限定 `test_batch`)、`ai.disclose_in_chat`、`ai.log_retention_days`、`ai.fallback_text`。 + +## 托管流程 + +1. `persistMessage` 成功写库并广播后,判断该会话对端是否为已启用的 AI 托管测试账号;是则写入 `ai_reply_jobs`(`run_after = now + rand(delay_min, delay_max)`)。判定和入队都不阻塞发送方的响应。 +2. 同一会话已有待处理任务时不新建,只把 `run_after` 后移(上限 15 秒)并更新 `trigger_message_id`。用户连发三条会得到一条回复,而不是三条。 +3. `Run()` 中新增 worker goroutine(与现有 `processDueAccountClosures` 同样的 ticker 模式),每秒领取到期任务。领取方式是一条 `UPDATE ... SET status=1,claim_token=? WHERE status=0 AND run_after<=NOW(3) ORDER BY run_after LIMIT n` 再按 token 回读:每行只能离开待处理状态一次,多实例不会重复消费。**没有用 `SELECT ... FOR UPDATE SKIP LOCKED`**——生产 compose 是 MySQL 8.4,但开发环境仍在 5.7,那里该语法直接解析失败。 +4. 组装上下文:取该会话最近 20 条消息,对端消息映射为 `user`、托管账号自己的消息映射为 `assistant`;图片和语音降级为 `[图片]` `[语音]` 占位符。系统提示词 = persona + 硬性约束(中文、口语化、不超过 40 字、不加 emoji 堆砌、不索取联系方式、不涉及金钱和线下见面、不输出链接)。 +5. 调用模型(超时取 `ai_models.timeout_ms`,失败重试一次并退避)。成功后清洗输出(去 Markdown、截断、过滤空回复),走 `persistMessageContext` 以托管账号身份发送。 +6. 回复经过与真人完全相同的落库和推送路径,因此未读数、WebSocket 推送、离线推送、会话列表排序全部复用现有逻辑,**客户端零改动**。 + +## 闸门 + +- `ai.enabled` 为总开关;生产环境禁止使用 `debug` 协议(与 `validateSMSProviderConfig` 中"生产环境禁止 debug 短信提供商"同一规则)。 +- 仅对 `is_test=1` 且在 `ai.allow_batches` 内的账号生效;AI 账号之间不互相触发,避免自激对话。 +- 三层配额:单账号每日回复上限、全局每日调用上限、并发信号量。任一超限即丢弃任务并记日志,不做排队堆积。 +- 发送前复查会话最后一条消息,若已是托管账号自己发的则跳过,避免重复回复。 +- `persistMessage` 已有的禁言检查(`isSanctionActive`)和媒体归属校验对 AI 同样生效。发送频控要用独立限流键,避免托管账号触发按人设计的"每日活跃聊天上限"(`dailyActiveChatLimitError`)。 +- 模型调用失败时默认静默,不发兜底话术;`ai.fallback_text` 非空时才发送,避免机械感。 + +## 后台 + +- `GET/POST/PUT/DELETE /admin/v1/ai/models`、`POST /admin/v1/ai/models/:id/test`(发一条探活提示词,返回延迟、tokens 和实际回复)。 +- `GET/PUT /admin/v1/ai/agents`:列出测试账号及其托管状态,支持按批次批量开关和批量绑定模型。 +- `GET /admin/v1/ai/logs`:调用日志,可按模型、场景、结果和时间范围筛选。 +- `GET /admin/v1/ai/usage`:今日用量、近 7 天按日与按模型汇总、任务队列状态、最近错误,以及服务端计算的告警列表。 +- `GET/PUT /admin/v1/integrations/ai`:全局开关,复用现有通用配置面板。 + +权限沿用 `permit("system:manage", ...)`;模型和托管开关的变更走 `a.audit(...)` 落审计日志。管理端前端只需新增一个 `ai-config.vue`(复用 `IntegrationConfigPanel`)、一个模型列表页,并把 `socialApi.integration` 的 group 联合类型加上 `'ai'`。 + +## 需要改动的既有代码 + +- `persistMessage(r *http.Request, ...)` 抽出 `persistMessageContext(ctx context.Context, ...)`,HTTP 处理器保留薄封装。worker 没有 `*http.Request`,这是唯一侵入式改造。 +- `app.go` 的 `Run()` 增加 worker goroutine 和退出时的优雅停止。 +- `integrationSpecs` 增加 `ai` 组,`adminTestIntegration` 增加 `ai` 分支。 + +## 分期 + +1. **P0 骨架(已完成)**:迁移、`ai_models` CRUD、五种协议适配器、后台配置页与探活按钮。不接 IM,可独立验证。 +2. **P1 托管闭环(已完成)**:`ai_agents` 绑定、任务表与 worker、`persistMessageContext` 重构、延迟与合并。默认 `ai.enabled=false` 发布。 +3. **P2 补齐(已完成)**:调用日志页、用量汇总与告警、日志与任务清理。 +4. **P3 可选**:AI 主动打招呼、动态与评论生成、多模型灰度路由。 + +## 风险与待定 + +- **合规**:测试账号资料页已标注"AI生成虚拟资料",但会话中是否显式标注由 `ai.disclose_in_chat` 控制,默认开启。对真实用户隐瞒机器人身份在部分地区属于监管红线,关闭前需要产品确认。 +- **成本**:单条回复约 500~1500 tokens。上线前应先用 `debug` 协议跑通链路,再用真实模型小批量放量,靠 `ai_call_logs` 观察日用量。 +- **数据留存**:`ai_call_logs.reply_text` 会保存对话内容,属于用户数据。默认只存托管账号自己的回复和用户消息摘要(`prompt_digest`),完整 prompt 不落库。 + +## P0 落地记录 + +已实现的文件: + +- `migrations/030_ai_models.sql`:两张表和 11 个 `ai.*` 配置项。 +- `internal/app/ai_providers.go`:`aiProvider` 接口与五个协议适配器,含地址校验(禁止账号、查询参数、片段,生产环境强制 HTTPS)和 512 KiB 响应上限。 +- `internal/app/admin_ai.go`:模型 CRUD、探活、调用日志写入、`validateAIConfig`。 +- `internal/app/integration.go`:`ai` 配置分组与 `POST /admin/v1/integrations/ai/test`(用默认模型真实调用一次)。 +- 管理端:`ai-config.vue`(复用通用配置面板)、`ai-models.vue`(模型表格、编辑弹窗、测试按钮)、两条路由。 + +验证方式:`internal/app/ai_providers_test.go` 用 `httptest` 覆盖四种线协议的请求组装与响应解析(system 提示词位置、鉴权头、assistant→model 角色映射、错误体与非 JSON 响应、密钥不进查询参数),以及地址校验、工厂守卫和 debug 协议的确定性。本地起服务后走通了「登录→建模型→探活→改配置→连通性测试→删除」,`ai_call_logs` 与 `admin_audit_logs` 均有对应记录。 + +P0 不改动任何既有业务路径:`persistMessage` 尚未重构,IM 流程没有接入点,总开关默认关闭。 + +## P1 落地记录 + +新增 `migrations/031_ai_agents.sql`(`ai_agents` + `ai_reply_jobs`)、`internal/app/ai_agent.go`(入队、worker、上下文组装、提示词、清洗、配额)、`internal/app/admin_ai_agents.go`(账号绑定与任务查询),并把 `persistMessage` 抽成 `persistMessageContext`。管理端新增 `ai-agents.vue`。 + +真库联调结果(本地 MySQL 5.7 + debug 协议): + +- 真实用户连发三条消息 → **只入队一个任务**(`trigger_message_id` 指向最后一条)→ 托管账号回一条,`ai_reply_jobs.status=2`、`ai_call_logs` 有 `scene='chat'` 记录、`ai_agents.reply_count` 递增。 +- 新会话里托管账号的第一条回复带 `[AI 助手]` 前缀(`ai.disclose_in_chat` 控制)。 +- 托管账号之间互发消息不入队;关闭绑定后不再入队。 + +一处设计修正:原计划用 `SELECT ... FOR UPDATE SKIP LOCKED` 领取任务,实测开发环境的 MySQL 5.7.26 直接报 `ERROR 1064 ... near 'SKIP LOCKED'`,已改为 claim_token 方案,5.7 与 8.x 都可用。 + +**尚未接入**:P2 的调用日志页、配额告警和日志清理任务;`ai.fallback_text` 已实现但默认留空(失败静默)。 + +## P2 落地记录 + +新增 `internal/app/admin_ai_logs.go`(日志查询、用量汇总、告警计算、清理任务)与管理端 `ai-logs.vue`。 + +**告警**在服务端计算(`aiUsageAlerts`),前端只负责显示,这样以后接入通知渠道时判定逻辑只有一份。触发条件:总开关已开但无可用模型(error)、没有任何账号开启托管(warning)、今日调用达到或超过配额 80%(warning/error)、当日调用 ≥5 次且失败率 >30%(error,附最近一条错误)、有任务超过 60 秒未处理即 worker 可能没在跑(warning)。总开关关闭时不产生任何告警。 + +**清理**由 `runAIMaintenance` 每小时执行一次,进程启动时也先跑一遍:按 `ai.log_retention_days` 删除过期调用日志、删除 7 天前的已完成/失败任务、把卡在「处理中」超过 5 分钟的任务标记为失败(进程在回复中途被杀时不会永久占用会话的待处理名额)。删除都按 1000 条一批,避免长事务。 + +实测确认:保留 30 天时,40 天和 31 天前的日志被清理、29 天前的保留;8 天前的已完成任务被清理、2 天前的失败任务保留;卡住 10 分钟的任务被标记为「处理超时,任务已中断」。用量接口在造出 7 次调用(含 1 次失败、耗时非零)后返回正确的今日汇总、按模型汇总和最近错误;非法筛选参数(`days=999`、`scene=hack`)返回 400。 + +一处实现修正:`AVG(latency_ms)` 在 MySQL 里返回 DECIMAL,扫进 `sql.NullInt64` 会失败——按模型汇总因此每行都被静默跳过、今日平均耗时恒为 0。已改为 `sql.NullFloat64` 后取整。 \ No newline at end of file diff --git a/im/backend/internal/app/admin_ai.go b/im/backend/internal/app/admin_ai.go new file mode 100644 index 0000000..8d089f3 --- /dev/null +++ b/im/backend/internal/app/admin_ai.go @@ -0,0 +1,424 @@ +package app + +import ( + "context" + "crypto/sha256" + "database/sql" + "encoding/hex" + "errors" + "fmt" + "net/http" + "strconv" + "strings" + "time" + + "github.com/go-sql-driver/mysql" +) + +type aiModelRequest struct { + Name string `json:"name"` + Protocol string `json:"protocol"` + BaseURL string `json:"baseUrl"` + ModelName string `json:"modelName"` + APIKey string `json:"apiKey"` + ClearAPIKey bool `json:"clearApiKey"` + Temperature *float64 `json:"temperature"` + MaxTokens *int `json:"maxTokens"` + TimeoutMS *int `json:"timeoutMs"` + Status *int `json:"status"` + IsDefault *bool `json:"isDefault"` +} + +const aiModelColumns = `id,name,protocol,base_url,model_name,api_key_cipher,temperature,max_tokens,timeout_ms,status,is_default,created_at,updated_at` + +func (a *App) scanAIModel(scan func(dest ...any) error) (aiModel, error) { + var model aiModel + var cipher string + var created, updated time.Time + if err := scan(&model.ID, &model.Name, &model.Protocol, &model.BaseURL, &model.ModelName, &cipher, + &model.Temperature, &model.MaxTokens, &model.TimeoutMS, &model.Status, &model.IsDefault, &created, &updated); err != nil { + return aiModel{}, err + } + model.HasAPIKey = cipher != "" + if cipher != "" { + plain, err := a.decryptSecret(cipher) + if err != nil { + // A key that cannot be decrypted (rotated encryption key) must not + // silently become an empty key that reaches the provider. + return model, fmt.Errorf("模型 %s 的密钥无法解密,请重新填写", model.Name) + } + model.APIKey = plain + } + model.CreatedAt = created.Format(time.RFC3339) + model.UpdatedAt = updated.Format(time.RFC3339) + return model, nil +} + +func (a *App) loadAIModel(ctx context.Context, id int64) (aiModel, error) { + row := a.db.QueryRowContext(ctx, `SELECT `+aiModelColumns+` FROM ai_models WHERE id=?`, id) + return a.scanAIModel(row.Scan) +} + +// The default model is what the reply worker uses when an agent has no explicit +// binding: the configured id first, then the flagged row. +func (a *App) defaultAIModel(ctx context.Context) (aiModel, error) { + if configured := strings.TrimSpace(a.configPlain(ctx, "ai.default_model_id", "")); configured != "" { + row := a.db.QueryRowContext(ctx, `SELECT `+aiModelColumns+` FROM ai_models WHERE id=? AND status=1`, configured) + if model, err := a.scanAIModel(row.Scan); err == nil { + return model, nil + } + } + row := a.db.QueryRowContext(ctx, `SELECT `+aiModelColumns+` FROM ai_models WHERE status=1 ORDER BY is_default DESC,id ASC LIMIT 1`) + model, err := a.scanAIModel(row.Scan) + if errors.Is(err, sql.ErrNoRows) { + return aiModel{}, fmt.Errorf("尚未配置可用的 AI 模型") + } + return model, err +} + +func (a *App) adminAIModels(w http.ResponseWriter, r *http.Request) { + rows, err := a.db.QueryContext(r.Context(), `SELECT `+aiModelColumns+` FROM ai_models ORDER BY is_default DESC,id ASC`) + if err != nil { + fail(w, http.StatusInternalServerError, 50001, "查询模型列表失败") + return + } + defer rows.Close() + items := []aiModel{} + for rows.Next() { + model, scanErr := a.scanAIModel(rows.Scan) + if scanErr != nil && model.ID == 0 { + fail(w, http.StatusInternalServerError, 50001, "查询模型列表失败") + return + } + model.APIKey = "" + items = append(items, model) + } + if rows.Err() != nil { + fail(w, http.StatusInternalServerError, 50001, "查询模型列表失败") + return + } + protocols := make([]map[string]string, 0, len(aiProtocols)) + for _, protocol := range aiProtocols { + protocols = append(protocols, map[string]string{"value": protocol, "label": aiProtocolName(protocol)}) + } + reply(w, map[string]any{ + "items": items, + "protocols": protocols, + "enabled": a.configBool(r.Context(), "ai.enabled", false), + "secretMask": maskedSecret, + }) +} + +func (a *App) normalizeAIModelRequest(req *aiModelRequest, existing *aiModel) error { + req.Name = strings.TrimSpace(req.Name) + req.Protocol = strings.ToLower(strings.TrimSpace(req.Protocol)) + req.BaseURL = strings.TrimSpace(req.BaseURL) + req.ModelName = strings.TrimSpace(req.ModelName) + if existing != nil { + if req.Name == "" { + req.Name = existing.Name + } + if req.Protocol == "" { + req.Protocol = existing.Protocol + } + } + if req.Name == "" || len([]rune(req.Name)) > 60 { + return fmt.Errorf("模型名称须为 1–60 字") + } + if !containsString(aiProtocols, req.Protocol) { + return fmt.Errorf("不支持的模型协议") + } + if a.config.Environment == "production" && req.Protocol == "debug" { + return fmt.Errorf("生产环境禁止使用本地调试模型") + } + if len(req.BaseURL) > 300 || len([]rune(req.ModelName)) > 120 { + return fmt.Errorf("接口地址最多 300 字,模型名称最多 120 字") + } + if _, err := parseAIBaseURL(req.Protocol, req.BaseURL, a.config.Environment == "production"); err != nil { + return err + } + if req.Protocol != "debug" && req.Protocol != "webhook" && req.ModelName == "" { + return fmt.Errorf("请填写模型名称,例如 gpt-4o-mini 或 deepseek-chat") + } + if req.Temperature != nil && (*req.Temperature < 0 || *req.Temperature > 2) { + return fmt.Errorf("温度必须在 0 到 2 之间") + } + if req.MaxTokens != nil && (*req.MaxTokens < 16 || *req.MaxTokens > 4096) { + return fmt.Errorf("最大回复长度必须在 16 到 4096 之间") + } + if req.TimeoutMS != nil && (*req.TimeoutMS < 1000 || *req.TimeoutMS > 60000) { + return fmt.Errorf("超时必须在 1000 到 60000 毫秒之间") + } + if req.Status != nil && *req.Status != 0 && *req.Status != 1 { + return fmt.Errorf("状态无效") + } + return nil +} + +func (a *App) adminCreateAIModel(w http.ResponseWriter, r *http.Request) { + var req aiModelRequest + if decode(r, &req) != nil { + fail(w, http.StatusBadRequest, 20001, "模型配置格式不正确") + return + } + if err := a.normalizeAIModelRequest(&req, nil); err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + cipher, err := a.encryptSecret(strings.TrimSpace(req.APIKey)) + if err != nil { + fail(w, http.StatusInternalServerError, 50001, "加密模型密钥失败") + return + } + temperature, maxTokens, timeout, status := 0.8, 256, 15000, 1 + if req.Temperature != nil { + temperature = *req.Temperature + } + if req.MaxTokens != nil { + maxTokens = *req.MaxTokens + } + if req.TimeoutMS != nil { + timeout = *req.TimeoutMS + } + if req.Status != nil { + status = *req.Status + } + isDefault := req.IsDefault != nil && *req.IsDefault + tx, err := a.db.BeginTx(r.Context(), nil) + if err != nil { + fail(w, http.StatusInternalServerError, 50001, "保存模型失败") + return + } + defer func() { _ = tx.Rollback() }() + if isDefault { + if _, err = tx.ExecContext(r.Context(), `UPDATE ai_models SET is_default=0 WHERE is_default=1`); err != nil { + fail(w, http.StatusInternalServerError, 50001, "保存模型失败") + return + } + } + result, err := tx.ExecContext(r.Context(), `INSERT INTO ai_models + (name,protocol,base_url,model_name,api_key_cipher,temperature,max_tokens,timeout_ms,status,is_default) + VALUES(?,?,?,?,?,?,?,?,?,?)`, + req.Name, req.Protocol, req.BaseURL, req.ModelName, cipher, temperature, maxTokens, timeout, status, btoi(isDefault)) + if err != nil { + if isDuplicateAIModelName(err) { + fail(w, http.StatusBadRequest, 20001, "模型名称已存在") + return + } + fail(w, http.StatusInternalServerError, 50001, "保存模型失败") + return + } + id, _ := result.LastInsertId() + if err = tx.Commit(); err != nil { + fail(w, http.StatusInternalServerError, 50001, "保存模型失败") + return + } + a.audit(r, "create", "ai_model", id, map[string]any{"name": req.Name, "protocol": req.Protocol}) + reply(w, map[string]any{"id": id}) +} + +func (a *App) adminUpdateAIModel(w http.ResponseWriter, r *http.Request) { + id, err := pathID(r) + if err != nil { + fail(w, http.StatusBadRequest, 20001, "模型不存在") + return + } + existing, err := a.loadAIModel(r.Context(), id) + if errors.Is(err, sql.ErrNoRows) { + fail(w, http.StatusNotFound, 30001, "模型不存在") + return + } + var req aiModelRequest + if decode(r, &req) != nil { + fail(w, http.StatusBadRequest, 20001, "模型配置格式不正确") + return + } + if err = a.normalizeAIModelRequest(&req, &existing); err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + // An empty key means "keep the stored one"; clearing is explicit. + cipherClause, cipherArgs := "", []any{} + if req.ClearAPIKey { + cipherClause = ",api_key_cipher=''" + } else if strings.TrimSpace(req.APIKey) != "" { + cipher, encryptErr := a.encryptSecret(strings.TrimSpace(req.APIKey)) + if encryptErr != nil { + fail(w, http.StatusInternalServerError, 50001, "加密模型密钥失败") + return + } + cipherClause = ",api_key_cipher=?" + cipherArgs = append(cipherArgs, cipher) + } + temperature, maxTokens, timeout, status := existing.Temperature, existing.MaxTokens, existing.TimeoutMS, existing.Status + if req.Temperature != nil { + temperature = *req.Temperature + } + if req.MaxTokens != nil { + maxTokens = *req.MaxTokens + } + if req.TimeoutMS != nil { + timeout = *req.TimeoutMS + } + if req.Status != nil { + status = *req.Status + } + isDefault := existing.IsDefault + if req.IsDefault != nil { + isDefault = *req.IsDefault + } + tx, err := a.db.BeginTx(r.Context(), nil) + if err != nil { + fail(w, http.StatusInternalServerError, 50001, "保存模型失败") + return + } + defer func() { _ = tx.Rollback() }() + if isDefault { + if _, err = tx.ExecContext(r.Context(), `UPDATE ai_models SET is_default=0 WHERE is_default=1 AND id<>?`, id); err != nil { + fail(w, http.StatusInternalServerError, 50001, "保存模型失败") + return + } + } + args := append([]any{req.Name, req.Protocol, req.BaseURL, req.ModelName, temperature, maxTokens, timeout, status, btoi(isDefault)}, cipherArgs...) + args = append(args, id) + if _, err = tx.ExecContext(r.Context(), `UPDATE ai_models SET name=?,protocol=?,base_url=?,model_name=?,temperature=?,max_tokens=?,timeout_ms=?,status=?,is_default=?`+cipherClause+` WHERE id=?`, args...); err != nil { + if isDuplicateAIModelName(err) { + fail(w, http.StatusBadRequest, 20001, "模型名称已存在") + return + } + fail(w, http.StatusInternalServerError, 50001, "保存模型失败") + return + } + if err = tx.Commit(); err != nil { + fail(w, http.StatusInternalServerError, 50001, "保存模型失败") + return + } + a.audit(r, "update", "ai_model", id, map[string]any{"name": req.Name, "protocol": req.Protocol, "status": status}) + reply(w, map[string]any{"success": true}) +} + +func (a *App) adminDeleteAIModel(w http.ResponseWriter, r *http.Request) { + id, err := pathID(r) + if err != nil { + fail(w, http.StatusBadRequest, 20001, "模型不存在") + return + } + if _, err = a.db.ExecContext(r.Context(), `DELETE FROM ai_models WHERE id=?`, id); err != nil { + fail(w, http.StatusInternalServerError, 50001, "删除模型失败") + return + } + a.audit(r, "delete", "ai_model", id, nil) + reply(w, map[string]any{"success": true}) +} + +// A probe call from the console: the operator sees latency, tokens and the +// actual reply before any conversation depends on the configuration. +func (a *App) adminTestAIModel(w http.ResponseWriter, r *http.Request) { + id, err := pathID(r) + if err != nil { + fail(w, http.StatusBadRequest, 20001, "模型不存在") + return + } + model, err := a.loadAIModel(r.Context(), id) + if errors.Is(err, sql.ErrNoRows) { + fail(w, http.StatusNotFound, 30001, "模型不存在") + return + } + if err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + provider, err := newAIProvider(model, a.config.Environment == "production") + if err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + request := aiChatRequest{ + System: "你在参与一次接口连通性测试,请直接给出简短的中文回复。", + Messages: []aiChatMessage{{Role: "user", Text: aiProbePrompt}}, + } + started := time.Now() + result, err := provider.Chat(r.Context(), request) + latency := int(time.Since(started).Milliseconds()) + a.logAICall(r.Context(), model.ID, "probe", 0, 0, latency, result, err) + if err != nil { + fail(w, http.StatusBadGateway, 50002, err.Error()) + return + } + a.audit(r, "test", "ai_model", id, map[string]any{"latencyMs": latency}) + reply(w, map[string]any{ + "reply": result.Text, + "latencyMs": latency, + "promptTokens": result.PromptTokens, + "completionTokens": result.CompletionTokens, + "finishReason": result.FinishReason, + }) +} + +func (a *App) logAICall(ctx context.Context, modelID int64, scene string, agentUserID, conversationID int64, latencyMS int, result aiChatResult, callErr error) { + status, message := 1, "" + if callErr != nil { + status = 0 + message = callErr.Error() + if len([]rune(message)) > 500 { + message = string([]rune(message)[:500]) + } + } + text := result.Text + if len([]rune(text)) > 1000 { + text = string([]rune(text)[:1000]) + } + digest := "" + if scene != "" { + sum := sha256.Sum256([]byte(fmt.Sprintf("%s:%d:%d", scene, agentUserID, conversationID))) + digest = hex.EncodeToString(sum[:])[:32] + } + _, _ = a.db.ExecContext(ctx, `INSERT INTO ai_call_logs + (model_id,scene,agent_user_id,conversation_id,latency_ms,prompt_tokens,completion_tokens,status,error,prompt_digest,reply_text) + VALUES(?,?,?,?,?,?,?,?,?,?,?)`, + modelID, scene, agentUserID, conversationID, latencyMS, result.PromptTokens, result.CompletionTokens, status, message, digest, text) +} + +func isDuplicateAIModelName(err error) bool { + var mysqlErr *mysql.MySQLError + return errors.As(err, &mysqlErr) && mysqlErr.Number == 1062 && strings.Contains(mysqlErr.Message, "uk_ai_models_name") +} + +// Global switches are validated together because several of them only make +// sense as a pair, and the console tests them before the worker ever runs. +func (a *App) validateAIConfig(ctx context.Context) error { + if !a.configBool(ctx, "ai.enabled", false) { + return fmt.Errorf("AI 托管当前未启用") + } + minDelay := a.configInt(ctx, "ai.reply_delay_min_ms", 2000) + maxDelay := a.configInt(ctx, "ai.reply_delay_max_ms", 6000) + if minDelay < 0 || maxDelay < minDelay || maxDelay > 120000 { + return fmt.Errorf("回复延迟区间无效,最大延迟须不小于最小延迟且不超过 120000 毫秒") + } + if concurrency := a.configInt(ctx, "ai.max_concurrency", 4); concurrency < 1 || concurrency > 64 { + return fmt.Errorf("并发调用上限必须在 1 到 64 之间") + } + if limit := a.configInt(ctx, "ai.daily_call_limit", 0); limit < 0 { + return fmt.Errorf("每日调用上限不能为负数") + } + if limit := a.configInt(ctx, "ai.agent_daily_reply_limit", 0); limit < 0 { + return fmt.Errorf("单账号每日回复上限不能为负数") + } + if days := a.configInt(ctx, "ai.log_retention_days", 30); days < 1 || days > 365 { + return fmt.Errorf("调用日志保留天数必须在 1 到 365 之间") + } + for _, batch := range strings.Split(a.configPlain(ctx, "ai.allow_batches", ""), ",") { + if len(strings.TrimSpace(batch)) > 64 { + return fmt.Errorf("测试批次名称长度不能超过 64 位") + } + } + return nil +} + +func (a *App) configInt(ctx context.Context, key string, fallback int) int { + value, err := strconv.Atoi(strings.TrimSpace(a.configPlain(ctx, key, ""))) + if err != nil { + return fallback + } + return value +} diff --git a/im/backend/internal/app/admin_ai_agents.go b/im/backend/internal/app/admin_ai_agents.go new file mode 100644 index 0000000..5921fff --- /dev/null +++ b/im/backend/internal/app/admin_ai_agents.go @@ -0,0 +1,248 @@ +package app + +import ( + "database/sql" + "fmt" + "net/http" + "strconv" + "strings" +) + +type aiAgentView struct { + UserID int64 `json:"userId"` + PublicID string `json:"publicId"` + Nickname string `json:"nickname"` + TestBatch string `json:"testBatch"` + Enabled bool `json:"enabled"` + ModelID int64 `json:"modelId"` + ModelName string `json:"modelName"` + Persona string `json:"persona"` + DailyReplyLimit int `json:"dailyReplyLimit"` + ReplyCount int `json:"replyCount"` + LastReplyAt string `json:"lastReplyAt"` + TodayReplies int `json:"todayReplies"` +} + +// Only test accounts can be managed, so the list is the test account list with +// its binding attached rather than a separate registry. +func (a *App) adminAIAgents(w http.ResponseWriter, r *http.Request) { + page, size, offset := pagination(r) + batch := strings.TrimSpace(r.URL.Query().Get("batch")) + where, args, err := testUserFilter("test", batch) + if err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + if strings.TrimSpace(r.URL.Query().Get("managed")) == "1" { + where += " AND g.user_id IS NOT NULL AND g.enabled=1" + } + var total int + countArgs := append([]any{}, args...) + if err = a.db.QueryRowContext(r.Context(), `SELECT COUNT(*) FROM users u + LEFT JOIN ai_agents g ON g.user_id=u.id + WHERE u.status=1 AND u.deleted_at IS NULL`+where, countArgs...).Scan(&total); err != nil { + fail(w, http.StatusInternalServerError, 50001, "查询托管账号失败") + return + } + rows, err := a.db.QueryContext(r.Context(), `SELECT u.id,u.public_id,u.test_batch,p.nickname, + COALESCE(g.enabled,0),COALESCE(g.model_id,0),COALESCE(g.persona,''),COALESCE(g.daily_reply_limit,0), + COALESCE(g.reply_count,0),g.last_reply_at,COALESCE(m.name,''), + (SELECT COUNT(*) FROM ai_call_logs l WHERE l.scene='chat' AND l.agent_user_id=u.id AND l.status=1 AND l.created_at>=CURRENT_DATE()) + FROM users u JOIN user_profiles p ON p.user_id=u.id + LEFT JOIN ai_agents g ON g.user_id=u.id + LEFT JOIN ai_models m ON m.id=g.model_id + WHERE u.status=1 AND u.deleted_at IS NULL`+where+` ORDER BY u.id ASC LIMIT ? OFFSET ?`, + append(args, size, offset)...) + if err != nil { + fail(w, http.StatusInternalServerError, 50001, "查询托管账号失败") + return + } + defer rows.Close() + items := []aiAgentView{} + for rows.Next() { + var item aiAgentView + var enabled int + var lastReplyAt sql.NullTime + if rows.Scan(&item.UserID, &item.PublicID, &item.TestBatch, &item.Nickname, &enabled, &item.ModelID, + &item.Persona, &item.DailyReplyLimit, &item.ReplyCount, &lastReplyAt, &item.ModelName, &item.TodayReplies) != nil { + continue + } + item.Enabled = enabled == 1 + if lastReplyAt.Valid { + item.LastReplyAt = lastReplyAt.Time.Format("2006-01-02 15:04") + } + items = append(items, item) + } + reply(w, map[string]any{"items": items, "page": page, "size": size, "total": total, "enabled": a.configBool(r.Context(), "ai.enabled", false)}) +} + +type aiAgentRequest struct { + UserIDs []int64 `json:"userIds"` + Batch string `json:"batch"` + Enabled *bool `json:"enabled"` + ModelID *int64 `json:"modelId"` + Persona *string `json:"persona"` + DailyReplyLimit *int `json:"dailyReplyLimit"` +} + +// One endpoint for a single account and for a whole batch: operations usually +// switch a seeded batch on or off in one go. +func (a *App) adminUpdateAIAgents(w http.ResponseWriter, r *http.Request) { + var req aiAgentRequest + if decode(r, &req) != nil { + fail(w, http.StatusBadRequest, 20001, "托管配置格式不正确") + return + } + if req.Persona != nil && len([]rune(*req.Persona)) > 1000 { + fail(w, http.StatusBadRequest, 20001, "人设描述最多 1000 字") + return + } + if req.DailyReplyLimit != nil && (*req.DailyReplyLimit < 0 || *req.DailyReplyLimit > 10000) { + fail(w, http.StatusBadRequest, 20001, "单账号每日回复上限必须在 0 到 10000 之间") + return + } + if req.ModelID != nil && *req.ModelID > 0 { + if _, err := a.loadAIModel(r.Context(), *req.ModelID); err != nil { + fail(w, http.StatusBadRequest, 20001, "指定的模型不存在") + return + } + } + targets, err := a.resolveAIAgentTargets(r, req) + if err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + if len(targets) == 0 { + fail(w, http.StatusBadRequest, 20001, "没有匹配的测试账号") + return + } + enabled := true + if req.Enabled != nil { + enabled = *req.Enabled + } + modelID := int64(0) + if req.ModelID != nil { + modelID = *req.ModelID + } + persona := "" + if req.Persona != nil { + persona = strings.TrimSpace(*req.Persona) + } + limit := 0 + if req.DailyReplyLimit != nil { + limit = *req.DailyReplyLimit + } + tx, err := a.db.BeginTx(r.Context(), nil) + if err != nil { + fail(w, http.StatusInternalServerError, 50001, "保存托管配置失败") + return + } + defer func() { _ = tx.Rollback() }() + for _, userID := range targets { + if _, err = tx.ExecContext(r.Context(), `INSERT INTO ai_agents(user_id,model_id,persona,enabled,daily_reply_limit) + VALUES(?,?,?,?,?) + ON DUPLICATE KEY UPDATE model_id=VALUES(model_id),persona=VALUES(persona),enabled=VALUES(enabled),daily_reply_limit=VALUES(daily_reply_limit)`, + userID, modelID, persona, btoi(enabled), limit); err != nil { + fail(w, http.StatusInternalServerError, 50001, "保存托管配置失败") + return + } + } + if err = tx.Commit(); err != nil { + fail(w, http.StatusInternalServerError, 50001, "保存托管配置失败") + return + } + a.audit(r, "update", "ai_agent", 0, map[string]any{"count": len(targets), "enabled": enabled, "modelId": modelID, "batch": req.Batch}) + reply(w, map[string]any{"success": true, "count": len(targets)}) +} + +// Managed accounts must be test accounts; a batch is resolved server-side so a +// stale client list cannot bind a real user. +func (a *App) resolveAIAgentTargets(r *http.Request, req aiAgentRequest) ([]int64, error) { + batch := strings.TrimSpace(req.Batch) + if batch == "" && len(req.UserIDs) == 0 { + return nil, fmt.Errorf("请选择测试账号或填写测试批次") + } + if len(req.UserIDs) > 500 { + return nil, fmt.Errorf("单次最多操作 500 个账号") + } + query := `SELECT u.id FROM users u WHERE u.is_test=1 AND u.status=1 AND u.deleted_at IS NULL` + args := []any{} + if batch != "" { + if len(batch) > 64 { + return nil, fmt.Errorf("测试批次长度不能超过64位") + } + query += ` AND u.test_batch=?` + args = append(args, batch) + } + if len(req.UserIDs) > 0 { + placeholders := make([]string, 0, len(req.UserIDs)) + for _, userID := range req.UserIDs { + placeholders = append(placeholders, "?") + args = append(args, userID) + } + query += ` AND u.id IN (` + strings.Join(placeholders, ",") + `)` + } + rows, err := a.db.QueryContext(r.Context(), query, args...) + if err != nil { + return nil, fmt.Errorf("查询测试账号失败") + } + defer rows.Close() + targets := []int64{} + for rows.Next() { + var userID int64 + if rows.Scan(&userID) == nil { + targets = append(targets, userID) + } + } + return targets, rows.Err() +} + +// The queue is the first place to look when replies stop arriving. +func (a *App) adminAIReplyJobs(w http.ResponseWriter, r *http.Request) { + page, size, offset := pagination(r) + where, args := "", []any{} + if status := strings.TrimSpace(r.URL.Query().Get("status")); status != "" { + value, err := strconv.Atoi(status) + if err != nil || value < 0 || value > 3 { + fail(w, http.StatusBadRequest, 20001, "任务状态无效") + return + } + where = " WHERE j.status=?" + args = append(args, value) + } + var total int + if err := a.db.QueryRowContext(r.Context(), `SELECT COUNT(*) FROM ai_reply_jobs j`+where, args...).Scan(&total); err != nil { + fail(w, http.StatusInternalServerError, 50001, "查询回复任务失败") + return + } + rows, err := a.db.QueryContext(r.Context(), `SELECT j.id,j.conversation_id,j.agent_user_id,COALESCE(p.nickname,''),j.status,j.attempts,j.run_after,j.last_error,j.created_at + FROM ai_reply_jobs j LEFT JOIN user_profiles p ON p.user_id=j.agent_user_id`+where+` ORDER BY j.id DESC LIMIT ? OFFSET ?`, + append(args, size, offset)...) + if err != nil { + fail(w, http.StatusInternalServerError, 50001, "查询回复任务失败") + return + } + defer rows.Close() + items := []map[string]any{} + for rows.Next() { + var id, conversationID, agentUserID int64 + var nickname, lastError string + var status, attempts int + var runAfter, createdAt sql.NullTime + if rows.Scan(&id, &conversationID, &agentUserID, &nickname, &status, &attempts, &runAfter, &lastError, &createdAt) != nil { + continue + } + item := map[string]any{ + "id": id, "conversationId": conversationID, "agentUserId": agentUserID, "nickname": nickname, + "status": status, "attempts": attempts, "lastError": lastError, + } + if runAfter.Valid { + item["runAfter"] = runAfter.Time.Format("2006-01-02 15:04:05") + } + if createdAt.Valid { + item["createdAt"] = createdAt.Time.Format("2006-01-02 15:04:05") + } + items = append(items, item) + } + reply(w, map[string]any{"items": items, "page": page, "size": size, "total": total}) +} diff --git a/im/backend/internal/app/admin_ai_logs.go b/im/backend/internal/app/admin_ai_logs.go new file mode 100644 index 0000000..373bfd5 --- /dev/null +++ b/im/backend/internal/app/admin_ai_logs.go @@ -0,0 +1,304 @@ +package app + +import ( + "context" + "database/sql" + "net/http" + "strconv" + "strings" + "time" +) + +const aiUsageDays = 7 + +func (a *App) adminAILogs(w http.ResponseWriter, r *http.Request) { + page, size, offset := pagination(r) + where := []string{} + args := []any{} + if modelID := strings.TrimSpace(r.URL.Query().Get("modelId")); modelID != "" { + value, err := strconv.ParseInt(modelID, 10, 64) + if err != nil { + fail(w, http.StatusBadRequest, 20001, "模型筛选无效") + return + } + where = append(where, "l.model_id=?") + args = append(args, value) + } + if scene := strings.TrimSpace(r.URL.Query().Get("scene")); scene != "" { + if !containsString([]string{"chat", "probe"}, scene) { + fail(w, http.StatusBadRequest, 20001, "场景筛选无效") + return + } + where = append(where, "l.scene=?") + args = append(args, scene) + } + if status := strings.TrimSpace(r.URL.Query().Get("status")); status != "" { + if status != "0" && status != "1" { + fail(w, http.StatusBadRequest, 20001, "状态筛选无效") + return + } + where = append(where, "l.status=?") + args = append(args, status) + } + if days := strings.TrimSpace(r.URL.Query().Get("days")); days != "" { + value, err := strconv.Atoi(days) + if err != nil || value < 1 || value > 90 { + fail(w, http.StatusBadRequest, 20001, "时间范围必须在 1 到 90 天之间") + return + } + where = append(where, "l.created_at>=DATE_SUB(CURRENT_DATE(),INTERVAL ? DAY)") + args = append(args, value-1) + } + clause := "" + if len(where) > 0 { + clause = " WHERE " + strings.Join(where, " AND ") + } + var total int + if err := a.db.QueryRowContext(r.Context(), `SELECT COUNT(*) FROM ai_call_logs l`+clause, args...).Scan(&total); err != nil { + fail(w, http.StatusInternalServerError, 50001, "查询调用日志失败") + return + } + rows, err := a.db.QueryContext(r.Context(), `SELECT l.id,l.model_id,COALESCE(m.name,''),l.scene,l.agent_user_id,COALESCE(p.nickname,''), + l.conversation_id,l.latency_ms,l.prompt_tokens,l.completion_tokens,l.status,l.error,l.reply_text,l.created_at + FROM ai_call_logs l LEFT JOIN ai_models m ON m.id=l.model_id LEFT JOIN user_profiles p ON p.user_id=l.agent_user_id`+ + clause+` ORDER BY l.id DESC LIMIT ? OFFSET ?`, append(args, size, offset)...) + if err != nil { + fail(w, http.StatusInternalServerError, 50001, "查询调用日志失败") + return + } + defer rows.Close() + items := []map[string]any{} + for rows.Next() { + var id, modelID, agentUserID, conversationID int64 + var modelName, scene, nickname, errText, replyText string + var latency, promptTokens, completionTokens, status int + var createdAt sql.NullTime + if rows.Scan(&id, &modelID, &modelName, &scene, &agentUserID, &nickname, &conversationID, + &latency, &promptTokens, &completionTokens, &status, &errText, &replyText, &createdAt) != nil { + continue + } + item := map[string]any{ + "id": id, "modelId": modelID, "modelName": modelName, "scene": scene, + "agentUserId": agentUserID, "nickname": nickname, "conversationId": conversationID, + "latencyMs": latency, "promptTokens": promptTokens, "completionTokens": completionTokens, + "status": status, "error": errText, "reply": replyText, + } + if createdAt.Valid { + item["createdAt"] = createdAt.Time.Format("2006-01-02 15:04:05") + } + items = append(items, item) + } + reply(w, map[string]any{"items": items, "page": page, "size": size, "total": total}) +} + +// Usage and health in one payload: the console needs both to answer "is it +// running and what is it costing". +func (a *App) adminAIUsage(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + today := map[string]any{"calls": 0, "success": 0, "failed": 0, "promptTokens": 0, "completionTokens": 0, "avgLatencyMs": 0} + var calls, success, promptTokens, completionTokens sql.NullInt64 + // AVG() comes back as DECIMAL, which will not scan into an integer. + var avgLatency sql.NullFloat64 + _ = a.db.QueryRowContext(ctx, `SELECT COUNT(*),SUM(status=1),SUM(prompt_tokens),SUM(completion_tokens),AVG(latency_ms) + FROM ai_call_logs WHERE created_at>=CURRENT_DATE()`).Scan(&calls, &success, &promptTokens, &completionTokens, &avgLatency) + today["calls"] = calls.Int64 + today["success"] = success.Int64 + today["failed"] = calls.Int64 - success.Int64 + today["promptTokens"] = promptTokens.Int64 + today["completionTokens"] = completionTokens.Int64 + today["avgLatencyMs"] = int64(avgLatency.Float64 + 0.5) + + daily := []map[string]any{} + dailyRows, err := a.db.QueryContext(ctx, `SELECT DATE(created_at),COUNT(*),SUM(status=1),SUM(prompt_tokens+completion_tokens) + FROM ai_call_logs WHERE created_at>=DATE_SUB(CURRENT_DATE(),INTERVAL ? DAY) + GROUP BY DATE(created_at) ORDER BY DATE(created_at)`, aiUsageDays-1) + if err == nil { + defer dailyRows.Close() + for dailyRows.Next() { + var date sql.NullTime + var dayCalls, daySuccess, dayTokens sql.NullInt64 + if dailyRows.Scan(&date, &dayCalls, &daySuccess, &dayTokens) != nil { + continue + } + entry := map[string]any{"calls": dayCalls.Int64, "success": daySuccess.Int64, "tokens": dayTokens.Int64} + if date.Valid { + entry["date"] = date.Time.Format("01-02") + } + daily = append(daily, entry) + } + } + + models := []map[string]any{} + modelRows, queryErr := a.db.QueryContext(ctx, `SELECT l.model_id,COALESCE(m.name,'已删除模型'),COUNT(*),SUM(l.status=1), + AVG(l.latency_ms),SUM(l.prompt_tokens+l.completion_tokens) + FROM ai_call_logs l LEFT JOIN ai_models m ON m.id=l.model_id + WHERE l.created_at>=DATE_SUB(CURRENT_DATE(),INTERVAL ? DAY) + GROUP BY l.model_id,m.name ORDER BY COUNT(*) DESC`, aiUsageDays-1) + if queryErr == nil { + defer modelRows.Close() + for modelRows.Next() { + var modelID sql.NullInt64 + var name string + var modelCalls, modelSuccess, modelTokens sql.NullInt64 + var modelLatency sql.NullFloat64 + if modelRows.Scan(&modelID, &name, &modelCalls, &modelSuccess, &modelLatency, &modelTokens) != nil { + continue + } + models = append(models, map[string]any{ + "modelId": modelID.Int64, "name": name, "calls": modelCalls.Int64, + "success": modelSuccess.Int64, "avgLatencyMs": int64(modelLatency.Float64 + 0.5), "tokens": modelTokens.Int64, + }) + } + } + + jobs := map[string]any{"pending": 0, "running": 0, "failed24h": 0, "overdue": 0} + var pending, running, failed, overdue int + _ = a.db.QueryRowContext(ctx, `SELECT SUM(status=0),SUM(status=1),SUM(status=3 AND updated_at>=DATE_SUB(NOW(3),INTERVAL 1 DAY)), + SUM(status=0 AND run_after 0 { + if calls >= dailyLimit { + alerts = append(alerts, map[string]string{"level": "error", "message": "今日调用已达上限,后续消息不会再回复"}) + } else if calls*10 >= dailyLimit*8 { + alerts = append(alerts, map[string]string{"level": "warning", "message": "今日调用已超过配额的 80%"}) + } + } + // A handful of failures is noise; a persistent failure rate is not. + if calls >= 5 && failed*10 >= calls*3 { + message := "今日调用失败率超过 30%" + if lastError != "" { + message += ",最近错误:" + lastError + } + alerts = append(alerts, map[string]string{"level": "error", "message": message}) + } + if overdue > 0 { + alerts = append(alerts, map[string]string{"level": "warning", "message": "有回复任务超过 60 秒未处理,请检查后端 worker 是否在运行"}) + } + return alerts +} + +// Call logs hold conversation replies, so they are trimmed on a schedule rather +// than kept forever. Deleting in batches keeps the lock short. +func (a *App) purgeAICallLogs(ctx context.Context) { + days := a.configInt(ctx, "ai.log_retention_days", 30) + if days < 1 || days > 365 { + days = 30 + } + for attempt := 0; attempt < 20; attempt++ { + result, err := a.db.ExecContext(ctx, `DELETE FROM ai_call_logs WHERE created_at 0 && len(haystack) >= len(needle) && indexOf(haystack, needle) >= 0 +} + +func indexOf(haystack, needle string) int { + for index := 0; index+len(needle) <= len(haystack); index++ { + if haystack[index:index+len(needle)] == needle { + return index + } + } + return -1 +} diff --git a/im/backend/internal/app/ai_agent.go b/im/backend/internal/app/ai_agent.go new file mode 100644 index 0000000..476b24c --- /dev/null +++ b/im/backend/internal/app/ai_agent.go @@ -0,0 +1,439 @@ +package app + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "log" + "math/rand" + "strings" + "sync" + "time" +) + +const ( + aiHistoryMessages = 20 + aiReplyMaxRunes = 120 + aiJobMaxAttempts = 2 + // A burst of incoming messages collapses into one reply, but the reader + // should not wait longer than this for it. + aiReplyMaxDelaySeconds = 15 +) + +type aiAgentBinding struct { + UserID int64 + ModelID int64 + Persona string + DailyReplyLimit int + Nickname string + Age int + City string + Bio string +} + +type aiReplyJob struct { + ID int64 + ConversationID int64 + AgentUserID int64 + Attempts int +} + +// Called right after a message is committed, from every send path. It must stay +// cheap and must never fail the send: a missing reply is better than a failed +// message. +func (a *App) enqueueAIReply(ctx context.Context, conversationID, senderID int64, members []int64, messageID int64) { + if !a.configBool(ctx, "ai.enabled", false) { + return + } + peerID := int64(0) + for _, member := range members { + if member != senderID { + if peerID != 0 { + return // Group conversations have no single managed recipient. + } + peerID = member + } + } + if peerID == 0 { + return + } + var agentUserID int64 + var testBatch string + // The sender must not itself be managed, otherwise two agents would keep + // each other talking forever. + err := a.db.QueryRowContext(ctx, `SELECT g.user_id,u.test_batch + FROM ai_agents g JOIN users u ON u.id=g.user_id + WHERE g.user_id=? AND g.enabled=1 AND u.is_test=1 AND u.status=1 AND u.deleted_at IS NULL + AND NOT EXISTS(SELECT 1 FROM ai_agents s WHERE s.user_id=? AND s.enabled=1)`, + peerID, senderID).Scan(&agentUserID, &testBatch) + if err != nil || agentUserID == 0 { + return + } + if !aiBatchAllowed(a.configPlain(ctx, "ai.allow_batches", ""), testBatch) { + return + } + delay := a.aiReplyDelay(ctx) + if _, err = a.db.ExecContext(ctx, `INSERT INTO ai_reply_jobs(conversation_id,agent_user_id,trigger_message_id,run_after) + VALUES(?,?,?,DATE_ADD(NOW(3),INTERVAL ? MICROSECOND)) + ON DUPLICATE KEY UPDATE trigger_message_id=VALUES(trigger_message_id), + run_after=LEAST(DATE_ADD(created_at,INTERVAL ? SECOND),VALUES(run_after))`, + conversationID, agentUserID, messageID, delay.Microseconds(), aiReplyMaxDelaySeconds); err != nil { + log.Printf("ai: 入队失败 conversation=%d: %v", conversationID, err) + } +} + +// An empty allow list means every test account; otherwise the account's batch +// has to be named explicitly. +func aiBatchAllowed(allowed, batch string) bool { + allowed = strings.TrimSpace(allowed) + if allowed == "" { + return true + } + for _, item := range strings.Split(allowed, ",") { + if strings.TrimSpace(item) == strings.TrimSpace(batch) && batch != "" { + return true + } + } + return false +} + +func (a *App) aiReplyDelay(ctx context.Context) time.Duration { + minDelay := a.configInt(ctx, "ai.reply_delay_min_ms", 2000) + maxDelay := a.configInt(ctx, "ai.reply_delay_max_ms", 6000) + if minDelay < 0 { + minDelay = 0 + } + if maxDelay <= minDelay { + maxDelay = minDelay + } + spread := maxDelay - minDelay + if spread > 0 { + minDelay += rand.Intn(spread) + } + return time.Duration(minDelay) * time.Millisecond +} + +// The worker follows the same ticker shape as processDueAccountClosures. +func (a *App) runAIReplyWorker(ctx context.Context) { + ticker := time.NewTicker(time.Second) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + a.processDueAIReplies(ctx) + } + } +} + +func (a *App) processDueAIReplies(ctx context.Context) { + if !a.configBool(ctx, "ai.enabled", false) { + return + } + concurrency := a.configInt(ctx, "ai.max_concurrency", 4) + if concurrency < 1 || concurrency > 64 { + concurrency = 4 + } + jobs, err := a.claimAIReplyJobs(ctx, concurrency) + if err != nil || len(jobs) == 0 { + return + } + var wait sync.WaitGroup + for _, job := range jobs { + wait.Add(1) + go func(job aiReplyJob) { + defer wait.Done() + if replyErr := a.processAIReplyJob(ctx, job); replyErr != nil { + a.failAIReplyJob(ctx, job, replyErr) + return + } + _, _ = a.db.ExecContext(ctx, `UPDATE ai_reply_jobs SET status=2,last_error='' WHERE id=?`, job.ID) + }(job) + } + wait.Wait() +} + +// Claiming stamps a token with a single UPDATE and reads the rows back: each +// row can only leave status 0 once, so two instances never take the same job. +// SELECT ... FOR UPDATE SKIP LOCKED would be the obvious alternative, but it +// does not parse on MySQL 5.7, which some deployments still run. +func (a *App) claimAIReplyJobs(ctx context.Context, limit int) ([]aiReplyJob, error) { + token := randomToken()[:32] + result, err := a.db.ExecContext(ctx, `UPDATE ai_reply_jobs SET status=1,attempts=attempts+1,claim_token=? + WHERE status=0 AND run_after<=NOW(3) ORDER BY run_after LIMIT ?`, token, limit) + if err != nil { + return nil, err + } + if affected, _ := result.RowsAffected(); affected == 0 { + return nil, nil + } + rows, err := a.db.QueryContext(ctx, `SELECT id,conversation_id,agent_user_id,attempts FROM ai_reply_jobs WHERE claim_token=? AND status=1`, token) + if err != nil { + return nil, err + } + defer rows.Close() + jobs := []aiReplyJob{} + for rows.Next() { + var job aiReplyJob + if rows.Scan(&job.ID, &job.ConversationID, &job.AgentUserID, &job.Attempts) == nil { + jobs = append(jobs, job) + } + } + return jobs, rows.Err() +} + +func (a *App) failAIReplyJob(ctx context.Context, job aiReplyJob, cause error) { + message := cause.Error() + if len([]rune(message)) > 500 { + message = string([]rune(message)[:500]) + } + if job.Attempts >= aiJobMaxAttempts { + _, _ = a.db.ExecContext(ctx, `UPDATE ai_reply_jobs SET status=3,last_error=? WHERE id=?`, message, job.ID) + return + } + // Returning the job to the queue can collide with a newer pending job for + // the same conversation; that newer job supersedes this one. + if _, err := a.db.ExecContext(ctx, `UPDATE ai_reply_jobs SET status=0,claim_token='',last_error=?,run_after=DATE_ADD(NOW(3),INTERVAL 5 SECOND) WHERE id=?`, message, job.ID); err != nil { + _, _ = a.db.ExecContext(ctx, `UPDATE ai_reply_jobs SET status=3,last_error=? WHERE id=?`, message, job.ID) + } +} + +func (a *App) processAIReplyJob(ctx context.Context, job aiReplyJob) error { + binding, err := a.loadAIAgent(ctx, job.AgentUserID) + if err != nil { + return err + } + if err = a.checkAIQuota(ctx, binding); err != nil { + return err + } + history, lastSenderID, err := a.aiConversationHistory(ctx, job.ConversationID, job.AgentUserID) + if err != nil { + return err + } + if len(history) == 0 || lastSenderID == job.AgentUserID { + // Someone already answered, or there is nothing to answer. + return nil + } + model, err := a.aiModelForAgent(ctx, binding) + if err != nil { + return err + } + provider, err := newAIProvider(model, a.config.Environment == "production") + if err != nil { + return err + } + first, err := a.aiFirstReply(ctx, job.ConversationID, job.AgentUserID) + if err != nil { + return err + } + request := aiChatRequest{ + System: aiSystemPrompt(binding, a.configBool(ctx, "ai.disclose_in_chat", true)), + Messages: history, + } + callCtx, cancel := context.WithTimeout(ctx, time.Duration(model.TimeoutMS+2000)*time.Millisecond) + defer cancel() + started := time.Now() + result, callErr := provider.Chat(callCtx, request) + latency := int(time.Since(started).Milliseconds()) + a.logAICall(ctx, model.ID, "chat", job.AgentUserID, job.ConversationID, latency, result, callErr) + if callErr != nil { + if fallback := strings.TrimSpace(a.configPlain(ctx, "ai.fallback_text", "")); fallback != "" { + _ = a.sendAIReply(ctx, job, fallback) + } + return callErr + } + text := sanitizeAIReply(result.Text) + if text == "" { + return fmt.Errorf("模型返回内容为空") + } + if first && a.configBool(ctx, "ai.disclose_in_chat", true) { + text = "[AI 助手] " + text + } + if err = a.sendAIReply(ctx, job, text); err != nil { + return err + } + _, _ = a.db.ExecContext(ctx, `UPDATE ai_agents SET reply_count=reply_count+1,last_reply_at=NOW(3) WHERE user_id=?`, job.AgentUserID) + return nil +} + +func (a *App) sendAIReply(ctx context.Context, job aiReplyJob, text string) error { + item, members, err := a.persistMessageContext(ctx, job.ConversationID, job.AgentUserID, + "ai-"+randomToken()[:20], 1, map[string]any{"text": text}) + if err != nil { + return err + } + // The reply reaches clients through the same push the sender's own messages + // use, so unread counts and conversation ordering need no special casing. + a.hub.broadcast(members, map[string]any{"command": "MESSAGE_PUSH", "data": item}) + return nil +} + +func (a *App) loadAIAgent(ctx context.Context, userID int64) (aiAgentBinding, error) { + binding := aiAgentBinding{UserID: userID} + var birthday sql.NullTime + err := a.db.QueryRowContext(ctx, `SELECT g.model_id,g.persona,g.daily_reply_limit,p.nickname,p.birthday,p.city_name,p.bio + FROM ai_agents g JOIN user_profiles p ON p.user_id=g.user_id JOIN users u ON u.id=g.user_id + WHERE g.user_id=? AND g.enabled=1 AND u.is_test=1 AND u.status=1 AND u.deleted_at IS NULL`, userID). + Scan(&binding.ModelID, &binding.Persona, &binding.DailyReplyLimit, &binding.Nickname, &birthday, &binding.City, &binding.Bio) + if errors.Is(err, sql.ErrNoRows) { + return binding, fmt.Errorf("账号未开启 AI 托管") + } + if birthday.Valid { + binding.Age = age(birthday.Time) + } + return binding, err +} + +func (a *App) aiModelForAgent(ctx context.Context, binding aiAgentBinding) (aiModel, error) { + if binding.ModelID > 0 { + model, err := a.loadAIModel(ctx, binding.ModelID) + if err == nil && model.Status == 1 { + return model, nil + } + } + return a.defaultAIModel(ctx) +} + +// Quotas are counted from the call log so a restart cannot reset them. +func (a *App) checkAIQuota(ctx context.Context, binding aiAgentBinding) error { + if limit := a.configInt(ctx, "ai.daily_call_limit", 0); limit > 0 { + var used int + _ = a.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM ai_call_logs WHERE scene='chat' AND created_at>=CURRENT_DATE()`).Scan(&used) + if used >= limit { + return fmt.Errorf("已达到今日全局调用上限 %d", limit) + } + } + limit := binding.DailyReplyLimit + if limit <= 0 { + limit = a.configInt(ctx, "ai.agent_daily_reply_limit", 0) + } + if limit > 0 { + var used int + _ = a.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM ai_call_logs WHERE scene='chat' AND agent_user_id=? AND status=1 AND created_at>=CURRENT_DATE()`, binding.UserID).Scan(&used) + if used >= limit { + return fmt.Errorf("账号已达到今日回复上限 %d", limit) + } + } + return nil +} + +func (a *App) aiFirstReply(ctx context.Context, conversationID, agentUserID int64) (bool, error) { + var replied int + err := a.db.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM im_messages WHERE conversation_id=? AND sender_id=?)`, conversationID, agentUserID).Scan(&replied) + return replied == 0, err +} + +// Only the tail of the conversation is sent: enough for continuity, bounded in +// cost. Media turns become placeholders because the model gets text only. +func (a *App) aiConversationHistory(ctx context.Context, conversationID, agentUserID int64) ([]aiChatMessage, int64, error) { + rows, err := a.db.QueryContext(ctx, `SELECT sender_id,message_type,body,recalled_at,admin_removed_at + FROM im_messages WHERE conversation_id=? ORDER BY seq DESC LIMIT ?`, conversationID, aiHistoryMessages) + if err != nil { + return nil, 0, err + } + defer rows.Close() + ordered := []aiChatMessage{} + lastSenderID := int64(0) + for rows.Next() { + var senderID int64 + var messageType int + var body []byte + var recalledAt, removedAt sql.NullTime + if rows.Scan(&senderID, &messageType, &body, &recalledAt, &removedAt) != nil { + continue + } + if lastSenderID == 0 { + lastSenderID = senderID + } + if recalledAt.Valid || removedAt.Valid { + continue + } + text := aiMessageText(messageType, body) + if text == "" { + continue + } + role := "user" + if senderID == agentUserID { + role = "assistant" + } + ordered = append([]aiChatMessage{{Role: role, Text: text}}, ordered...) + } + return ordered, lastSenderID, rows.Err() +} + +func aiMessageText(messageType int, body []byte) string { + switch messageType { + case 2: + return "[图片]" + case 3: + return "[语音]" + } + var content struct { + Text string `json:"text"` + } + if json.Unmarshal(body, &content) != nil { + return "" + } + return strings.TrimSpace(content.Text) +} + +func aiSystemPrompt(binding aiAgentBinding, disclose bool) string { + persona := strings.TrimSpace(binding.Persona) + if persona == "" { + // Falling back to the profile keeps the replies consistent with what the + // other person can see on the page. + details := []string{} + if binding.Nickname != "" { + details = append(details, "昵称"+binding.Nickname) + } + if binding.Age > 0 { + details = append(details, fmt.Sprintf("%d 岁", binding.Age)) + } + if binding.City != "" { + details = append(details, "在"+binding.City) + } + if bio := strings.TrimSpace(binding.Bio); bio != "" { + details = append(details, "个人简介:"+bio) + } + persona = "你是一位社交软件用户," + strings.Join(details, ",") + "。" + } + rules := []string{ + "用中文口语回复,像真人聊天一样简短自然,不超过 40 个字。", + "一次只说一件事,不要分点、不要 Markdown、不要堆叠表情。", + "不索取或提供手机号、微信等联系方式,不谈金钱、转账、投资。", + "不主动约线下见面,不输出链接。", + "不知道的事就说不知道,不要编造经历细节。", + } + if disclose { + rules = append(rules, "如果对方询问你是不是机器人或 AI,必须如实说明你是 AI 助手。") + } + return persona + "\n\n对话要求:\n- " + strings.Join(rules, "\n- ") +} + +// Models drift into assistant formatting; the chat bubble wants one plain line. +func sanitizeAIReply(text string) string { + text = strings.TrimSpace(text) + if text == "" { + return "" + } + if index := strings.Index(text, "```"); index >= 0 { + text = text[:index] + } + replacer := strings.NewReplacer("**", "", "__", "", "##", "", "> ", "", "\r", "\n") + text = replacer.Replace(text) + lines := []string{} + for _, line := range strings.Split(text, "\n") { + line = strings.TrimSpace(strings.TrimLeft(strings.TrimSpace(line), "-*•0123456789. ")) + if line != "" { + lines = append(lines, line) + } + } + text = strings.TrimSpace(strings.Join(lines, " ")) + if runes := []rune(text); len(runes) > aiReplyMaxRunes { + text = strings.TrimSpace(string(runes[:aiReplyMaxRunes])) + } + return text +} diff --git a/im/backend/internal/app/ai_agent_test.go b/im/backend/internal/app/ai_agent_test.go new file mode 100644 index 0000000..627b975 --- /dev/null +++ b/im/backend/internal/app/ai_agent_test.go @@ -0,0 +1,87 @@ +package app + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestSanitizeAIReplyFlattensAssistantFormatting(t *testing.T) { + tests := []struct { + name, input, want string + }{ + {name: "trim", input: " 在的 ", want: "在的"}, + {name: "markdown", input: "**在的**,__刚看到__", want: "在的,刚看到"}, + {name: "bullets", input: "- 第一点\n- 第二点", want: "第一点 第二点"}, + {name: "numbered", input: "1. 先这样\n2. 再那样", want: "先这样 再那样"}, + {name: "quote", input: "> 引用\n正文", want: "引用 正文"}, + {name: "code fence", input: "好的\n```go\nfmt.Println()\n```", want: "好的"}, + {name: "blank", input: "\n\n \n", want: ""}, + } + for _, test := range tests { + if got := sanitizeAIReply(test.input); got != test.want { + t.Errorf("sanitizeAIReply(%s) = %q, want %q", test.name, got, test.want) + } + } + long := strings.Repeat("字", aiReplyMaxRunes+40) + if got := []rune(sanitizeAIReply(long)); len(got) != aiReplyMaxRunes { + t.Errorf("a long reply must be cut to %d runes, got %d", aiReplyMaxRunes, len(got)) + } +} + +func TestAISystemPromptFallsBackToProfile(t *testing.T) { + binding := aiAgentBinding{Nickname: "季若宁", Age: 27, City: "青岛", Bio: "喜欢电影与绘画"} + prompt := aiSystemPrompt(binding, true) + for _, fragment := range []string{"季若宁", "27 岁", "青岛", "喜欢电影与绘画"} { + if !strings.Contains(prompt, fragment) { + t.Errorf("profile-derived persona must mention %q: %s", fragment, prompt) + } + } + // Disclosure is a monitored requirement, so it has to be switchable and + // present by default. + if !strings.Contains(prompt, "AI 助手") { + t.Error("disclosure must add the honesty rule") + } + if strings.Contains(aiSystemPrompt(binding, false), "如果对方询问") { + t.Error("disclosure off must drop the rule") + } + custom := aiSystemPrompt(aiAgentBinding{Persona: "你是一个健身教练", Nickname: "忽略我"}, false) + if !strings.Contains(custom, "健身教练") || strings.Contains(custom, "忽略我") { + t.Errorf("an explicit persona must replace the profile text: %s", custom) + } + if !strings.Contains(custom, "不索取或提供手机号") { + t.Error("safety rules must survive a custom persona") + } +} + +func TestAIMessageTextDegradesMedia(t *testing.T) { + text, _ := json.Marshal(map[string]any{"text": " 你好 "}) + if got := aiMessageText(1, text); got != "你好" { + t.Errorf("text message = %q", got) + } + if got := aiMessageText(2, []byte(`{"url":"https://example.com/a.jpg"}`)); got != "[图片]" { + t.Errorf("image placeholder = %q", got) + } + if got := aiMessageText(3, []byte(`{"url":"https://example.com/a.m4a","duration":3}`)); got != "[语音]" { + t.Errorf("audio placeholder = %q", got) + } + if got := aiMessageText(1, []byte(`not json`)); got != "" { + t.Errorf("a broken body must be skipped, got %q", got) + } +} + +func TestAIBatchAllowList(t *testing.T) { + if !aiBatchAllowed("", "any-batch") || !aiBatchAllowed("", "") { + t.Error("an empty allow list must accept every test account") + } + if !aiBatchAllowed("batch-a, batch-b", "batch-b") { + t.Error("a listed batch must be accepted") + } + if aiBatchAllowed("batch-a", "batch-c") { + t.Error("an unlisted batch must be refused") + } + // An account with no batch must not slip through a non-empty allow list. + if aiBatchAllowed("batch-a", "") { + t.Error("an empty batch must not match a configured allow list") + } +} diff --git a/im/backend/internal/app/ai_providers.go b/im/backend/internal/app/ai_providers.go new file mode 100644 index 0000000..df05b1a --- /dev/null +++ b/im/backend/internal/app/ai_providers.go @@ -0,0 +1,488 @@ +package app + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" +) + +const ( + aiResponseLimit = 512 << 10 + aiProbePrompt = "用一句不超过 20 个字的中文打个招呼。" +) + +// Vendors are many but wire protocols are few, so adapters are written per +// protocol. Anything speaking the OpenAI chat format needs no new code, only a +// new row in ai_models. +var aiProtocols = []string{"openai", "anthropic", "gemini", "webhook", "debug"} + +type aiChatMessage struct { + Role string // user 或 assistant + Text string +} + +type aiChatRequest struct { + System string + Messages []aiChatMessage + MaxTokens int + Temperature float64 +} + +type aiChatResult struct { + Text string + PromptTokens int + CompletionTokens int + FinishReason string + RequestID string +} + +type aiProvider interface { + Chat(ctx context.Context, req aiChatRequest) (aiChatResult, error) +} + +type aiModel struct { + ID int64 `json:"id"` + Name string `json:"name"` + Protocol string `json:"protocol"` + BaseURL string `json:"baseUrl"` + ModelName string `json:"modelName"` + Temperature float64 `json:"temperature"` + MaxTokens int `json:"maxTokens"` + TimeoutMS int `json:"timeoutMs"` + Status int `json:"status"` + IsDefault bool `json:"isDefault"` + APIKey string `json:"-"` + HasAPIKey bool `json:"hasApiKey"` + CreatedAt string `json:"createdAt"` + UpdatedAt string `json:"updatedAt"` +} + +func aiProtocolName(protocol string) string { + switch protocol { + case "openai": + return "OpenAI 兼容" + case "anthropic": + return "Anthropic Messages" + case "gemini": + return "Google Gemini" + case "webhook": + return "自定义 Webhook" + case "debug": + return "本地调试" + } + return protocol +} + +// The endpoint is operator-supplied, so it is validated the same way the SMS +// endpoints are: no credentials, no query string, no fragment. +func parseAIBaseURL(protocol, raw string, production bool) (*url.URL, error) { + trimmed := strings.TrimSpace(raw) + if protocol == "debug" { + return nil, nil + } + if trimmed == "" { + return nil, fmt.Errorf("%s 需要填写接口地址", aiProtocolName(protocol)) + } + parsed, err := url.Parse(trimmed) + if err != nil || parsed.Host == "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { + return nil, fmt.Errorf("接口地址必须是无账号、查询参数和片段的有效地址") + } + if parsed.Scheme != "https" && !(parsed.Scheme == "http" && !production) { + return nil, fmt.Errorf("接口地址必须使用 HTTPS") + } + return parsed, nil +} + +func aiEndpoint(base *url.URL, suffix string) string { + trimmed := strings.TrimRight(base.String(), "/") + if suffix == "" || strings.HasSuffix(trimmed, suffix) { + return trimmed + } + return trimmed + suffix +} + +func newAIProvider(model aiModel, production bool) (aiProvider, error) { + if !containsString(aiProtocols, model.Protocol) { + return nil, fmt.Errorf("不支持的模型协议 %q", model.Protocol) + } + if model.Protocol == "debug" { + if production { + return nil, fmt.Errorf("生产环境禁止使用本地调试模型") + } + return debugAIProvider{}, nil + } + base, err := parseAIBaseURL(model.Protocol, model.BaseURL, production) + if err != nil { + return nil, err + } + if strings.TrimSpace(model.APIKey) == "" && model.Protocol != "webhook" { + return nil, fmt.Errorf("%s 需要填写密钥", aiProtocolName(model.Protocol)) + } + if strings.TrimSpace(model.ModelName) == "" && model.Protocol != "webhook" { + return nil, fmt.Errorf("请填写模型名称") + } + timeout := time.Duration(model.TimeoutMS) * time.Millisecond + if timeout < time.Second || timeout > time.Minute { + timeout = 15 * time.Second + } + client := &http.Client{Timeout: timeout} + switch model.Protocol { + case "openai": + return openAIProvider{model: model, base: base, client: client}, nil + case "anthropic": + return anthropicProvider{model: model, base: base, client: client}, nil + case "gemini": + return geminiProvider{model: model, base: base, client: client}, nil + } + return webhookAIProvider{model: model, base: base, client: client}, nil +} + +func aiRequestTemperature(req aiChatRequest, model aiModel) float64 { + if req.Temperature > 0 { + return req.Temperature + } + return model.Temperature +} + +func aiRequestMaxTokens(req aiChatRequest, model aiModel) int { + if req.MaxTokens > 0 { + return req.MaxTokens + } + if model.MaxTokens > 0 { + return model.MaxTokens + } + return 256 +} + +func aiPostJSON(ctx context.Context, client *http.Client, endpoint string, headers map[string]string, payload any) ([]byte, error) { + body, err := json.Marshal(payload) + if err != nil { + return nil, err + } + request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) + if err != nil { + return nil, err + } + request.Header.Set("Content-Type", "application/json") + request.Header.Set("Accept", "application/json") + for key, value := range headers { + request.Header.Set(key, value) + } + response, err := client.Do(request) + if err != nil { + return nil, fmt.Errorf("模型服务请求失败:%w", err) + } + defer response.Body.Close() + raw, err := io.ReadAll(io.LimitReader(response.Body, aiResponseLimit)) + if err != nil { + return nil, fmt.Errorf("读取模型响应失败:%w", err) + } + if response.StatusCode < 200 || response.StatusCode >= 300 { + return nil, fmt.Errorf("模型服务返回 HTTP %d:%s", response.StatusCode, aiErrorSnippet(raw)) + } + return raw, nil +} + +// Provider error bodies are quoted back to the operator, so they are trimmed to +// something a config page can display. +func aiErrorSnippet(raw []byte) string { + text := strings.TrimSpace(string(raw)) + text = strings.Join(strings.Fields(text), " ") + if len([]rune(text)) > 200 { + return string([]rune(text)[:200]) + "…" + } + if text == "" { + return "无响应内容" + } + return text +} + +type openAIProvider struct { + model aiModel + base *url.URL + client *http.Client +} + +func (p openAIProvider) Chat(ctx context.Context, req aiChatRequest) (aiChatResult, error) { + messages := make([]map[string]string, 0, len(req.Messages)+1) + if strings.TrimSpace(req.System) != "" { + messages = append(messages, map[string]string{"role": "system", "content": req.System}) + } + for _, message := range req.Messages { + messages = append(messages, map[string]string{"role": message.Role, "content": message.Text}) + } + raw, err := aiPostJSON(ctx, p.client, aiEndpoint(p.base, "/chat/completions"), + map[string]string{"Authorization": "Bearer " + p.model.APIKey}, + map[string]any{ + "model": p.model.ModelName, + "messages": messages, + "max_tokens": aiRequestMaxTokens(req, p.model), + "temperature": aiRequestTemperature(req, p.model), + "stream": false, + }) + if err != nil { + return aiChatResult{}, err + } + var payload struct { + ID string `json:"id"` + Choices []struct { + Message struct { + Content string `json:"content"` + } `json:"message"` + FinishReason string `json:"finish_reason"` + } `json:"choices"` + Usage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + } `json:"usage"` + Error struct { + Message string `json:"message"` + } `json:"error"` + } + if err = json.Unmarshal(raw, &payload); err != nil { + return aiChatResult{}, fmt.Errorf("模型响应不是有效 JSON:%s", aiErrorSnippet(raw)) + } + if payload.Error.Message != "" { + return aiChatResult{}, fmt.Errorf("模型服务返回错误:%s", payload.Error.Message) + } + if len(payload.Choices) == 0 { + return aiChatResult{}, fmt.Errorf("模型没有返回任何内容") + } + return aiChatResult{ + Text: strings.TrimSpace(payload.Choices[0].Message.Content), + PromptTokens: payload.Usage.PromptTokens, + CompletionTokens: payload.Usage.CompletionTokens, + FinishReason: payload.Choices[0].FinishReason, + RequestID: payload.ID, + }, nil +} + +type anthropicProvider struct { + model aiModel + base *url.URL + client *http.Client +} + +func (p anthropicProvider) Chat(ctx context.Context, req aiChatRequest) (aiChatResult, error) { + messages := make([]map[string]any, 0, len(req.Messages)) + for _, message := range req.Messages { + messages = append(messages, map[string]any{ + "role": message.Role, + "content": []map[string]string{{"type": "text", "text": message.Text}}, + }) + } + payloadBody := map[string]any{ + "model": p.model.ModelName, + "messages": messages, + "max_tokens": aiRequestMaxTokens(req, p.model), + "temperature": aiRequestTemperature(req, p.model), + } + // Anthropic carries the system prompt outside the message list. + if strings.TrimSpace(req.System) != "" { + payloadBody["system"] = req.System + } + raw, err := aiPostJSON(ctx, p.client, aiEndpoint(p.base, "/v1/messages"), + map[string]string{"x-api-key": p.model.APIKey, "anthropic-version": "2023-06-01"}, payloadBody) + if err != nil { + return aiChatResult{}, err + } + var payload struct { + ID string `json:"id"` + Content []struct { + Type string `json:"type"` + Text string `json:"text"` + } `json:"content"` + StopReason string `json:"stop_reason"` + Usage struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + } `json:"usage"` + Error struct { + Message string `json:"message"` + } `json:"error"` + } + if err = json.Unmarshal(raw, &payload); err != nil { + return aiChatResult{}, fmt.Errorf("模型响应不是有效 JSON:%s", aiErrorSnippet(raw)) + } + if payload.Error.Message != "" { + return aiChatResult{}, fmt.Errorf("模型服务返回错误:%s", payload.Error.Message) + } + text := "" + for _, block := range payload.Content { + if block.Type == "text" { + text += block.Text + } + } + if strings.TrimSpace(text) == "" { + return aiChatResult{}, fmt.Errorf("模型没有返回任何内容") + } + return aiChatResult{ + Text: strings.TrimSpace(text), + PromptTokens: payload.Usage.InputTokens, + CompletionTokens: payload.Usage.OutputTokens, + FinishReason: payload.StopReason, + RequestID: payload.ID, + }, nil +} + +type geminiProvider struct { + model aiModel + base *url.URL + client *http.Client +} + +func (p geminiProvider) Chat(ctx context.Context, req aiChatRequest) (aiChatResult, error) { + contents := make([]map[string]any, 0, len(req.Messages)) + for _, message := range req.Messages { + role := "user" + if message.Role == "assistant" { + role = "model" + } + contents = append(contents, map[string]any{ + "role": role, + "parts": []map[string]string{{"text": message.Text}}, + }) + } + payloadBody := map[string]any{ + "contents": contents, + "generationConfig": map[string]any{ + "temperature": aiRequestTemperature(req, p.model), + "maxOutputTokens": aiRequestMaxTokens(req, p.model), + }, + } + if strings.TrimSpace(req.System) != "" { + payloadBody["systemInstruction"] = map[string]any{"parts": []map[string]string{{"text": req.System}}} + } + endpoint := fmt.Sprintf("%s/v1beta/models/%s:generateContent", strings.TrimRight(p.base.String(), "/"), url.PathEscape(p.model.ModelName)) + // The key travels in a header rather than the query string so it stays out + // of proxy logs; Gemini accepts both. + raw, err := aiPostJSON(ctx, p.client, endpoint, map[string]string{"x-goog-api-key": p.model.APIKey}, payloadBody) + if err != nil { + return aiChatResult{}, err + } + var payload struct { + Candidates []struct { + Content struct { + Parts []struct { + Text string `json:"text"` + } `json:"parts"` + } `json:"content"` + FinishReason string `json:"finishReason"` + } `json:"candidates"` + UsageMetadata struct { + PromptTokenCount int `json:"promptTokenCount"` + CandidatesTokenCount int `json:"candidatesTokenCount"` + } `json:"usageMetadata"` + Error struct { + Message string `json:"message"` + } `json:"error"` + } + if err = json.Unmarshal(raw, &payload); err != nil { + return aiChatResult{}, fmt.Errorf("模型响应不是有效 JSON:%s", aiErrorSnippet(raw)) + } + if payload.Error.Message != "" { + return aiChatResult{}, fmt.Errorf("模型服务返回错误:%s", payload.Error.Message) + } + if len(payload.Candidates) == 0 { + return aiChatResult{}, fmt.Errorf("模型没有返回任何内容") + } + text := "" + for _, part := range payload.Candidates[0].Content.Parts { + text += part.Text + } + return aiChatResult{ + Text: strings.TrimSpace(text), + PromptTokens: payload.UsageMetadata.PromptTokenCount, + CompletionTokens: payload.UsageMetadata.CandidatesTokenCount, + FinishReason: payload.Candidates[0].FinishReason, + }, nil +} + +// Orchestration platforms (Coze, Dify, self-hosted flows) each have their own +// shape, so they get one normalised contract instead of one adapter each. +type webhookAIProvider struct { + model aiModel + base *url.URL + client *http.Client +} + +func (p webhookAIProvider) Chat(ctx context.Context, req aiChatRequest) (aiChatResult, error) { + messages := make([]map[string]string, 0, len(req.Messages)) + for _, message := range req.Messages { + messages = append(messages, map[string]string{"role": message.Role, "content": message.Text}) + } + headers := map[string]string{} + if strings.TrimSpace(p.model.APIKey) != "" { + headers["Authorization"] = "Bearer " + p.model.APIKey + } + raw, err := aiPostJSON(ctx, p.client, aiEndpoint(p.base, ""), headers, map[string]any{ + "model": p.model.ModelName, + "system": req.System, + "messages": messages, + "maxTokens": aiRequestMaxTokens(req, p.model), + "temperature": aiRequestTemperature(req, p.model), + }) + if err != nil { + return aiChatResult{}, err + } + var payload struct { + Text string `json:"text"` + Reply string `json:"reply"` + Content string `json:"content"` + Usage struct { + PromptTokens int `json:"promptTokens"` + CompletionTokens int `json:"completionTokens"` + } `json:"usage"` + Error string `json:"error"` + } + if err = json.Unmarshal(raw, &payload); err != nil { + return aiChatResult{}, fmt.Errorf("模型响应不是有效 JSON:%s", aiErrorSnippet(raw)) + } + if payload.Error != "" { + return aiChatResult{}, fmt.Errorf("模型服务返回错误:%s", payload.Error) + } + text := payload.Text + if text == "" { + text = payload.Reply + } + if text == "" { + text = payload.Content + } + if strings.TrimSpace(text) == "" { + return aiChatResult{}, fmt.Errorf("Webhook 必须返回 text 字段") + } + return aiChatResult{ + Text: strings.TrimSpace(text), + PromptTokens: payload.Usage.PromptTokens, + CompletionTokens: payload.Usage.CompletionTokens, + }, nil +} + +// Deterministic and offline: local development and tests exercise the whole +// pipeline without credentials or network access. +type debugAIProvider struct{} + +func (debugAIProvider) Chat(_ context.Context, req aiChatRequest) (aiChatResult, error) { + last := "" + for index := len(req.Messages) - 1; index >= 0; index-- { + if req.Messages[index].Role == "user" { + last = strings.TrimSpace(req.Messages[index].Text) + break + } + } + reply := "在的,刚看到消息" + if last != "" { + runes := []rune(last) + if len(runes) > 12 { + runes = runes[:12] + } + reply = "收到:" + string(runes) + } + return aiChatResult{Text: reply, PromptTokens: len([]rune(last)), CompletionTokens: len([]rune(reply)), FinishReason: "stop"}, nil +} diff --git a/im/backend/internal/app/ai_providers_test.go b/im/backend/internal/app/ai_providers_test.go new file mode 100644 index 0000000..b8215d3 --- /dev/null +++ b/im/backend/internal/app/ai_providers_test.go @@ -0,0 +1,218 @@ +package app + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func captureAIRequest(t *testing.T, response string, status int) (*httptest.Server, *map[string]any, *http.Header) { + t.Helper() + body := map[string]any{} + header := http.Header{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + raw, _ := io.ReadAll(io.LimitReader(r.Body, 1<<20)) + _ = json.Unmarshal(raw, &body) + for key, values := range r.Header { + header[key] = values + } + body["__path"] = r.URL.Path + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _, _ = w.Write([]byte(response)) + })) + t.Cleanup(server.Close) + return server, &body, &header +} + +func testModel(protocol, baseURL string) aiModel { + return aiModel{ + ID: 1, Name: "test", Protocol: protocol, BaseURL: baseURL, ModelName: "demo-model", + APIKey: "secret-key", Temperature: 0.5, MaxTokens: 128, TimeoutMS: 5000, Status: 1, + } +} + +func chatOnce(t *testing.T, protocol, baseURL string) (aiChatResult, error) { + t.Helper() + provider, err := newAIProvider(testModel(protocol, baseURL), false) + if err != nil { + t.Fatalf("newAIProvider(%s) error: %v", protocol, err) + } + return provider.Chat(context.Background(), aiChatRequest{ + System: "你是测试人设", + Messages: []aiChatMessage{{Role: "user", Text: "在吗"}, {Role: "assistant", Text: "在的"}, {Role: "user", Text: "吃了吗"}}, + }) +} + +func TestOpenAIProviderRequestAndResponse(t *testing.T) { + server, body, header := captureAIRequest(t, `{"id":"cmpl-1","choices":[{"message":{"content":" 吃过啦 "},"finish_reason":"stop"}],"usage":{"prompt_tokens":12,"completion_tokens":4}}`, http.StatusOK) + result, err := chatOnce(t, "openai", server.URL) + if err != nil { + t.Fatalf("openai chat failed: %v", err) + } + if result.Text != "吃过啦" || result.PromptTokens != 12 || result.CompletionTokens != 4 || result.FinishReason != "stop" { + t.Fatalf("unexpected openai result: %+v", result) + } + if (*body)["__path"] != "/chat/completions" { + t.Fatalf("unexpected openai path: %v", (*body)["__path"]) + } + if header.Get("Authorization") != "Bearer secret-key" { + t.Fatalf("unexpected openai authorization: %q", header.Get("Authorization")) + } + // The system prompt is the first message for this protocol. + messages, _ := (*body)["messages"].([]any) + if len(messages) != 4 { + t.Fatalf("expected system + 3 messages, got %d", len(messages)) + } + first, _ := messages[0].(map[string]any) + if first["role"] != "system" || first["content"] != "你是测试人设" { + t.Fatalf("unexpected first message: %v", first) + } + if (*body)["max_tokens"] != float64(128) || (*body)["temperature"] != 0.5 { + t.Fatalf("model defaults were not applied: %v %v", (*body)["max_tokens"], (*body)["temperature"]) + } +} + +func TestAnthropicProviderUsesTopLevelSystem(t *testing.T) { + server, body, header := captureAIRequest(t, `{"id":"msg-1","content":[{"type":"text","text":"吃过啦"}],"stop_reason":"end_turn","usage":{"input_tokens":9,"output_tokens":3}}`, http.StatusOK) + result, err := chatOnce(t, "anthropic", server.URL) + if err != nil { + t.Fatalf("anthropic chat failed: %v", err) + } + if result.Text != "吃过啦" || result.PromptTokens != 9 || result.CompletionTokens != 3 { + t.Fatalf("unexpected anthropic result: %+v", result) + } + if (*body)["__path"] != "/v1/messages" { + t.Fatalf("unexpected anthropic path: %v", (*body)["__path"]) + } + if header.Get("X-Api-Key") != "secret-key" || header.Get("Anthropic-Version") == "" { + t.Fatalf("unexpected anthropic headers: %v", *header) + } + if (*body)["system"] != "你是测试人设" { + t.Fatalf("system prompt must stay outside the message list: %v", (*body)["system"]) + } + if messages, _ := (*body)["messages"].([]any); len(messages) != 3 { + t.Fatalf("expected 3 messages, got %d", len(messages)) + } +} + +func TestGeminiProviderMapsAssistantToModel(t *testing.T) { + server, body, header := captureAIRequest(t, `{"candidates":[{"content":{"parts":[{"text":"吃过啦"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":7,"candidatesTokenCount":2}}`, http.StatusOK) + result, err := chatOnce(t, "gemini", server.URL) + if err != nil { + t.Fatalf("gemini chat failed: %v", err) + } + if result.Text != "吃过啦" || result.PromptTokens != 7 || result.CompletionTokens != 2 { + t.Fatalf("unexpected gemini result: %+v", result) + } + if (*body)["__path"] != "/v1beta/models/demo-model:generateContent" { + t.Fatalf("unexpected gemini path: %v", (*body)["__path"]) + } + // The key must not travel in the query string where proxies log it. + if header.Get("X-Goog-Api-Key") != "secret-key" { + t.Fatalf("gemini key must be sent as a header: %v", *header) + } + contents, _ := (*body)["contents"].([]any) + if len(contents) != 3 { + t.Fatalf("expected 3 contents, got %d", len(contents)) + } + second, _ := contents[1].(map[string]any) + if second["role"] != "model" { + t.Fatalf("assistant turns must be mapped to role model: %v", second["role"]) + } +} + +func TestWebhookProviderAcceptsAlternateReplyFields(t *testing.T) { + server, _, _ := captureAIRequest(t, `{"reply":"吃过啦","usage":{"promptTokens":5,"completionTokens":2}}`, http.StatusOK) + result, err := chatOnce(t, "webhook", server.URL) + if err != nil { + t.Fatalf("webhook chat failed: %v", err) + } + if result.Text != "吃过啦" || result.PromptTokens != 5 { + t.Fatalf("unexpected webhook result: %+v", result) + } +} + +func TestAIProviderSurfacesUpstreamErrors(t *testing.T) { + server, _, _ := captureAIRequest(t, `{"error":{"message":"insufficient quota"}}`, http.StatusTooManyRequests) + if _, err := chatOnce(t, "openai", server.URL); err == nil || !strings.Contains(err.Error(), "HTTP 429") { + t.Fatalf("upstream status must be reported: %v", err) + } + okServer, _, _ := captureAIRequest(t, `{"error":{"message":"model not found"}}`, http.StatusOK) + _, err := chatOnce(t, "openai", okServer.URL) + if err == nil || !strings.Contains(err.Error(), "model not found") { + t.Fatalf("an error body with HTTP 200 must still fail: %v", err) + } + emptyServer, _, _ := captureAIRequest(t, `{"choices":[]}`, http.StatusOK) + if _, err = chatOnce(t, "openai", emptyServer.URL); err == nil { + t.Fatal("an empty choice list must fail rather than send an empty reply") + } + brokenServer, _, _ := captureAIRequest(t, `not json`, http.StatusOK) + if _, err = chatOnce(t, "openai", brokenServer.URL); err == nil { + t.Fatal("a non-JSON body must fail") + } +} + +func TestParseAIBaseURL(t *testing.T) { + tests := []struct { + protocol string + endpoint string + production bool + valid bool + }{ + {protocol: "openai", endpoint: "https://api.openai.com/v1", production: true, valid: true}, + {protocol: "openai", endpoint: "http://127.0.0.1:11434/v1", production: false, valid: true}, + {protocol: "openai", endpoint: "http://127.0.0.1:11434/v1", production: true, valid: false}, + {protocol: "openai", endpoint: "https://user:pass@api.example.com/v1", production: true, valid: false}, + {protocol: "openai", endpoint: "https://api.example.com/v1?key=secret", production: true, valid: false}, + {protocol: "openai", endpoint: "", production: false, valid: false}, + {protocol: "debug", endpoint: "", production: false, valid: true}, + } + for _, test := range tests { + _, err := parseAIBaseURL(test.protocol, test.endpoint, test.production) + if (err == nil) != test.valid { + t.Errorf("parseAIBaseURL(%q, %q, production=%v) error = %v", test.protocol, test.endpoint, test.production, err) + } + } +} + +func TestAIProviderFactoryGuards(t *testing.T) { + if _, err := newAIProvider(testModel("debug", ""), true); err == nil { + t.Fatal("production must refuse the debug protocol") + } + if _, err := newAIProvider(testModel("unknown", "https://example.com"), false); err == nil { + t.Fatal("unknown protocols must be refused") + } + missingKey := testModel("openai", "https://example.com") + missingKey.APIKey = "" + if _, err := newAIProvider(missingKey, false); err == nil { + t.Fatal("a missing key must be refused before the request is made") + } + // A webhook may be protected by network rules instead of a key. + if _, err := newAIProvider(func() aiModel { m := testModel("webhook", "https://example.com/hook"); m.APIKey = ""; return m }(), false); err != nil { + t.Fatalf("webhook without a key must be allowed: %v", err) + } +} + +func TestDebugProviderIsDeterministicAndOffline(t *testing.T) { + provider, err := newAIProvider(testModel("debug", ""), false) + if err != nil { + t.Fatalf("debug provider unavailable: %v", err) + } + request := aiChatRequest{Messages: []aiChatMessage{{Role: "user", Text: "今天天气怎么样呀在不在"}}} + first, err := provider.Chat(context.Background(), request) + if err != nil { + t.Fatalf("debug chat failed: %v", err) + } + second, _ := provider.Chat(context.Background(), request) + if first.Text != second.Text || first.Text == "" { + t.Fatalf("debug replies must be deterministic: %q vs %q", first.Text, second.Text) + } + if empty, _ := provider.Chat(context.Background(), aiChatRequest{}); empty.Text == "" { + t.Fatal("debug provider must answer even without history") + } +} diff --git a/im/backend/internal/app/app.go b/im/backend/internal/app/app.go index 7d82df3..2ed450b 100644 --- a/im/backend/internal/app/app.go +++ b/im/backend/internal/app/app.go @@ -97,6 +97,23 @@ func New(config Config) (*App, error) { func (a *App) Close() { _ = a.db.Close() } func (a *App) Run() { + workerContext, stopWorkers := context.WithCancel(context.Background()) + defer stopWorkers() + go func() { + a.processDueAccountClosures(workerContext) + ticker := time.NewTicker(time.Minute) + defer ticker.Stop() + for { + select { + case <-workerContext.Done(): + return + case <-ticker.C: + a.processDueAccountClosures(workerContext) + } + } + }() + go a.runAIReplyWorker(workerContext) + go a.runAIMaintenance(workerContext) server := rest.MustNewServer(rest.RestConf{ Host: a.config.Host, Port: a.config.Port, @@ -156,6 +173,7 @@ func (a *App) routes() []rest.Route { {Method: http.MethodPost, Path: "/api/v1/auth/login/sms", Handler: a.loginSMS}, {Method: http.MethodGet, Path: "/api/v1/auth/oauth/providers", Handler: a.userOAuthProviders}, {Method: http.MethodPost, Path: "/api/v1/auth/oauth/start", Handler: a.userOAuthStart}, + {Method: http.MethodPost, Path: "/api/v1/auth/oauth/native", Handler: a.userOAuthNative}, {Method: http.MethodGet, Path: "/api/v1/auth/oauth/callback", Handler: a.oauthCallback}, {Method: http.MethodPost, Path: "/api/v1/auth/oauth/exchange", Handler: a.userOAuthExchange}, {Method: http.MethodPost, Path: "/api/v1/auth/oauth/link", Handler: a.userOAuthLink}, @@ -169,6 +187,7 @@ func (a *App) routes() []rest.Route { {Method: http.MethodHead, Path: "/uploads/:name", Handler: a.serveMedia}, {Method: http.MethodPost, Path: "/admin/v1/auth/login", Handler: a.adminLogin}, {Method: http.MethodPost, Path: "/admin/v1/auth/refresh", Handler: a.adminRefresh}, + {Method: http.MethodPost, Path: "/admin/v1/auth/logout", Handler: a.adminLogout}, {Method: http.MethodGet, Path: "/admin/v1/auth/oauth/providers", Handler: a.adminOAuthProviders}, {Method: http.MethodPost, Path: "/admin/v1/auth/oauth/start", Handler: a.adminOAuthStart}, {Method: http.MethodGet, Path: "/admin/v1/auth/oauth/callback", Handler: a.oauthCallback}, @@ -187,13 +206,32 @@ func (a *App) userRoutes() []rest.Route { {Method: http.MethodGet, Path: "/api/v1/me", Handler: auth(a.me)}, {Method: http.MethodGet, Path: "/api/v1/me/profile", Handler: auth(a.me)}, {Method: http.MethodPatch, Path: "/api/v1/me/profile", Handler: auth(a.updateProfile)}, + {Method: http.MethodPut, Path: "/api/v1/me/password", Handler: auth(a.changeUserPassword)}, + {Method: http.MethodPut, Path: "/api/v1/me/phone", Handler: auth(a.changeUserPhone)}, + {Method: http.MethodGet, Path: "/api/v1/me/devices", Handler: auth(a.myDevices)}, + {Method: http.MethodDelete, Path: "/api/v1/me/devices/:id", Handler: auth(a.revokeDevice)}, + {Method: http.MethodGet, Path: "/api/v1/me/blocks", Handler: auth(a.blockedUsers)}, + {Method: http.MethodGet, Path: "/api/v1/me/notification-settings", Handler: auth(a.notificationSettings)}, + {Method: http.MethodPut, Path: "/api/v1/me/notification-settings", Handler: auth(a.updateNotificationSettings)}, + {Method: http.MethodPost, Path: "/api/v1/me/push-tokens", Handler: auth(a.registerPushToken)}, + {Method: http.MethodDelete, Path: "/api/v1/me/push-tokens", Handler: auth(a.deletePushToken)}, + {Method: http.MethodGet, Path: "/api/v1/me/feedback", Handler: auth(a.feedback)}, + {Method: http.MethodPost, Path: "/api/v1/me/feedback", Handler: auth(a.feedback)}, + {Method: http.MethodGet, Path: "/api/v1/me/account-closure", Handler: auth(a.accountClosure)}, + {Method: http.MethodPost, Path: "/api/v1/me/account-closure", Handler: auth(a.accountClosure)}, + {Method: http.MethodDelete, Path: "/api/v1/me/account-closure", Handler: auth(a.accountClosure)}, + {Method: http.MethodGet, Path: "/api/v1/me/data-export", Handler: auth(a.exportMyData)}, + {Method: http.MethodPost, Path: "/api/v1/me/consents", Handler: auth(a.recordConsent)}, {Method: http.MethodGet, Path: "/api/v1/users/search", Handler: auth(a.searchUsers)}, + {Method: http.MethodGet, Path: "/api/v1/tags", Handler: auth(a.availableTags)}, {Method: http.MethodGet, Path: "/api/v1/users/:id", Handler: auth(a.userProfile)}, {Method: http.MethodGet, Path: "/api/v1/me/following", Handler: auth(a.followingList)}, {Method: http.MethodGet, Path: "/api/v1/me/followers", Handler: auth(a.followerList)}, {Method: http.MethodGet, Path: "/api/v1/me/visitors", Handler: auth(a.visitorList)}, {Method: http.MethodGet, Path: "/api/v1/me/privacy", Handler: auth(a.getPrivacy)}, {Method: http.MethodPut, Path: "/api/v1/me/privacy", Handler: auth(a.updatePrivacy)}, + {Method: http.MethodGet, Path: "/api/v1/me/verification", Handler: auth(a.myVerification)}, + {Method: http.MethodPost, Path: "/api/v1/me/verification", Handler: auth(a.submitVerification)}, {Method: http.MethodGet, Path: "/api/v1/discover/recommendations", Handler: auth(a.discover)}, {Method: http.MethodGet, Path: "/api/v1/nearby/users", Handler: auth(a.nearby)}, {Method: http.MethodPut, Path: "/api/v1/location", Handler: auth(a.updateLocation)}, @@ -207,76 +245,114 @@ func (a *App) userRoutes() []rest.Route { {Method: http.MethodGet, Path: "/api/v1/posts/:id", Handler: auth(a.postDetail)}, {Method: http.MethodGet, Path: "/api/v1/users/:id/posts", Handler: auth(a.userPosts)}, {Method: http.MethodPost, Path: "/api/v1/posts", Handler: auth(a.createPost)}, + {Method: http.MethodPatch, Path: "/api/v1/posts/:id", Handler: auth(a.updatePost)}, + {Method: http.MethodDelete, Path: "/api/v1/posts/:id", Handler: auth(a.deleteOwnPost)}, {Method: http.MethodPost, Path: "/api/v1/media/upload", Handler: auth(a.uploadMedia)}, {Method: http.MethodPost, Path: "/api/v1/posts/:id/like", Handler: auth(a.likePost)}, {Method: http.MethodDelete, Path: "/api/v1/posts/:id/like", Handler: auth(a.unlikePost)}, {Method: http.MethodGet, Path: "/api/v1/posts/:id/comments", Handler: auth(a.comments)}, {Method: http.MethodPost, Path: "/api/v1/posts/:id/comments", Handler: auth(a.createComment)}, + {Method: http.MethodDelete, Path: "/api/v1/comments/:id", Handler: auth(a.deleteOwnComment)}, {Method: http.MethodPost, Path: "/api/v1/im/conversations/direct", Handler: auth(a.directConversation)}, {Method: http.MethodGet, Path: "/api/v1/im/conversations", Handler: auth(a.conversations)}, {Method: http.MethodGet, Path: "/api/v1/im/conversations/:id/messages", Handler: auth(a.messages)}, {Method: http.MethodPost, Path: "/api/v1/im/conversations/:id/messages", Handler: auth(a.sendMessageHTTP)}, + {Method: http.MethodPost, Path: "/api/v1/im/messages/:id/recall", Handler: auth(a.recallMessage)}, {Method: http.MethodPatch, Path: "/api/v1/im/conversations/:id/settings", Handler: auth(a.conversationSettings)}, {Method: http.MethodGet, Path: "/api/v1/membership/status", Handler: auth(a.membershipStatus)}, {Method: http.MethodPost, Path: "/api/v1/orders", Handler: auth(a.createOrder)}, {Method: http.MethodPost, Path: "/api/v1/orders/:id/pay", Handler: auth(a.payOrder)}, {Method: http.MethodGet, Path: "/api/v1/orders/:id", Handler: auth(a.orderStatus)}, + {Method: http.MethodPost, Path: "/api/v1/orders/:id/close", Handler: auth(a.closeOwnOrder)}, + {Method: http.MethodPost, Path: "/api/v1/orders/:id/refund", Handler: auth(a.requestOrderRefund)}, {Method: http.MethodGet, Path: "/api/v1/me/orders", Handler: auth(a.myOrders)}, {Method: http.MethodGet, Path: "/api/v1/notifications", Handler: auth(a.notifications)}, {Method: http.MethodPost, Path: "/api/v1/notifications/read-all", Handler: auth(a.readAllNotifications)}, + {Method: http.MethodPost, Path: "/api/v1/notifications/:id/read", Handler: auth(a.readNotification)}, {Method: http.MethodPost, Path: "/api/v1/reports", Handler: auth(a.createReport)}, + {Method: http.MethodGet, Path: "/api/v1/me/reports", Handler: auth(a.myReports)}, {Method: http.MethodGet, Path: "/ws", Handler: a.websocket}, } } func (a *App) adminRoutes() []rest.Route { auth := func(next http.HandlerFunc) http.HandlerFunc { return a.requireAuth("admin", next) } + permit := func(permission string, next http.HandlerFunc) http.HandlerFunc { + return a.requireAdminPermission(permission, next) + } return []rest.Route{ - {Method: http.MethodPost, Path: "/admin/v1/auth/logout", Handler: auth(a.adminLogout)}, {Method: http.MethodPut, Path: "/admin/v1/me/password", Handler: auth(a.adminChangePassword)}, {Method: http.MethodGet, Path: "/admin/v1/auth/codes", Handler: auth(a.adminCodes)}, {Method: http.MethodGet, Path: "/admin/v1/user/info", Handler: auth(a.adminInfo)}, - {Method: http.MethodGet, Path: "/admin/v1/dashboard/overview", Handler: auth(a.dashboard)}, - {Method: http.MethodGet, Path: "/admin/v1/users", Handler: auth(a.adminUsers)}, - {Method: http.MethodGet, Path: "/admin/v1/users/:id", Handler: auth(a.adminUserDetail)}, - {Method: http.MethodPut, Path: "/admin/v1/users/:id/profile", Handler: auth(a.adminUpdateUserProfile)}, - {Method: http.MethodPut, Path: "/admin/v1/users/:id/verification", Handler: auth(a.adminUpdateVerification)}, - {Method: http.MethodPut, Path: "/admin/v1/users/:id/membership", Handler: auth(a.adminUpdateMembership)}, - {Method: http.MethodPost, Path: "/admin/v1/users/:id/password-reset", Handler: auth(a.adminResetUserPassword)}, - {Method: http.MethodPost, Path: "/admin/v1/users/:id/force-logout", Handler: auth(a.adminForceLogout)}, - {Method: http.MethodGet, Path: "/admin/v1/users/:id/sanctions", Handler: auth(a.adminUserSanctions)}, - {Method: http.MethodPost, Path: "/admin/v1/users/:id/sanctions", Handler: auth(a.adminUserSanctions)}, - {Method: http.MethodPost, Path: "/admin/v1/sanctions/:id/revoke", Handler: auth(a.adminRevokeSanction)}, - {Method: http.MethodPost, Path: "/admin/v1/users/:id/freeze", Handler: auth(a.adminUserStatus)}, - {Method: http.MethodPost, Path: "/admin/v1/users/:id/unfreeze", Handler: auth(a.adminUserStatus)}, - {Method: http.MethodPost, Path: "/admin/v1/users/:id/ban", Handler: auth(a.adminUserStatus)}, - {Method: http.MethodPost, Path: "/admin/v1/users/:id/unban", Handler: auth(a.adminUserStatus)}, - {Method: http.MethodGet, Path: "/admin/v1/posts", Handler: auth(a.adminPosts)}, - {Method: http.MethodGet, Path: "/admin/v1/posts/:id", Handler: auth(a.adminPostDetail)}, - {Method: http.MethodDelete, Path: "/admin/v1/posts/:id", Handler: auth(a.adminDeletePost)}, - {Method: http.MethodGet, Path: "/admin/v1/reports", Handler: auth(a.adminReports)}, - {Method: http.MethodPost, Path: "/admin/v1/reports/:id/handle", Handler: auth(a.adminHandleReport)}, - {Method: http.MethodGet, Path: "/admin/v1/risk/users", Handler: auth(a.adminRiskUsers)}, - {Method: http.MethodGet, Path: "/admin/v1/risk/events", Handler: auth(a.adminRiskEvents)}, - {Method: http.MethodGet, Path: "/admin/v1/membership/plans", Handler: auth(a.adminPlans)}, - {Method: http.MethodPost, Path: "/admin/v1/membership/plans", Handler: auth(a.adminCreatePlan)}, - {Method: http.MethodPatch, Path: "/admin/v1/membership/plans/:id", Handler: auth(a.adminUpdatePlan)}, - {Method: http.MethodPut, Path: "/admin/v1/membership/plans/:id", Handler: auth(a.adminUpdatePlan)}, - {Method: http.MethodDelete, Path: "/admin/v1/membership/plans/:id", Handler: auth(a.adminDeletePlan)}, - {Method: http.MethodGet, Path: "/admin/v1/orders", Handler: auth(a.adminOrders)}, - {Method: http.MethodPut, Path: "/admin/v1/orders/:id", Handler: auth(a.adminUpdateOrder)}, - {Method: http.MethodDelete, Path: "/admin/v1/orders/:id", Handler: auth(a.adminDeleteOrder)}, - {Method: http.MethodPost, Path: "/admin/v1/orders/:id/pay", Handler: auth(a.adminMarkOrderPaid)}, - {Method: http.MethodPost, Path: "/admin/v1/orders/:id/close", Handler: auth(a.adminCloseOrder)}, - {Method: http.MethodPost, Path: "/admin/v1/orders/:id/refund", Handler: auth(a.adminRefund)}, - {Method: http.MethodGet, Path: "/admin/v1/messages", Handler: auth(a.adminMessages)}, - {Method: http.MethodGet, Path: "/admin/v1/system/configs", Handler: auth(a.adminConfigs)}, - {Method: http.MethodPatch, Path: "/admin/v1/system/configs/:key", Handler: auth(a.adminUpdateConfig)}, - {Method: http.MethodPut, Path: "/admin/v1/system/configs/:key", Handler: auth(a.adminUpdateConfig)}, - {Method: http.MethodGet, Path: "/admin/v1/integrations/:group", Handler: auth(a.adminIntegration)}, - {Method: http.MethodPut, Path: "/admin/v1/integrations/:group", Handler: auth(a.adminUpdateIntegration)}, - {Method: http.MethodPost, Path: "/admin/v1/integrations/:group/test", Handler: auth(a.adminTestIntegration)}, - {Method: http.MethodGet, Path: "/admin/v1/audit-logs", Handler: auth(a.adminAuditLogs)}, + {Method: http.MethodGet, Path: "/admin/v1/dashboard/overview", Handler: permit("dashboard:view", a.dashboard)}, + {Method: http.MethodGet, Path: "/admin/v1/users", Handler: permit("users:view", a.adminUsers)}, + {Method: http.MethodPost, Path: "/admin/v1/users", Handler: permit("users:create", a.adminCreateUser)}, + {Method: http.MethodGet, Path: "/admin/v1/users/:id", Handler: permit("users:view", a.adminUserDetail)}, + {Method: http.MethodPut, Path: "/admin/v1/users/:id/profile", Handler: permit("users:manage", a.adminUpdateUserProfile)}, + {Method: http.MethodPut, Path: "/admin/v1/users/:id/verification", Handler: permit("verification:manage", a.adminUpdateVerification)}, + {Method: http.MethodPut, Path: "/admin/v1/users/:id/membership", Handler: permit("users:manage", a.adminUpdateMembership)}, + {Method: http.MethodPost, Path: "/admin/v1/users/:id/password-reset", Handler: permit("users:security", a.adminResetUserPassword)}, + {Method: http.MethodPost, Path: "/admin/v1/users/:id/force-logout", Handler: permit("users:security", a.adminForceLogout)}, + {Method: http.MethodGet, Path: "/admin/v1/users/:id/sanctions", Handler: permit("violations:manage", a.adminUserSanctions)}, + {Method: http.MethodPost, Path: "/admin/v1/users/:id/sanctions", Handler: permit("violations:manage", a.adminUserSanctions)}, + {Method: http.MethodPost, Path: "/admin/v1/sanctions/:id/revoke", Handler: permit("violations:manage", a.adminRevokeSanction)}, + {Method: http.MethodPost, Path: "/admin/v1/users/:id/freeze", Handler: permit("violations:manage", a.adminUserStatus)}, + {Method: http.MethodPost, Path: "/admin/v1/users/:id/unfreeze", Handler: permit("violations:manage", a.adminUserStatus)}, + {Method: http.MethodPost, Path: "/admin/v1/users/:id/ban", Handler: permit("violations:manage", a.adminUserStatus)}, + {Method: http.MethodPost, Path: "/admin/v1/users/:id/unban", Handler: permit("violations:manage", a.adminUserStatus)}, + {Method: http.MethodGet, Path: "/admin/v1/posts", Handler: permit("content:view", a.adminPosts)}, + {Method: http.MethodGet, Path: "/admin/v1/posts/:id", Handler: permit("content:view", a.adminPostDetail)}, + {Method: http.MethodDelete, Path: "/admin/v1/posts/:id", Handler: permit("content:manage", a.adminDeletePost)}, + {Method: http.MethodGet, Path: "/admin/v1/reports", Handler: permit("reports:handle", a.adminReports)}, + {Method: http.MethodPost, Path: "/admin/v1/reports/:id/handle", Handler: permit("reports:handle", a.adminHandleReport)}, + {Method: http.MethodGet, Path: "/admin/v1/risk/users", Handler: permit("risk:view", a.adminRiskUsers)}, + {Method: http.MethodGet, Path: "/admin/v1/risk/events", Handler: permit("risk:view", a.adminRiskEvents)}, + {Method: http.MethodGet, Path: "/admin/v1/membership/plans", Handler: permit("membership:manage", a.adminPlans)}, + {Method: http.MethodPost, Path: "/admin/v1/membership/plans", Handler: permit("membership:manage", a.adminCreatePlan)}, + {Method: http.MethodPatch, Path: "/admin/v1/membership/plans/:id", Handler: permit("membership:manage", a.adminUpdatePlan)}, + {Method: http.MethodPut, Path: "/admin/v1/membership/plans/:id", Handler: permit("membership:manage", a.adminUpdatePlan)}, + {Method: http.MethodDelete, Path: "/admin/v1/membership/plans/:id", Handler: permit("membership:manage", a.adminDeletePlan)}, + {Method: http.MethodGet, Path: "/admin/v1/orders", Handler: permit("orders:view", a.adminOrders)}, + {Method: http.MethodPut, Path: "/admin/v1/orders/:id", Handler: permit("orders:manage", a.adminUpdateOrder)}, + {Method: http.MethodDelete, Path: "/admin/v1/orders/:id", Handler: permit("orders:manage", a.adminDeleteOrder)}, + {Method: http.MethodPost, Path: "/admin/v1/orders/:id/pay", Handler: permit("orders:manage", a.adminMarkOrderPaid)}, + {Method: http.MethodPost, Path: "/admin/v1/orders/:id/close", Handler: permit("orders:manage", a.adminCloseOrder)}, + {Method: http.MethodPost, Path: "/admin/v1/orders/:id/refund", Handler: permit("orders:manage", a.adminRefund)}, + {Method: http.MethodGet, Path: "/admin/v1/messages", Handler: permit("messages:view", a.adminMessages)}, + {Method: http.MethodPost, Path: "/admin/v1/messages/:id/moderate", Handler: permit("messages:manage", a.adminModerateMessage)}, + {Method: http.MethodGet, Path: "/admin/v1/client-feedback", Handler: permit("users:view", a.adminFeedback)}, + {Method: http.MethodPut, Path: "/admin/v1/client-feedback/:id", Handler: permit("users:manage", a.adminFeedback)}, + {Method: http.MethodGet, Path: "/admin/v1/account-closures", Handler: permit("users:view", a.adminAccountClosures)}, + {Method: http.MethodPut, Path: "/admin/v1/account-closures/:id", Handler: permit("users:manage", a.adminAccountClosures)}, + {Method: http.MethodGet, Path: "/admin/v1/app-versions", Handler: permit("system:manage", a.adminAppVersions)}, + {Method: http.MethodPost, Path: "/admin/v1/app-versions", Handler: permit("system:manage", a.adminAppVersions)}, + {Method: http.MethodPut, Path: "/admin/v1/app-versions/:id", Handler: permit("system:manage", a.adminAppVersions)}, + {Method: http.MethodDelete, Path: "/admin/v1/app-versions/:id", Handler: permit("system:manage", a.adminAppVersions)}, + {Method: http.MethodGet, Path: "/admin/v1/system/configs", Handler: permit("system:manage", a.adminConfigs)}, + {Method: http.MethodPatch, Path: "/admin/v1/system/configs/:key", Handler: permit("system:manage", a.adminUpdateConfig)}, + {Method: http.MethodPut, Path: "/admin/v1/system/configs/:key", Handler: permit("system:manage", a.adminUpdateConfig)}, + {Method: http.MethodGet, Path: "/admin/v1/integrations/:group", Handler: permit("system:manage", a.adminIntegration)}, + {Method: http.MethodPut, Path: "/admin/v1/integrations/:group", Handler: permit("system:manage", a.adminUpdateIntegration)}, + {Method: http.MethodPost, Path: "/admin/v1/integrations/:group/test", Handler: permit("system:manage", a.adminTestIntegration)}, + {Method: http.MethodGet, Path: "/admin/v1/ai/models", Handler: permit("system:manage", a.adminAIModels)}, + {Method: http.MethodPost, Path: "/admin/v1/ai/models", Handler: permit("system:manage", a.adminCreateAIModel)}, + {Method: http.MethodPut, Path: "/admin/v1/ai/models/:id", Handler: permit("system:manage", a.adminUpdateAIModel)}, + {Method: http.MethodDelete, Path: "/admin/v1/ai/models/:id", Handler: permit("system:manage", a.adminDeleteAIModel)}, + {Method: http.MethodPost, Path: "/admin/v1/ai/models/:id/test", Handler: permit("system:manage", a.adminTestAIModel)}, + {Method: http.MethodGet, Path: "/admin/v1/ai/agents", Handler: permit("system:manage", a.adminAIAgents)}, + {Method: http.MethodPut, Path: "/admin/v1/ai/agents", Handler: permit("system:manage", a.adminUpdateAIAgents)}, + {Method: http.MethodGet, Path: "/admin/v1/ai/reply-jobs", Handler: permit("system:manage", a.adminAIReplyJobs)}, + {Method: http.MethodGet, Path: "/admin/v1/ai/logs", Handler: permit("system:manage", a.adminAILogs)}, + {Method: http.MethodGet, Path: "/admin/v1/ai/usage", Handler: permit("system:manage", a.adminAIUsage)}, + {Method: http.MethodGet, Path: "/admin/v1/audit-logs", Handler: permit("system:manage", a.adminAuditLogs)}, + {Method: http.MethodGet, Path: "/admin/v1/admin-users", Handler: permit("system:manage", a.adminAccounts)}, + {Method: http.MethodPost, Path: "/admin/v1/admin-users", Handler: permit("system:manage", a.adminCreateAccount)}, + {Method: http.MethodPut, Path: "/admin/v1/admin-users/:id", Handler: permit("system:manage", a.adminUpdateAccount)}, + {Method: http.MethodGet, Path: "/admin/v1/admin-roles", Handler: permit("system:manage", a.adminRoles)}, + {Method: http.MethodPost, Path: "/admin/v1/admin-roles", Handler: permit("system:manage", a.adminCreateRole)}, + {Method: http.MethodPut, Path: "/admin/v1/admin-roles/:id", Handler: permit("system:manage", a.adminUpdateRole)}, + {Method: http.MethodDelete, Path: "/admin/v1/admin-roles/:id", Handler: permit("system:manage", a.adminDeleteRole)}, + {Method: http.MethodGet, Path: "/admin/v1/admin-permissions", Handler: permit("system:manage", a.adminPermissions)}, } } @@ -394,32 +470,9 @@ func validateConfig(config Config) error { return fmt.Errorf("生产环境来源必须使用 HTTPS: %s", origin) } } - if config.BootstrapAdminPassword != "" && !strongAdminPassword(config.BootstrapAdminPassword) { - return fmt.Errorf("IM_BOOTSTRAP_ADMIN_PASSWORD 至少 12 位,且必须包含大小写字母、数字和特殊字符") - } return nil } -func strongAdminPassword(value string) bool { - if len(value) < 12 { - return false - } - var lower, upper, digit, special bool - for _, char := range value { - switch { - case char >= 'a' && char <= 'z': - lower = true - case char >= 'A' && char <= 'Z': - upper = true - case char >= '0' && char <= '9': - digit = true - default: - special = true - } - } - return lower && upper && digit && special -} - func (a *App) originAllowed(origin string) bool { origin = strings.TrimRight(strings.TrimSpace(origin), "/") if origin == "" { diff --git a/im/backend/internal/app/im.go b/im/backend/internal/app/im.go index a90514f..31a5244 100644 --- a/im/backend/internal/app/im.go +++ b/im/backend/internal/app/im.go @@ -1,14 +1,20 @@ package app import ( + "context" "database/sql" "encoding/json" "errors" "fmt" "net/http" + "net/url" + "regexp" "sort" + "strconv" + "strings" "sync" "time" + "unicode/utf8" "github.com/gorilla/websocket" ) @@ -21,15 +27,107 @@ type messageView struct { ClientMsgID string `json:"clientMsgId"` Type int `json:"type"` Content any `json:"content"` + Recalled bool `json:"recalled"` CreatedAt time.Time `json:"createdAt"` } +const ( + maxMessageBodyBytes = 16 << 10 + maxMessageTextRunes = 2000 +) + +var clientMessageIDPattern = regexp.MustCompile(`^[A-Za-z0-9_.:-]{1,64}$`) + +func validMessageMediaURL(raw string, production bool) bool { + raw = strings.TrimSpace(raw) + if raw == "" || len(raw) > 2048 { + return false + } + if strings.HasPrefix(raw, "/uploads/") { + return !strings.Contains(raw, "..") + } + parsed, err := url.Parse(raw) + if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { + return false + } + return !production || parsed.Scheme == "https" +} + +func numericDuration(value any) (int, bool) { + switch duration := value.(type) { + case float64: + return int(duration), duration == float64(int(duration)) + case int: + return duration, true + case int64: + return int(duration), true + case json.Number: + parsed, err := strconv.Atoi(duration.String()) + return parsed, err == nil + default: + return 0, false + } +} + +func (a *App) validateMessagePayload(clientMsgID string, messageType int, content any) (string, any, []byte, error) { + clientMsgID = strings.TrimSpace(clientMsgID) + if clientMsgID == "" { + clientMsgID = randomToken()[:26] + } + if !clientMessageIDPattern.MatchString(clientMsgID) { + return "", nil, nil, fmt.Errorf("客户端消息 ID 格式无效") + } + if messageType == 0 { + messageType = 1 + } + data, ok := content.(map[string]any) + if !ok { + return "", nil, nil, fmt.Errorf("消息内容格式错误") + } + switch messageType { + case 1: + text, ok := data["text"].(string) + text = strings.TrimSpace(text) + if !ok || text == "" || !utf8.ValidString(text) || utf8.RuneCountInString(text) > maxMessageTextRunes { + return "", nil, nil, fmt.Errorf("文本消息应为 1 至 %d 个字符", maxMessageTextRunes) + } + data = map[string]any{"text": text} + case 2: + mediaURL, ok := data["url"].(string) + if !ok || !validMessageMediaURL(mediaURL, a.config.Environment == "production") { + return "", nil, nil, fmt.Errorf("图片地址无效") + } + data = map[string]any{"url": strings.TrimSpace(mediaURL)} + case 3: + mediaURL, ok := data["url"].(string) + duration, durationOK := numericDuration(data["duration"]) + if !ok || !validMessageMediaURL(mediaURL, a.config.Environment == "production") || !durationOK || duration < 1 || duration > 60 { + return "", nil, nil, fmt.Errorf("语音消息地址或时长无效") + } + data = map[string]any{"duration": duration, "url": strings.TrimSpace(mediaURL)} + default: + return "", nil, nil, fmt.Errorf("不支持的消息类型") + } + body, err := json.Marshal(data) + if err != nil || len(body) > maxMessageBodyBytes { + return "", nil, nil, fmt.Errorf("消息内容过大") + } + return clientMsgID, data, body, nil +} + type wsClient struct { userID int64 conn *websocket.Conn mu sync.Mutex } +func (c *wsClient) writeJSON(payload any) error { + c.mu.Lock() + defer c.mu.Unlock() + _ = c.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) + return c.conn.WriteJSON(payload) +} + type Hub struct { mu sync.RWMutex clients map[int64]map[*wsClient]struct{} @@ -37,22 +135,47 @@ type Hub struct { func NewHub() *Hub { return &Hub{clients: make(map[int64]map[*wsClient]struct{})} } -func (h *Hub) add(client *wsClient) { +func (h *Hub) add(client *wsClient) bool { h.mu.Lock() defer h.mu.Unlock() + becameOnline := len(h.clients[client.userID]) == 0 if h.clients[client.userID] == nil { h.clients[client.userID] = make(map[*wsClient]struct{}) } h.clients[client.userID][client] = struct{}{} + return becameOnline } -func (h *Hub) remove(client *wsClient) { +func (h *Hub) remove(client *wsClient) bool { h.mu.Lock() defer h.mu.Unlock() + if _, exists := h.clients[client.userID][client]; !exists { + return false + } delete(h.clients[client.userID], client) if len(h.clients[client.userID]) == 0 { delete(h.clients, client.userID) + return true } + return false +} + +func (h *Hub) online(userID int64) bool { + h.mu.RLock() + defer h.mu.RUnlock() + return len(h.clients[userID]) > 0 +} + +func (h *Hub) onlineUsers() []int64 { + h.mu.RLock() + defer h.mu.RUnlock() + users := make([]int64, 0, len(h.clients)) + for userID, clients := range h.clients { + if len(clients) > 0 { + users = append(users, userID) + } + } + return users } func (h *Hub) broadcast(userIDs []int64, payload any) { @@ -65,9 +188,7 @@ func (h *Hub) broadcast(userIDs []int64, payload any) { } h.mu.RUnlock() for _, client := range targets { - client.mu.Lock() - _ = client.conn.WriteJSON(payload) - client.mu.Unlock() + _ = client.writeJSON(payload) } } @@ -86,6 +207,41 @@ func (h *Hub) disconnect(userID int64) { } } +func (a *App) onlineNow(userID int64) bool { + return a.hub.online(userID) +} + +func (a *App) publishPresence(userID int64, online bool) { + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + var visible int + if err := a.db.QueryRowContext(ctx, `SELECT online_visible FROM user_privacy_settings WHERE user_id=?`, userID).Scan(&visible); err != nil { + return + } + rows, err := a.db.QueryContext(ctx, `SELECT DISTINCT peer.user_id + FROM im_conversation_members mine JOIN im_conversation_members peer + ON peer.conversation_id=mine.conversation_id AND peer.user_id<>mine.user_id + JOIN im_conversations conversation ON conversation.id=mine.conversation_id + WHERE mine.user_id=? AND mine.status=1 AND peer.status=1 AND conversation.status=1`, userID) + if err != nil { + return + } + defer rows.Close() + peers := []int64{} + for rows.Next() { + var peerID int64 + if rows.Scan(&peerID) == nil { + peers = append(peers, peerID) + } + } + if len(peers) == 0 { + return + } + a.hub.broadcast(peers, map[string]any{"command": "PRESENCE", "data": map[string]any{ + "userId": userID, "online": online && visible == 1, "lastActiveAt": time.Now(), + }}) +} + func (a *App) directConversation(w http.ResponseWriter, r *http.Request) { var req struct { UserID int64 `json:"userId"` @@ -98,16 +254,47 @@ func (a *App) directConversation(w http.ResponseWriter, r *http.Request) { if first > second { first, second = second, first } + var targetStatus, allowStranger int + if err := a.db.QueryRowContext(r.Context(), `SELECT u.status,p.allow_stranger_message FROM users u JOIN user_privacy_settings p ON p.user_id=u.id WHERE u.id=? AND u.deleted_at IS NULL`, req.UserID).Scan(&targetStatus, &allowStranger); err != nil || targetStatus != 1 { + fail(w, http.StatusNotFound, 30001, "聊天对象不存在或不可用") + return + } + var blocked int + if err := a.db.QueryRowContext(r.Context(), `SELECT EXISTS(SELECT 1 FROM user_blocks WHERE (user_id=? AND blocked_user_id=?) OR (user_id=? AND blocked_user_id=?))`, current(r).ID, req.UserID, req.UserID, current(r).ID).Scan(&blocked); err != nil || blocked == 1 { + fail(w, http.StatusForbidden, 30002, "当前无法与该用户聊天") + return + } var id int64 err := a.db.QueryRowContext(r.Context(), `SELECT conversation_id FROM im_direct_conversations WHERE user1_id=? AND user2_id=?`, first, second).Scan(&id) if err == nil { reply(w, map[string]any{"id": id}) return } - tx, _ := a.db.BeginTx(r.Context(), nil) + if allowStranger == 0 { + var mutualFollow int + _ = a.db.QueryRowContext(r.Context(), `SELECT EXISTS( + SELECT 1 FROM user_follows mine JOIN user_follows target + ON target.user_id=mine.target_user_id AND target.target_user_id=mine.user_id + WHERE mine.user_id=? AND mine.target_user_id=?)`, current(r).ID, req.UserID).Scan(&mutualFollow) + if mutualFollow != 1 { + fail(w, http.StatusForbidden, 30002, "对方仅允许互相关注的人发起私信") + return + } + } + tx, beginErr := a.db.BeginTx(r.Context(), nil) + if beginErr != nil { + fail(w, http.StatusInternalServerError, 50001, "创建会话失败") + return + } + defer func() { _ = tx.Rollback() }() result, err := tx.ExecContext(r.Context(), `INSERT INTO im_conversations(conversation_type)VALUES(1)`) if err != nil { _ = tx.Rollback() + // Concurrent requests may have created the same unique direct pair. + if a.db.QueryRowContext(r.Context(), `SELECT conversation_id FROM im_direct_conversations WHERE user1_id=? AND user2_id=?`, first, second).Scan(&id) == nil { + reply(w, map[string]any{"id": id}) + return + } fail(w, 500, 50001, "创建会话失败") return } @@ -127,12 +314,16 @@ func (a *App) directConversation(w http.ResponseWriter, r *http.Request) { func (a *App) conversations(w http.ResponseWriter, r *http.Request) { who := current(r) - rows, err := a.db.QueryContext(r.Context(), `SELECT c.id,c.last_seq,c.last_message_at,m.read_seq,m.pinned,m.muted, - other.user_id,p.nickname,p.avatar_url,p.is_vip,p.last_active_at,COALESCE(CAST(msg.body AS CHAR CHARACTER SET utf8mb4),'') + // A sequence is shared by both senders, so last_seq-read_seq is not an unread count. + rows, err := a.db.QueryContext(r.Context(), `SELECT c.id,c.last_seq,c.last_message_at, + (SELECT COUNT(*) FROM im_messages unread_msg WHERE unread_msg.conversation_id=c.id + AND unread_msg.seq>GREATEST(m.read_seq,m.clear_seq,m.join_seq) AND unread_msg.sender_id<>m.user_id + AND unread_msg.recalled_at IS NULL AND unread_msg.admin_removed_at IS NULL),m.pinned,m.muted, + other.user_id,p.nickname,p.avatar_url,p.is_vip,p.last_active_at,privacy.online_visible,COALESCE(CAST(msg.body AS CHAR CHARACTER SET utf8mb4),''),msg.recalled_at,msg.admin_removed_at FROM im_conversation_members m JOIN im_conversations c ON c.id=m.conversation_id JOIN im_conversation_members other ON other.conversation_id=c.id AND other.user_id<>m.user_id - JOIN user_profiles p ON p.user_id=other.user_id LEFT JOIN im_messages msg ON msg.id=c.last_message_id - WHERE m.user_id=? AND m.status=1 ORDER BY m.pinned DESC,c.last_message_at DESC`, who.ID) + JOIN user_profiles p ON p.user_id=other.user_id JOIN user_privacy_settings privacy ON privacy.user_id=other.user_id LEFT JOIN im_messages msg ON msg.id=c.last_message_id + WHERE m.user_id=? AND m.status=1 AND c.status=1 ORDER BY m.pinned DESC,c.last_message_at DESC`, who.ID) if err != nil { fail(w, 500, 50001, err.Error()) return @@ -140,20 +331,31 @@ func (a *App) conversations(w http.ResponseWriter, r *http.Request) { defer rows.Close() items := []map[string]any{} for rows.Next() { - var id, lastSeq, readSeq, otherID int64 + var id, lastSeq, unread, otherID int64 var lastAt sql.NullTime var pinned, muted, vip int var nick, avatar, body string var active sql.NullTime - _ = rows.Scan(&id, &lastSeq, &lastAt, &readSeq, &pinned, &muted, &otherID, &nick, &avatar, &vip, &active, &body) + var recalledAt, adminRemovedAt sql.NullTime + var onlineVisible int + if err := rows.Scan(&id, &lastSeq, &lastAt, &unread, &pinned, &muted, &otherID, &nick, &avatar, &vip, &active, &onlineVisible, &body, &recalledAt, &adminRemovedAt); err != nil { + fail(w, 500, 50001, "读取会话失败") + return + } preview := "开始聊天吧" var content map[string]any - if json.Unmarshal([]byte(body), &content) == nil { + if recalledAt.Valid || adminRemovedAt.Valid { + preview = "消息已撤回" + } else if json.Unmarshal([]byte(body), &content) == nil { if text, ok := content["text"].(string); ok { preview = text } } - items = append(items, map[string]any{"id": id, "lastSeq": lastSeq, "unread": max64(lastSeq-readSeq, 0), "lastMessageAt": lastAt, "pinned": pinned == 1, "muted": muted == 1, "lastMessage": preview, "user": map[string]any{"id": otherID, "nickname": nick, "avatar": avatar, "vip": vip == 1, "online": active.Valid && time.Since(active.Time) < 15*time.Minute}}) + items = append(items, map[string]any{"id": id, "lastSeq": lastSeq, "unread": unread, "lastMessageAt": nullableTime(lastAt), "pinned": pinned == 1, "muted": muted == 1, "lastMessage": preview, "user": map[string]any{"id": otherID, "nickname": nick, "avatar": avatar, "avatarThumbnail": avatarThumbnailURL(avatar), "vip": vip == 1, "online": onlineVisible == 1 && a.onlineNow(otherID)}}) + } + if err := rows.Err(); err != nil { + fail(w, 500, 50001, "读取会话失败") + return } reply(w, map[string]any{"items": items, "total": len(items)}) } @@ -168,7 +370,31 @@ func (a *App) messages(w http.ResponseWriter, r *http.Request) { fail(w, 403, 30002, "不是会话成员") return } - rows, err := a.db.QueryContext(r.Context(), `SELECT id,conversation_id,seq,sender_id,client_msg_id,message_type,body,created_at FROM im_messages WHERE conversation_id=? AND recalled_at IS NULL ORDER BY seq DESC LIMIT 100`, id) + beforeSeq, _ := strconv.ParseInt(r.URL.Query().Get("beforeSeq"), 10, 64) + afterSeq, _ := strconv.ParseInt(r.URL.Query().Get("afterSeq"), 10, 64) + if beforeSeq > 0 && afterSeq > 0 { + fail(w, http.StatusBadRequest, 20001, "beforeSeq 和 afterSeq 不能同时使用") + return + } + limit, _ := strconv.Atoi(r.URL.Query().Get("limit")) + if limit <= 0 || limit > 100 { + limit = 50 + } + query := `SELECT id,conversation_id,seq,sender_id,client_msg_id,message_type,body,recalled_at,admin_removed_at,created_at FROM im_messages WHERE conversation_id=?` + args := []any{id} + if beforeSeq > 0 { + query += ` AND seq 0 { + query += ` AND seq>?` + args = append(args, afterSeq) + query += ` ORDER BY seq ASC LIMIT ?` + } else { + query += ` ORDER BY seq DESC LIMIT ?` + } + args = append(args, limit+1) + rows, err := a.db.QueryContext(r.Context(), query, args...) if err != nil { fail(w, 500, 50001, "查询消息失败") return @@ -178,18 +404,47 @@ func (a *App) messages(w http.ResponseWriter, r *http.Request) { for rows.Next() { var item messageView var body []byte - _ = rows.Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &item.CreatedAt) + var recalledAt, adminRemovedAt sql.NullTime + if err := rows.Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &recalledAt, &adminRemovedAt, &item.CreatedAt); err != nil { + fail(w, 500, 50001, "读取消息失败") + return + } var content any - _ = json.Unmarshal(body, &content) + item.Recalled = recalledAt.Valid || adminRemovedAt.Valid + if !item.Recalled { + _ = json.Unmarshal(body, &content) + } item.Content = content items = append(items, item) } - sort.Slice(items, func(i, j int) bool { return items[i].Seq < items[j].Seq }) + if err := rows.Err(); err != nil { + fail(w, 500, 50001, "读取消息失败") + return + } + _ = rows.Close() + hasMore := len(items) > limit + if hasMore { + items = items[:limit] + } + if afterSeq <= 0 { + sort.Slice(items, func(i, j int) bool { return items[i].Seq < items[j].Seq }) + } if len(items) > 0 { last := items[len(items)-1].Seq - _, _ = a.db.ExecContext(r.Context(), `UPDATE im_conversation_members SET read_seq=GREATEST(read_seq,?) WHERE conversation_id=? AND user_id=?`, last, id, current(r).ID) + if _, _, err := a.markConversationRead(r.Context(), id, current(r).ID, last); err != nil { + fail(w, 500, 50001, "更新已读状态失败") + return + } } - reply(w, map[string]any{"items": items, "hasMore": false}) + nextBeforeSeq := int64(0) + if len(items) > 0 { + nextBeforeSeq = items[0].Seq + } + nextAfterSeq := afterSeq + if len(items) > 0 { + nextAfterSeq = items[len(items)-1].Seq + } + reply(w, map[string]any{"items": items, "hasMore": hasMore, "nextBeforeSeq": nextBeforeSeq, "nextAfterSeq": nextAfterSeq}) } func (a *App) sendMessageHTTP(w http.ResponseWriter, r *http.Request) { @@ -207,16 +462,6 @@ func (a *App) sendMessageHTTP(w http.ResponseWriter, r *http.Request) { fail(w, 400, 20001, "消息格式错误") return } - if req.ClientMsgID == "" { - req.ClientMsgID = fmt.Sprintf("%026d", time.Now().UnixNano()) - } - if len(req.ClientMsgID) > 64 { - fail(w, 400, 20001, "客户端消息 ID 长度不能超过 64 个字符") - return - } - if req.Type == 0 { - req.Type = 1 - } item, members, err := a.persistMessage(r, conversationID, current(r).ID, req.ClientMsgID, req.Type, req.Content) if err != nil { var limitErr *dailyActiveChatLimitError @@ -231,78 +476,195 @@ func (a *App) sendMessageHTTP(w http.ResponseWriter, r *http.Request) { reply(w, item) } -func (a *App) persistMessage(r *http.Request, conversationID, senderID int64, clientMsgID string, messageType int, content any) (messageView, []int64, error) { - item := messageView{} - if !a.allowRequest(r.Context(), "message_send", fmt.Sprintf("%d", senderID), 120, time.Minute) { - return item, nil, fmt.Errorf("消息发送过于频繁,请稍后再试") +func (a *App) recallMessage(w http.ResponseWriter, r *http.Request) { + messageID, err := pathID(r) + if err != nil { + fail(w, http.StatusBadRequest, 20001, "消息编号无效") + return } - if a.isSanctionActive(r.Context(), senderID, "MUTE") { - return item, nil, fmt.Errorf("账号处于禁言期,暂时无法发送消息") + window, _ := strconv.Atoi(a.configPlain(r.Context(), "im.recall_seconds", "120")) + if window < 1 || window > 86400 { + window = 120 } tx, err := a.db.BeginTx(r.Context(), nil) + if err != nil { + fail(w, 500, 50001, "撤回失败") + return + } + defer func() { _ = tx.Rollback() }() + var conversationID, seq int64 + var senderID int64 + var created time.Time + var recalled sql.NullTime + if err = tx.QueryRowContext(r.Context(), `SELECT conversation_id,seq,sender_id,created_at,recalled_at FROM im_messages WHERE id=? FOR UPDATE`, messageID).Scan(&conversationID, &seq, &senderID, &created, &recalled); err != nil { + fail(w, 404, 30001, "消息不存在") + return + } + if senderID != current(r).ID { + fail(w, 403, 30002, "只能撤回自己发送的消息") + return + } + if recalled.Valid { + reply(w, map[string]bool{"success": true}) + return + } + if time.Since(created) > time.Duration(window)*time.Second { + fail(w, 409, 30004, "已超过消息撤回时限") + return + } + _, err = tx.ExecContext(r.Context(), `UPDATE im_messages SET recalled_at=NOW(3) WHERE id=?`, messageID) + if err != nil || tx.Commit() != nil { + fail(w, 500, 50001, "撤回失败") + return + } + rows, _ := a.db.QueryContext(r.Context(), `SELECT user_id FROM im_conversation_members WHERE conversation_id=? AND status=1`, conversationID) + members := []int64{} + if rows != nil { + defer rows.Close() + for rows.Next() { + var id int64 + _ = rows.Scan(&id) + members = append(members, id) + } + } + a.hub.broadcast(members, map[string]any{"command": "MESSAGE_RECALLED", "data": map[string]any{"id": messageID, "conversationId": conversationID, "seq": seq}}) + reply(w, map[string]bool{"success": true}) +} + +// The reply worker has no *http.Request, so message writing is expressed in +// terms of a context and the HTTP handlers wrap it. +func (a *App) persistMessage(r *http.Request, conversationID, senderID int64, clientMsgID string, messageType int, content any) (messageView, []int64, error) { + return a.persistMessageContext(r.Context(), conversationID, senderID, clientMsgID, messageType, content) +} + +func (a *App) persistMessageContext(ctx context.Context, conversationID, senderID int64, clientMsgID string, messageType int, content any) (messageView, []int64, error) { + item := messageView{} + clientMsgID, content, body, err := a.validateMessagePayload(clientMsgID, messageType, content) + if err != nil { + return item, nil, err + } + if messageType == 0 { + messageType = 1 + } + if messageType == 2 || messageType == 3 { + contentMap, _ := content.(map[string]any) + mediaURL, _ := contentMap["url"].(string) + expectedType := "image" + if messageType == 3 { + expectedType = "audio" + } + var owned int + if queryErr := a.db.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM media_assets WHERE owner_user_id=? AND public_url=? AND media_type=? AND status=1 AND moderation_status=1)`, senderID, mediaURL, expectedType).Scan(&owned); queryErr != nil || owned != 1 { + return item, nil, fmt.Errorf("消息媒体必须由当前账号上传") + } + } + if !a.allowRequest(ctx, "message_send", fmt.Sprintf("%d", senderID), 120, time.Minute) { + return item, nil, fmt.Errorf("消息发送过于频繁,请稍后再试") + } + if a.isSanctionActive(ctx, senderID, "MUTE") { + return item, nil, fmt.Errorf("账号处于禁言期,暂时无法发送消息") + } + tx, err := a.db.BeginTx(ctx, nil) if err != nil { return item, nil, err } defer func() { _ = tx.Rollback() }() var lastSeq int64 - if err = tx.QueryRowContext(r.Context(), `SELECT last_seq FROM im_conversations WHERE id=? AND status=1 FOR UPDATE`, conversationID).Scan(&lastSeq); err != nil { + if err = tx.QueryRowContext(ctx, `SELECT last_seq FROM im_conversations WHERE id=? AND status=1 FOR UPDATE`, conversationID).Scan(&lastSeq); err != nil { return item, nil, fmt.Errorf("会话不存在") } var memberCount int - if err = tx.QueryRowContext(r.Context(), `SELECT COUNT(*) FROM im_conversation_members WHERE conversation_id=? AND user_id=? AND status=1`, conversationID, senderID).Scan(&memberCount); err != nil || memberCount == 0 { + if err = tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM im_conversation_members WHERE conversation_id=? AND user_id=? AND status=1`, conversationID, senderID).Scan(&memberCount); err != nil || memberCount == 0 { return item, nil, fmt.Errorf("不是会话成员") } - if err = a.reserveDailyActiveChat(r.Context(), tx, conversationID, senderID); err != nil { + var blocked int + if err = tx.QueryRowContext(ctx, `SELECT EXISTS( + SELECT 1 FROM im_direct_conversations d JOIN user_blocks b + ON (b.user_id=d.user1_id AND b.blocked_user_id=d.user2_id) OR (b.user_id=d.user2_id AND b.blocked_user_id=d.user1_id) + WHERE d.conversation_id=?)`, conversationID).Scan(&blocked); err != nil { return item, nil, err } - body, err := json.Marshal(content) - if err != nil { - return item, nil, fmt.Errorf("消息内容错误") + if blocked == 1 { + return item, nil, fmt.Errorf("当前无法向该用户发送消息") + } + if err = a.reserveDailyActiveChat(ctx, tx, conversationID, senderID); err != nil { + return item, nil, err } seq := lastSeq + 1 - result, err := tx.ExecContext(r.Context(), `INSERT INTO im_messages(conversation_id,seq,sender_id,client_msg_id,message_type,body)VALUES(?,?,?,?,?,?)`, conversationID, seq, senderID, clientMsgID, messageType, body) + result, err := tx.ExecContext(ctx, `INSERT INTO im_messages(conversation_id,seq,sender_id,client_msg_id,message_type,body)VALUES(?,?,?,?,?,?)`, conversationID, seq, senderID, clientMsgID, messageType, body) if err != nil { var existingID, existingSeq int64 - existingErr := tx.QueryRowContext(r.Context(), `SELECT id,seq FROM im_messages WHERE sender_id=? AND client_msg_id=?`, senderID, clientMsgID).Scan(&existingID, &existingSeq) + existingErr := tx.QueryRowContext(ctx, `SELECT id,seq FROM im_messages WHERE sender_id=? AND client_msg_id=?`, senderID, clientMsgID).Scan(&existingID, &existingSeq) if existingErr == nil { _ = tx.Rollback() - return a.loadMessage(r, existingID), nil, nil + return a.loadMessageContext(ctx, existingID), nil, nil } return item, nil, err } messageID, _ := result.LastInsertId() - _, err = tx.ExecContext(r.Context(), `UPDATE im_conversations SET last_seq=?,last_message_id=?,last_message_at=NOW(3) WHERE id=?`, seq, messageID, conversationID) + _, err = tx.ExecContext(ctx, `UPDATE im_conversations SET last_seq=?,last_message_id=?,last_message_at=NOW(3) WHERE id=?`, seq, messageID, conversationID) if err != nil { return item, nil, err } - _, _ = tx.ExecContext(r.Context(), `UPDATE im_conversation_members SET delivered_seq=GREATEST(delivered_seq,?),updated_at=NOW(3) WHERE conversation_id=?`, seq, conversationID) - memberRows, err := tx.QueryContext(r.Context(), `SELECT user_id FROM im_conversation_members WHERE conversation_id=? AND status=1`, conversationID) + if _, err = tx.ExecContext(ctx, `UPDATE im_conversation_members SET delivered_seq=GREATEST(delivered_seq,?),updated_at=NOW(3) WHERE conversation_id=?`, seq, conversationID); err != nil { + return item, nil, err + } + memberRows, err := tx.QueryContext(ctx, `SELECT user_id FROM im_conversation_members WHERE conversation_id=? AND status=1`, conversationID) if err != nil { return item, nil, err } members := []int64{} for memberRows.Next() { var userID int64 - _ = memberRows.Scan(&userID) + if err = memberRows.Scan(&userID); err != nil { + _ = memberRows.Close() + return item, nil, err + } members = append(members, userID) } - _ = memberRows.Close() + if err = memberRows.Err(); err != nil { + _ = memberRows.Close() + return item, nil, err + } + if err = memberRows.Close(); err != nil { + return item, nil, err + } + sort.Slice(members, func(i, j int) bool { return members[i] < members[j] }) + for _, userID := range members { + var lockedUserID int64 + if err = tx.QueryRowContext(ctx, `SELECT id FROM users WHERE id=? FOR UPDATE`, userID).Scan(&lockedUserID); err != nil { + return item, nil, err + } + } for _, userID := range members { var next int64 - _ = tx.QueryRowContext(r.Context(), `SELECT COALESCE(MAX(event_seq),0)+1 FROM im_user_sync_events WHERE user_id=?`, userID).Scan(&next) - _, _ = tx.ExecContext(r.Context(), `INSERT INTO im_user_sync_events(user_id,event_seq,event_type,conversation_id,message_seq,event_data)VALUES(?, ?,12,?,?,?)`, userID, next, conversationID, seq, body) + if err = tx.QueryRowContext(ctx, `SELECT COALESCE(MAX(event_seq),0)+1 FROM im_user_sync_events WHERE user_id=?`, userID).Scan(&next); err != nil { + return item, nil, err + } + if _, err = tx.ExecContext(ctx, `INSERT INTO im_user_sync_events(user_id,event_seq,event_type,conversation_id,message_seq,event_data)VALUES(?, ?,12,?,?,?)`, userID, next, conversationID, seq, body); err != nil { + return item, nil, err + } } if err = tx.Commit(); err != nil { return item, nil, err } + a.enqueueAIReply(ctx, conversationID, senderID, members, messageID) return messageView{ID: messageID, ConversationID: conversationID, Seq: seq, SenderID: senderID, ClientMsgID: clientMsgID, Type: messageType, Content: content, CreatedAt: time.Now()}, members, nil } func (a *App) loadMessage(r *http.Request, id int64) messageView { + return a.loadMessageContext(r.Context(), id) +} + +func (a *App) loadMessageContext(ctx context.Context, id int64) messageView { var item messageView var body []byte - _ = a.db.QueryRowContext(r.Context(), `SELECT id,conversation_id,seq,sender_id,client_msg_id,message_type,body,created_at FROM im_messages WHERE id=?`, id).Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &item.CreatedAt) - _ = json.Unmarshal(body, &item.Content) + var recalledAt, adminRemovedAt sql.NullTime + _ = a.db.QueryRowContext(ctx, `SELECT id,conversation_id,seq,sender_id,client_msg_id,message_type,body,recalled_at,admin_removed_at,created_at FROM im_messages WHERE id=?`, id).Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &recalledAt, &adminRemovedAt, &item.CreatedAt) + item.Recalled = recalledAt.Valid || adminRemovedAt.Valid + if !item.Recalled { + _ = json.Unmarshal(body, &item.Content) + } return item } func (a *App) isMember(r *http.Request, conversationID, userID int64) bool { @@ -312,13 +674,17 @@ func (a *App) isMember(r *http.Request, conversationID, userID int64) bool { } func (a *App) conversationSettings(w http.ResponseWriter, r *http.Request) { - id, _ := pathID(r) + id, pathErr := pathID(r) + if pathErr != nil { + fail(w, http.StatusBadRequest, 20001, "会话 ID 无效") + return + } var req struct { Pinned *bool `json:"pinned"` Muted *bool `json:"muted"` ReadSeq int64 `json:"readSeq"` } - if decode(r, &req) != nil { + if decode(r, &req) != nil || req.ReadSeq < 0 { fail(w, 400, 20001, "invalid settings") return } @@ -329,12 +695,72 @@ func (a *App) conversationSettings(w http.ResponseWriter, r *http.Request) { if req.Muted != nil { muted = btoi(*req.Muted) } - _, err := a.db.ExecContext(r.Context(), `UPDATE im_conversation_members SET pinned=IF(?>=0,?,pinned),muted=IF(?>=0,?,muted),read_seq=GREATEST(read_seq,?) WHERE conversation_id=? AND user_id=?`, pinned, pinned, muted, muted, req.ReadSeq, id, current(r).ID) + var lastSeq int64 + if err := a.db.QueryRowContext(r.Context(), `SELECT c.last_seq FROM im_conversations c JOIN im_conversation_members m ON m.conversation_id=c.id WHERE c.id=? AND m.user_id=? AND m.status=1`, id, current(r).ID).Scan(&lastSeq); err != nil { + fail(w, http.StatusForbidden, 30002, "不是会话成员") + return + } + readSeq := req.ReadSeq + if readSeq > lastSeq { + readSeq = lastSeq + } + _, err := a.db.ExecContext(r.Context(), `UPDATE im_conversation_members SET pinned=IF(?>=0,?,pinned),muted=IF(?>=0,?,muted),read_seq=GREATEST(read_seq,?) WHERE conversation_id=? AND user_id=?`, pinned, pinned, muted, muted, readSeq, id, current(r).ID) if err != nil { fail(w, 500, 50001, "保存失败") return } - reply(w, map[string]bool{"success": true}) + savedReadSeq, unread, err := a.conversationReadState(r.Context(), id, current(r).ID) + if err != nil { + fail(w, 500, 50001, "读取已读状态失败") + return + } + if req.ReadSeq > 0 { + a.publishConversationRead(id, current(r).ID, savedReadSeq, unread) + } + reply(w, map[string]any{"success": true, "readSeq": savedReadSeq, "unread": unread}) +} + +// Notify this user's other devices after a successful read; never notify the peer +// or advance the sender's read marker merely because they sent a message. +func (a *App) publishConversationRead(conversationID, userID, readSeq, unread int64) { + a.hub.broadcast([]int64{userID}, map[string]any{"command": "READ_ACK", "data": map[string]any{ + "conversationId": conversationID, "readSeq": readSeq, "unread": unread, + }}) +} + +func (a *App) conversationReadState(ctx context.Context, conversationID, userID int64) (int64, int64, error) { + var savedReadSeq, lastSeq, unread int64 + err := a.db.QueryRowContext(ctx, `SELECT m.read_seq,c.last_seq, + (SELECT COUNT(*) FROM im_messages unread_msg WHERE unread_msg.conversation_id=m.conversation_id + AND unread_msg.seq>GREATEST(m.read_seq,m.clear_seq,m.join_seq) AND unread_msg.sender_id<>m.user_id + AND unread_msg.recalled_at IS NULL AND unread_msg.admin_removed_at IS NULL) + FROM im_conversation_members m JOIN im_conversations c ON c.id=m.conversation_id + WHERE m.conversation_id=? AND m.user_id=? AND m.status=1`, conversationID, userID).Scan(&savedReadSeq, &lastSeq, &unread) + // Older clients compare READ_ACK with the conversation sequence, which also + // includes the user's own messages. Cover those sequences only when the + // authoritative unread count is zero; a newer incoming message stays safe. + if err == nil && unread == 0 && lastSeq > savedReadSeq { + savedReadSeq = lastSeq + } + return savedReadSeq, unread, err +} + +func (a *App) markConversationRead(ctx context.Context, conversationID, userID, readSeq int64) (int64, int64, error) { + if readSeq < 0 { + return 0, 0, fmt.Errorf("已读序号无效") + } + _, err := a.db.ExecContext(ctx, `UPDATE im_conversation_members m JOIN im_conversations c ON c.id=m.conversation_id + SET m.read_seq=GREATEST(m.read_seq,LEAST(?,c.last_seq)) + WHERE m.conversation_id=? AND m.user_id=? AND m.status=1`, readSeq, conversationID, userID) + if err != nil { + return 0, 0, err + } + savedReadSeq, unread, err := a.conversationReadState(ctx, conversationID, userID) + if err != nil { + return 0, 0, err + } + a.publishConversationRead(conversationID, userID, savedReadSeq, unread) + return savedReadSeq, unread, nil } func (a *App) websocket(w http.ResponseWriter, r *http.Request) { @@ -342,7 +768,19 @@ func (a *App) websocket(w http.ResponseWriter, r *http.Request) { fail(w, http.StatusForbidden, 10006, "WebSocket 来源不允许") return } - raw := r.URL.Query().Get("token") + selectedProtocol := "" + raw := "" + for _, protocol := range websocket.Subprotocols(r) { + if strings.HasPrefix(protocol, "xingyu.jwt.") { + selectedProtocol = protocol + raw = strings.TrimPrefix(protocol, "xingyu.jwt.") + break + } + } + // Query-token compatibility is development-only because URLs may be written to proxy logs. + if raw == "" && a.config.Environment != "production" { + raw = r.URL.Query().Get("token") + } who, err := a.parseToken(raw) if err != nil || who.Role != "user" { fail(w, 401, 10001, "invalid token") @@ -366,14 +804,27 @@ func (a *App) websocket(w http.ResponseWriter, r *http.Request) { upgrader := websocket.Upgrader{CheckOrigin: func(request *http.Request) bool { return a.originAllowed(request.Header.Get("Origin")) }} + if selectedProtocol != "" { + upgrader.Subprotocols = []string{selectedProtocol} + } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } client := &wsClient{userID: who.ID, conn: conn} - a.hub.add(client) - defer func() { a.hub.remove(client); _ = conn.Close() }() - _ = conn.WriteJSON(map[string]any{"command": "AUTH_ACK", "data": map[string]any{"heartbeatSeconds": 25, "serverTime": time.Now().UnixMilli()}}) + conn.SetReadLimit(64 << 10) + _ = conn.SetReadDeadline(time.Now().Add(75 * time.Second)) + a.touchUserActivity(r.Context(), who.ID) + if a.hub.add(client) { + a.publishPresence(who.ID, true) + } + defer func() { + if a.hub.remove(client) { + a.publishPresence(who.ID, false) + } + _ = conn.Close() + }() + _ = client.writeJSON(map[string]any{"command": "AUTH_ACK", "data": map[string]any{"heartbeatSeconds": 25, "serverTime": time.Now().UnixMilli()}}) for { var frame struct { Command string `json:"command"` @@ -388,24 +839,33 @@ func (a *App) websocket(w http.ResponseWriter, r *http.Request) { } switch frame.Command { case "PING": - _ = conn.WriteJSON(map[string]any{"command": "PONG", "timestamp": time.Now().UnixMilli()}) + a.touchUserActivity(r.Context(), who.ID) + _ = conn.SetReadDeadline(time.Now().Add(75 * time.Second)) + _ = client.writeJSON(map[string]any{"command": "PONG", "timestamp": time.Now().UnixMilli()}) case "SEND_MESSAGE": item, members, persistErr := a.persistMessage(r, frame.ConversationID, who.ID, frame.ClientMsgID, frame.Type, frame.Content) if persistErr != nil { - _ = conn.WriteJSON(map[string]any{"command": "ERROR", "message": persistErr.Error()}) + _ = client.writeJSON(map[string]any{"command": "ERROR", "message": persistErr.Error()}) continue } a.hub.broadcast(members, map[string]any{"command": "MESSAGE_PUSH", "data": item}) case "READ": - _, _ = a.db.ExecContext(r.Context(), `UPDATE im_conversation_members SET read_seq=GREATEST(read_seq,?) WHERE conversation_id=? AND user_id=?`, frame.ReadSeq, frame.ConversationID, who.ID) - a.hub.broadcast([]int64{who.ID}, map[string]any{"command": "READ_ACK", "data": frame}) + if frame.ReadSeq < 0 { + _ = client.writeJSON(map[string]any{"command": "ERROR", "message": "已读序号无效"}) + continue + } + var lastSeq int64 + if a.db.QueryRowContext(r.Context(), `SELECT c.last_seq FROM im_conversations c JOIN im_conversation_members m ON m.conversation_id=c.id WHERE c.id=? AND m.user_id=? AND m.status=1`, frame.ConversationID, who.ID).Scan(&lastSeq) != nil { + _ = client.writeJSON(map[string]any{"command": "ERROR", "message": "不是会话成员"}) + continue + } + if frame.ReadSeq > lastSeq { + frame.ReadSeq = lastSeq + } + if _, _, err := a.markConversationRead(r.Context(), frame.ConversationID, who.ID, frame.ReadSeq); err != nil { + _ = client.writeJSON(map[string]any{"command": "ERROR", "message": "更新已读状态失败"}) + continue + } } } } - -func max64(a, b int64) int64 { - if a > b { - return a - } - return b -} diff --git a/im/backend/internal/app/integration.go b/im/backend/internal/app/integration.go index e27a59e..00646de 100644 --- a/im/backend/internal/app/integration.go +++ b/im/backend/internal/app/integration.go @@ -44,11 +44,20 @@ type integrationFieldView struct { var integrationSpecs = map[string][]integrationFieldSpec{ "oauth": { - {Key: "oauth.admin.frontend_callback_url", Label: "管理端登录结果页", Input: "text", Required: true, Description: "第三方授权完成后返回的管理端页面;生产环境必须使用 HTTPS"}, - {Key: "oauth.user.frontend_callback_url", Label: "客户端 H5 登录结果页", Input: "text", Required: true, Description: "第三方授权完成后返回的 uni-app H5 页面;生产环境必须使用 HTTPS"}, + {Key: "oauth.admin.frontend_callback_url", Label: "管理端登录结果页", Input: "text", Description: "管理端启用第三方登录时必填;生产环境必须使用 HTTPS"}, + {Key: "oauth.user.frontend_callback_url", Label: "客户端 H5 登录结果页", Input: "text", Description: "H5 启用第三方登录时必填;生产环境必须使用 HTTPS"}, + {Key: "oauth.app.frontend_callback_url", Label: "App GitHub 登录结果地址", Input: "text", Providers: []string{"github"}, Description: "固定为 xingyuim://oauth/callback;Android/iOS 打包须注册 xingyuim URL Scheme,GitHub 平台仍填写后端 HTTPS 回调地址"}, + {Key: "oauth.app.wechat.enabled", Label: "App 启用微信登录", Input: "boolean", Required: true, Providers: []string{"wechat"}, Description: "后端控制 App 原生微信登录;还需在 manifest 中配置微信 SDK 并重新打包"}, + {Key: "oauth.app.qq.enabled", Label: "App 启用 QQ 登录", Input: "boolean", Required: true, Providers: []string{"qq"}, Description: "后端控制 App 原生 QQ 登录;还需在 manifest 中配置 QQ SDK 并重新打包"}, + {Key: "oauth.app.github.enabled", Label: "App 启用 GitHub 登录", Input: "boolean", Required: true, Providers: []string{"github"}, Description: "系统浏览器授权后返回 App;复用下方 GitHub OAuth 应用参数"}, + {Key: "oauth.app.google.enabled", Label: "App 启用 Google 登录", Input: "boolean", Required: true, Providers: []string{"google"}, Description: "后端控制 App 原生 Google 登录;需配置 Google SDK 和允许的客户端 ID"}, + {Key: "oauth.app.wechat.client_id", Label: "App 微信 AppID", Input: "text", Providers: []string{"wechat"}, Description: "微信开放平台移动应用 AppID,需与 App 打包配置一致;独立于网站应用"}, + {Key: "oauth.app.wechat.client_secret", Label: "App 微信 AppSecret", Input: "secret", Providers: []string{"wechat"}, Description: "移动应用密钥,AES-GCM 加密保存;不得写入 App 包"}, + {Key: "oauth.app.qq.client_id", Label: "App QQ AppID", Input: "text", Providers: []string{"qq"}, Description: "QQ 互联移动应用 AppID,需与 App 打包配置一致"}, + {Key: "oauth.app.google.client_ids", Label: "App Google Client ID 白名单", Input: "text", Providers: []string{"google"}, Description: "Android/iOS OAuth 客户端 ID,以英文逗号分隔;后端验证令牌所属应用"}, {Key: "oauth.wechat.enabled", Label: "管理端启用微信登录", Input: "boolean", Required: true, Providers: []string{"wechat"}, Description: "开启后且必填参数完整时,管理端登录页显示微信入口"}, - {Key: "oauth.user.wechat.enabled", Label: "客户端启用微信登录", Input: "boolean", Required: true, Providers: []string{"wechat"}, Description: "开启后且必填参数完整时,uni-app H5 登录页显示微信入口"}, + {Key: "oauth.user.wechat.enabled", Label: "H5 启用微信登录", Input: "boolean", Required: true, Providers: []string{"wechat"}, Description: "开启后且必填参数完整时,uni-app H5 登录页显示微信入口"}, {Key: "oauth.wechat.client_id", Label: "微信 AppID", Input: "text", Providers: []string{"wechat"}, Description: "微信开放平台网站应用 AppID"}, {Key: "oauth.wechat.client_secret", Label: "微信 AppSecret", Input: "secret", Providers: []string{"wechat"}, Description: "AES-GCM 加密保存"}, {Key: "oauth.wechat.authorization_url", Label: "微信授权地址", Input: "text", Providers: []string{"wechat"}, Description: "默认使用微信开放平台 qrconnect 地址"}, @@ -58,7 +67,7 @@ var integrationSpecs = map[string][]integrationFieldSpec{ {Key: "oauth.wechat.redirect_uri", Label: "微信回调地址", Input: "text", Providers: []string{"wechat"}, Description: "必须与微信开放平台登记值完全一致"}, {Key: "oauth.qq.enabled", Label: "管理端启用 QQ 登录", Input: "boolean", Required: true, Providers: []string{"qq"}, Description: "开启后且必填参数完整时,管理端登录页显示 QQ 入口"}, - {Key: "oauth.user.qq.enabled", Label: "客户端启用 QQ 登录", Input: "boolean", Required: true, Providers: []string{"qq"}, Description: "开启后且必填参数完整时,uni-app H5 登录页显示 QQ 入口"}, + {Key: "oauth.user.qq.enabled", Label: "H5 启用 QQ 登录", Input: "boolean", Required: true, Providers: []string{"qq"}, Description: "开启后且必填参数完整时,uni-app H5 登录页显示 QQ 入口"}, {Key: "oauth.qq.client_id", Label: "QQ AppID", Input: "text", Providers: []string{"qq"}, Description: "QQ 互联应用 AppID"}, {Key: "oauth.qq.client_secret", Label: "QQ AppKey", Input: "secret", Providers: []string{"qq"}, Description: "AES-GCM 加密保存"}, {Key: "oauth.qq.authorization_url", Label: "QQ 授权地址", Input: "text", Providers: []string{"qq"}, Description: "QQ OAuth 2.0 authorize 地址"}, @@ -69,7 +78,7 @@ var integrationSpecs = map[string][]integrationFieldSpec{ {Key: "oauth.qq.redirect_uri", Label: "QQ 回调地址", Input: "text", Providers: []string{"qq"}, Description: "必须与 QQ 互联登记值完全一致"}, {Key: "oauth.github.enabled", Label: "管理端启用 GitHub 登录", Input: "boolean", Required: true, Providers: []string{"github"}, Description: "开启后且必填参数完整时,管理端登录页显示 GitHub 入口"}, - {Key: "oauth.user.github.enabled", Label: "客户端启用 GitHub 登录", Input: "boolean", Required: true, Providers: []string{"github"}, Description: "开启后且必填参数完整时,uni-app H5 登录页显示 GitHub 入口"}, + {Key: "oauth.user.github.enabled", Label: "H5 启用 GitHub 登录", Input: "boolean", Required: true, Providers: []string{"github"}, Description: "开启后且必填参数完整时,uni-app H5 登录页显示 GitHub 入口"}, {Key: "oauth.github.client_id", Label: "GitHub Client ID", Input: "text", Providers: []string{"github"}, Description: "GitHub OAuth App Client ID"}, {Key: "oauth.github.client_secret", Label: "GitHub Client Secret", Input: "secret", Providers: []string{"github"}, Description: "AES-GCM 加密保存"}, {Key: "oauth.github.authorization_url", Label: "GitHub 授权地址", Input: "text", Providers: []string{"github"}, Description: "GitHub OAuth authorize 地址"}, @@ -79,7 +88,7 @@ var integrationSpecs = map[string][]integrationFieldSpec{ {Key: "oauth.github.redirect_uri", Label: "GitHub 回调地址", Input: "text", Providers: []string{"github"}, Description: "必须与 OAuth App 的 callback URL 完全一致"}, {Key: "oauth.google.enabled", Label: "管理端启用 Google 登录", Input: "boolean", Required: true, Providers: []string{"google"}, Description: "开启后且必填参数完整时,管理端登录页显示 Google 入口"}, - {Key: "oauth.user.google.enabled", Label: "客户端启用 Google 登录", Input: "boolean", Required: true, Providers: []string{"google"}, Description: "开启后且必填参数完整时,uni-app H5 登录页显示 Google 入口"}, + {Key: "oauth.user.google.enabled", Label: "H5 启用 Google 登录", Input: "boolean", Required: true, Providers: []string{"google"}, Description: "开启后且必填参数完整时,uni-app H5 登录页显示 Google 入口"}, {Key: "oauth.google.client_id", Label: "Google Client ID", Input: "text", Providers: []string{"google"}, Description: "Google OAuth 2.0 Client ID"}, {Key: "oauth.google.client_secret", Label: "Google Client Secret", Input: "secret", Providers: []string{"google"}, Description: "AES-GCM 加密保存"}, {Key: "oauth.google.authorization_url", Label: "Google 授权地址", Input: "text", Providers: []string{"google"}, Description: "Google OAuth authorization endpoint"}, @@ -169,6 +178,19 @@ var integrationSpecs = map[string][]integrationFieldSpec{ {Key: "sms.webhook_url", Label: "Webhook 地址", Input: "text", Required: true, Providers: []string{"webhook"}, Description: "接收 JSON POST;生产环境必须使用 HTTPS"}, {Key: "sms.webhook_token", Label: "Webhook Token", Input: "secret", Providers: []string{"webhook"}, Description: "以 Bearer Token 发送,仅显示保存状态"}, }, + "ai": { + {Key: "ai.enabled", Label: "启用 AI 托管", Input: "boolean", Required: true, Description: "总开关;关闭后不调用任何模型,测试账号不再自动回复"}, + {Key: "ai.default_model_id", Label: "默认模型 ID", Input: "text", Description: "留空则使用模型列表中标记为默认的一条"}, + {Key: "ai.max_concurrency", Label: "并发调用上限", Input: "number", Description: "同时进行的模型调用数,建议 2 到 16"}, + {Key: "ai.daily_call_limit", Label: "每日调用上限", Input: "number", Description: "全局每日调用次数,0 表示不限制"}, + {Key: "ai.agent_daily_reply_limit", Label: "单账号每日回复上限", Input: "number", Description: "单个托管账号每日回复条数,0 表示不限制"}, + {Key: "ai.reply_delay_min_ms", Label: "回复最小延迟(毫秒)", Input: "number", Description: "避免秒回失真,建议 1000 以上"}, + {Key: "ai.reply_delay_max_ms", Label: "回复最大延迟(毫秒)", Input: "number", Description: "必须不小于最小延迟"}, + {Key: "ai.allow_batches", Label: "允许托管的测试批次", Input: "text", Description: "逗号分隔;留空表示全部测试账号"}, + {Key: "ai.disclose_in_chat", Label: "会话内标注 AI 身份", Input: "boolean", Required: true, Description: "关闭前需确认所在地区对机器人身份披露的监管要求"}, + {Key: "ai.fallback_text", Label: "失败兜底文案", Input: "text", Description: "模型调用失败时发送的内容,留空则静默不回复"}, + {Key: "ai.log_retention_days", Label: "调用日志保留天数", Input: "number", Description: "建议 7 到 180 天"}, + }, "payment": { {Key: "payment.mode", Label: "支付模式", Input: "select", Required: true, Options: []string{"sandbox", "live"}, Description: "sandbox 可直接完成本地支付闭环"}, {Key: "payment.gateway.create_url", Label: "支付网关下单地址", Input: "text", Description: "live 模式必填,生产环境必须使用 HTTPS"}, @@ -486,6 +508,33 @@ func (a *App) adminTestIntegration(w http.ResponseWriter, r *http.Request) { } } reply(w, map[string]any{"success": true, "message": "支付配置校验通过", "mode": mode, "channels": channels}) + case "ai": + if err := a.validateAIConfig(r.Context()); err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + model, err := a.defaultAIModel(r.Context()) + if err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + provider, err := newAIProvider(model, a.config.Environment == "production") + if err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + started := time.Now() + result, chatErr := provider.Chat(r.Context(), aiChatRequest{ + System: "你在参与一次接口连通性测试,请直接给出简短的中文回复。", + Messages: []aiChatMessage{{Role: "user", Text: aiProbePrompt}}, + }) + latency := int(time.Since(started).Milliseconds()) + a.logAICall(r.Context(), model.ID, "probe", 0, 0, latency, result, chatErr) + if chatErr != nil { + fail(w, http.StatusBadGateway, 50002, chatErr.Error()) + return + } + reply(w, map[string]any{"success": true, "message": model.Name + " 调用成功", "model": model.Name, "protocol": model.Protocol, "latencyMs": latency, "reply": result.Text}) case "storage": provider := a.configPlain(r.Context(), "storage.provider", "local") if err := a.validateStorageProviderConfig(r.Context()); err != nil { @@ -504,11 +553,24 @@ func (a *App) adminTestIntegration(w http.ResponseWriter, r *http.Request) { fail(w, http.StatusBadRequest, 20001, err.Error()) return } - if _, err = a.userOAuthFrontendURL(r.Context()); err != nil { + appProviders, err := a.enabledAppOAuthProviders(r.Context()) + if err != nil { fail(w, http.StatusBadRequest, 20001, err.Error()) return } - if len(adminProviders) == 0 && len(userProviders) == 0 { + if len(userProviders) > 0 { + if _, err = a.userOAuthFrontendURL(r.Context()); err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + } + if len(adminProviders) > 0 { + if _, err = a.adminOAuthFrontendURL(r.Context()); err != nil { + fail(w, http.StatusBadRequest, 20001, err.Error()) + return + } + } + if len(adminProviders) == 0 && len(userProviders) == 0 && len(appProviders) == 0 { reply(w, map[string]any{"success": true, "message": "当前未启用第三方登录,登录页不会显示第三方入口", "providers": []any{}}) return } @@ -519,6 +581,9 @@ func (a *App) adminTestIntegration(w http.ResponseWriter, r *http.Request) { for _, provider := range userProviders { items = append(items, map[string]string{"audience": "user", "code": provider.Code, "name": provider.Name}) } + for _, provider := range appProviders { + items = append(items, map[string]string{"audience": "app", "code": provider.Code, "name": provider.Name}) + } reply(w, map[string]any{"success": true, "message": "第三方登录配置校验通过", "providers": items}) default: fail(w, http.StatusNotFound, 30001, "配置分组不存在") diff --git a/im/backend/migrations/030_ai_models.sql b/im/backend/migrations/030_ai_models.sql new file mode 100644 index 0000000..298cfa9 --- /dev/null +++ b/im/backend/migrations/030_ai_models.sql @@ -0,0 +1,58 @@ +SET NAMES utf8mb4; + +-- Model records rather than a single provider block: several protocols have to +-- coexist so operations can switch models without a release. +CREATE TABLE IF NOT EXISTS ai_models ( + id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT, + name VARCHAR(60) NOT NULL, + protocol VARCHAR(20) NOT NULL DEFAULT 'openai', + base_url VARCHAR(300) NOT NULL DEFAULT '', + model_name VARCHAR(120) NOT NULL DEFAULT '', + api_key_cipher TEXT NOT NULL, + temperature DECIMAL(3,2) NOT NULL DEFAULT 0.80, + max_tokens INT UNSIGNED NOT NULL DEFAULT 256, + timeout_ms INT UNSIGNED NOT NULL DEFAULT 15000, + status TINYINT UNSIGNED NOT NULL DEFAULT 1, + is_default TINYINT(1) NOT NULL DEFAULT 0, + extra_json JSON NULL, + created_at DATETIME(3) NOT NULL DEFAULT CURRENT_TIMESTAMP(3), + updated_at DATETIME(3) NOT NULL DEFAULT CURRENT_TIMESTAMP(3) ON UPDATE CURRENT_TIMESTAMP(3), + PRIMARY KEY (id), + UNIQUE KEY uk_ai_models_name (name), + KEY idx_ai_models_status (status, is_default) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci; + +-- Every call is recorded for cost accounting and after-the-fact review. The +-- prompt is kept as a digest so conversation content does not accumulate here. +CREATE TABLE IF NOT EXISTS ai_call_logs ( + id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT, + model_id BIGINT UNSIGNED NOT NULL, + scene VARCHAR(30) NOT NULL DEFAULT 'chat', + agent_user_id BIGINT UNSIGNED NOT NULL DEFAULT 0, + conversation_id BIGINT UNSIGNED NOT NULL DEFAULT 0, + latency_ms INT UNSIGNED NOT NULL DEFAULT 0, + prompt_tokens INT UNSIGNED NOT NULL DEFAULT 0, + completion_tokens INT UNSIGNED NOT NULL DEFAULT 0, + status TINYINT UNSIGNED NOT NULL DEFAULT 1, + error VARCHAR(500) NOT NULL DEFAULT '', + prompt_digest VARCHAR(64) NOT NULL DEFAULT '', + reply_text VARCHAR(1000) NOT NULL DEFAULT '', + created_at DATETIME(3) NOT NULL DEFAULT CURRENT_TIMESTAMP(3), + PRIMARY KEY (id), + KEY idx_ai_call_logs_model_created (model_id, created_at), + KEY idx_ai_call_logs_created (created_at) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci; + +INSERT INTO system_configs (config_key, config_value, value_type, description) VALUES + ('ai.enabled', 'false', 'boolean', '总开关:关闭后不调用任何模型,测试账号不再自动回复'), + ('ai.default_model_id', '', 'string', '默认模型 ID,留空则使用模型列表中标记为默认的一条'), + ('ai.max_concurrency', '4', 'string', '同时进行的模型调用上限'), + ('ai.daily_call_limit', '2000', 'string', '全局每日调用次数上限,0 表示不限制'), + ('ai.agent_daily_reply_limit', '50', 'string', '单个托管账号每日回复条数上限,0 表示不限制'), + ('ai.reply_delay_min_ms', '2000', 'string', '回复前的最小延迟毫秒数,避免秒回失真'), + ('ai.reply_delay_max_ms', '6000', 'string', '回复前的最大延迟毫秒数'), + ('ai.allow_batches', '', 'string', '允许托管的测试批次,逗号分隔;留空表示全部测试账号'), + ('ai.disclose_in_chat', 'true', 'boolean', '会话内标注 AI 身份;关闭前需确认所在地区的监管要求'), + ('ai.fallback_text', '', 'string', '模型调用失败时发送的兜底文案,留空则静默不回复'), + ('ai.log_retention_days', '30', 'string', '调用日志保留天数') +ON DUPLICATE KEY UPDATE description = VALUES(description), value_type = VALUES(value_type); diff --git a/im/backend/migrations/031_ai_agents.sql b/im/backend/migrations/031_ai_agents.sql new file mode 100644 index 0000000..ad86923 --- /dev/null +++ b/im/backend/migrations/031_ai_agents.sql @@ -0,0 +1,44 @@ +SET NAMES utf8mb4; + +-- One row per managed test account. model_id 0 falls back to the configured +-- default model so operations can switch every agent at once. +CREATE TABLE IF NOT EXISTS ai_agents ( + user_id BIGINT UNSIGNED NOT NULL, + model_id BIGINT UNSIGNED NOT NULL DEFAULT 0, + persona VARCHAR(1000) NOT NULL DEFAULT '', + enabled TINYINT(1) NOT NULL DEFAULT 1, + daily_reply_limit INT UNSIGNED NOT NULL DEFAULT 0, + reply_count INT UNSIGNED NOT NULL DEFAULT 0, + last_reply_at DATETIME(3) NULL, + created_at DATETIME(3) NOT NULL DEFAULT CURRENT_TIMESTAMP(3), + updated_at DATETIME(3) NOT NULL DEFAULT CURRENT_TIMESTAMP(3) ON UPDATE CURRENT_TIMESTAMP(3), + PRIMARY KEY (user_id), + KEY idx_ai_agents_enabled (enabled) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci; + +-- A durable queue rather than an in-memory one: replies survive a restart and +-- several instances can share it. Claiming is an UPDATE that stamps claim_token +-- instead of SELECT ... FOR UPDATE SKIP LOCKED, because deployments still run +-- MySQL 5.7 where SKIP LOCKED does not parse. +-- pending_key is NULL once a job leaves the pending state, so the unique index +-- allows exactly one waiting job per conversation and collapses a burst of +-- incoming messages into a single reply. +CREATE TABLE IF NOT EXISTS ai_reply_jobs ( + id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT, + conversation_id BIGINT UNSIGNED NOT NULL, + agent_user_id BIGINT UNSIGNED NOT NULL, + trigger_message_id BIGINT UNSIGNED NOT NULL DEFAULT 0, + status TINYINT UNSIGNED NOT NULL DEFAULT 0, + attempts TINYINT UNSIGNED NOT NULL DEFAULT 0, + run_after DATETIME(3) NOT NULL DEFAULT CURRENT_TIMESTAMP(3), + claim_token CHAR(32) NOT NULL DEFAULT '', + last_error VARCHAR(500) NOT NULL DEFAULT '', + created_at DATETIME(3) NOT NULL DEFAULT CURRENT_TIMESTAMP(3), + updated_at DATETIME(3) NOT NULL DEFAULT CURRENT_TIMESTAMP(3) ON UPDATE CURRENT_TIMESTAMP(3), + pending_key BIGINT UNSIGNED GENERATED ALWAYS AS (IF(status = 0, conversation_id, NULL)) STORED, + PRIMARY KEY (id), + UNIQUE KEY uk_ai_reply_jobs_pending (pending_key), + KEY idx_ai_reply_jobs_due (status, run_after), + KEY idx_ai_reply_jobs_claim (claim_token), + KEY idx_ai_reply_jobs_agent (agent_user_id, created_at) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;