5 Commits
Author SHA1 Message Date
wangdefa 9309ad1ffc 修复 AI 网关流式断流:错误透传、自动降级,max_tokens 可缺省
CI / test (push) Successful in 33s
Release / release (push) Successful in 56s
2026-07-13 15:06:01 +08:00
wangdefa 81d4650f3d 代理批量关联租户接口、关于页运行时与资源指标
CI / test (push) Successful in 31s
Release / release (push) Successful in 57s
2026-07-13 12:23:18 +08:00
wangdefa c7cc5616ed 发布 0.3.0:恢复 Chat 接口并增强云端事件通知
CI / test (push) Successful in 31s
Release / release (push) Successful in 48s
2026-07-13 10:13:44 +08:00
wangdefa 489cb49cb3 AI 网关切换 OpenAI 兼容面并移除 chat 端点,新增模型黑白名单
CI / test (push) Successful in 31s
Release / release (push) Successful in 54s
2026-07-12 17:48:28 +08:00
wangdefa 7706f59549 发布 0.2.0:模型池自愈、探测修正、任务异步触发与删除加固
CI / test (push) Successful in 32s
Release / release (push) Successful in 52s
2026-07-10 20:25:37 +08:00
67 changed files with 5422 additions and 3975 deletions
+2 -2
View File
@@ -23,7 +23,7 @@ jobs:
echo "IMAGE=${REGISTRY}/${{ gitea.repository }}" >> "$GITHUB_ENV" echo "IMAGE=${REGISTRY}/${{ gitea.repository }}" >> "$GITHUB_ENV"
echo "BUILD_TIME=$(date -u +%Y-%m-%dT%H:%M:%SZ)" >> "$GITHUB_ENV" echo "BUILD_TIME=$(date -u +%Y-%m-%dT%H:%M:%SZ)" >> "$GITHUB_ENV"
- name: 下载 DASH_VERSION 固定版本的前端 dist.zip,校验后解压进嵌入目录 - name: 下载前端 dist.zip
env: env:
TOKEN: ${{ secrets.BUILD_TOKEN }} TOKEN: ${{ secrets.BUILD_TOKEN }}
run: | run: |
@@ -36,7 +36,7 @@ jobs:
unzip -q dist.zip -d internal/webui/dist unzip -q dist.zip -d internal/webui/dist
test -f internal/webui/dist/index.html test -f internal/webui/dist/index.html
- name: 构建双架构二进制(artifact stage 导出) - name: 构建双架构二进制
run: | run: |
docker buildx build --pull \ docker buildx build --pull \
--platform linux/amd64,linux/arm64 \ --platform linux/amd64,linux/arm64 \
+1 -1
View File
@@ -1 +1 @@
0.6.5 0.6.6
+13 -13
View File
@@ -183,7 +183,7 @@ Complex task: ask the user if you can create a Trellis task and enter the planni
- 1.0 Create task `[required · once]` (only after task-creation consent) - 1.0 Create task `[required · once]` (only after task-creation consent)
- 1.1 Requirement exploration `[required · repeatable]` (`prd.md`; complex tasks also need `design.md` + `implement.md`) - 1.1 Requirement exploration `[required · repeatable]` (`prd.md`; complex tasks also need `design.md` + `implement.md`)
- 1.2 Research `[optional · repeatable]` - 1.2 Research `[optional · repeatable]`
- 1.3 Configure context `[required · once]` — Claude Code, Cursor, OpenCode, Codex, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix (sub-agent-dispatch platforms only; inline platforms skip) - 1.3 Configure context `[required · once]` — Claude Code, Cursor, OpenCode, Codex, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix (sub-agent-dispatch platforms only; inline platforms skip)
- 1.4 Activate task `[required · once]` (review gate, then `task.py start`; status → in_progress) - 1.4 Activate task `[required · once]` (review gate, then `task.py start`; status → in_progress)
- 1.5 Completion criteria - 1.5 Completion criteria
@@ -272,13 +272,13 @@ Code committed. Run `/trellis:finish-work`; if dirty, return to Phase 3.4 first.
When a user request matches one of these intents inside an active task, route first, then load the detailed phase step if needed. When a user request matches one of these intents inside an active task, route first, then load the detailed phase step if needed.
[Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
- Planning or unclear requirements -> `trellis-brainstorm`. - Planning or unclear requirements -> `trellis-brainstorm`.
- `in_progress` implementation/check -> dispatch `trellis-implement` / `trellis-check`. - `in_progress` implementation/check -> dispatch `trellis-implement` / `trellis-check`.
- Repeated debugging -> `trellis-break-loop`; spec updates -> `trellis-update-spec`. - Repeated debugging -> `trellis-break-loop`; spec updates -> `trellis-update-spec`.
[/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
[codex-inline, Kilo, Antigravity, Devin] [codex-inline, Kilo, Antigravity, Devin]
@@ -353,7 +353,7 @@ Return to this step whenever requirements change and revise the relevant artifac
Research can happen at any time during requirement exploration. It isn't limited to local code — you can use any available tool (MCP servers, skills, web search, etc.) to look up external information, including third-party library docs, industry practices, API references, etc. Research can happen at any time during requirement exploration. It isn't limited to local code — you can use any available tool (MCP servers, skills, web search, etc.) to look up external information, including third-party library docs, industry practices, API references, etc.
[Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
Spawn the research sub-agent: Spawn the research sub-agent:
@@ -361,7 +361,7 @@ Spawn the research sub-agent:
- **Task description**: Research <specific question> - **Task description**: Research <specific question>
- **Key requirement**: Research output MUST be persisted to `{TASK_DIR}/research/` - **Key requirement**: Research output MUST be persisted to `{TASK_DIR}/research/`
[/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
[codex-inline, Kilo, Antigravity, Devin] [codex-inline, Kilo, Antigravity, Devin]
@@ -380,7 +380,7 @@ Brainstorm and research can interleave freely — pause to research a technical
#### 1.3 Configure context `[required · once]` #### 1.3 Configure context `[required · once]`
[Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
Curate `implement.jsonl` and `check.jsonl` so the Phase 2 sub-agents get the right spec/research context. These files were seeded on `task create` with a single self-describing `_example` line; your job here is to fill in real entries. Curate `implement.jsonl` and `check.jsonl` so the Phase 2 sub-agents get the right spec/research context. These files were seeded on `task create` with a single self-describing `_example` line; your job here is to fill in real entries.
@@ -425,7 +425,7 @@ Ready gate: both `implement.jsonl` and `check.jsonl` must contain at least one r
Skip this step only when both files already have real curated entries. Skip this step only when both files already have real curated entries.
[/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
[codex-inline, Kilo, Antigravity, Devin] [codex-inline, Kilo, Antigravity, Devin]
@@ -458,11 +458,11 @@ If `task.py start` errors with a session-identity message (no context key from h
| `design.md` exists (complex tasks) | ✅ | | `design.md` exists (complex tasks) | ✅ |
| `implement.md` exists (complex tasks) | ✅ | | `implement.md` exists (complex tasks) | ✅ |
[Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
| `implement.jsonl` and `check.jsonl` each contain at least one real curated entry (seed row does not count) | ✅ | | `implement.jsonl` and `check.jsonl` each contain at least one real curated entry (seed row does not count) | ✅ |
[/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
--- ---
@@ -472,7 +472,7 @@ Goal: turn reviewed planning artifacts into code that passes quality checks.
#### 2.1 Implement `[required · repeatable]` #### 2.1 Implement `[required · repeatable]`
[Claude Code, Cursor, OpenCode, CodeBuddy, Droid, Pi] [Claude Code, Cursor, OpenCode, CodeBuddy, Droid, Pi, Oh My Pi]
Spawn the implement sub-agent: Spawn the implement sub-agent:
@@ -484,7 +484,7 @@ The platform hook/plugin auto-handles:
- Reads `implement.jsonl` and injects referenced spec/research files into the agent prompt - Reads `implement.jsonl` and injects referenced spec/research files into the agent prompt
- Injects `prd.md`, `design.md` if present, and `implement.md` if present - Injects `prd.md`, `design.md` if present, and `implement.md` if present
[/Claude Code, Cursor, OpenCode, CodeBuddy, Droid, Pi] [/Claude Code, Cursor, OpenCode, CodeBuddy, Droid, Pi, Oh My Pi]
[codex-sub-agent, Gemini, Qoder, Copilot, ZCode, Reasonix, Trae] [codex-sub-agent, Gemini, Qoder, Copilot, ZCode, Reasonix, Trae]
@@ -526,7 +526,7 @@ The platform prelude auto-handles the context load requirement:
#### 2.2 Quality check `[required · repeatable]` #### 2.2 Quality check `[required · repeatable]`
[Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
Spawn the check sub-agent: Spawn the check sub-agent:
@@ -540,7 +540,7 @@ The check agent's job:
- Auto-fix issues it finds - Auto-fix issues it finds
- Run lint and typecheck to verify - Run lint and typecheck to verify
[/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, ZCode, Reasonix, Trae] [/Claude Code, Cursor, OpenCode, codex-sub-agent, Kiro, Gemini, Qoder, CodeBuddy, Copilot, Droid, Pi, Oh My Pi, ZCode, Reasonix, Trae]
[codex-inline, Kilo, Antigravity, Devin] [codex-inline, Kilo, Antigravity, Devin]
+58
View File
@@ -2,6 +2,64 @@
格式遵循 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/),版本号遵循语义化版本。 格式遵循 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/),版本号遵循语义化版本。
## [0.3.1] - 2026-07-13
### Added
- AI 网关流式断流自动降级:Messages 与 Chat Completions 流式请求在客户端尚未收到任何输出时遭遇上游断流,自动降级为非流式重做,结果按标准事件 / chunk 序列一次推送,调用日志记 `retries=1` 与降级标记。实测 OCI 兼容面对 `instructions``tools` 合计超约 64.5KB 的流式请求会静默断连(无错误事件,非流式正常,消息正文不计入),Claude Code 等大体量系统提示客户端极易触发;Responses 直通因初始事件已转发无法透明降级,日志记「上游流提前终止」;README 增补「已知上游限制:大 system 区流式断流」小节
### Changed
- Messages 的 `max_tokens` 改为可缺省:缺省或 ≤0 时按默认值 8192 放行(此前返回 400;部分客户端将该字段视为选填)
### Fixed
- AI 网关流式假成功:上游 error / `response.failed` 事件此前被吞掉,流提前 EOF 也被伪装成正常结束,调用日志呈现 200 · 0/0 且无错误信息;现 Messages 将上游失败转为 Anthropic `error` 事件,三个流式端点日志均记录上游错误消息或「上游流提前终止,未返回终态事件」
## [0.3.0] - 2026-07-13
### Added
- 恢复 OpenAI Chat Completions 兼容端点 `/ai/v1/chat/completions`:请求经 OCI Responses 接口转换转发,支持非流式与 SSE、文本与图片、function 工具调用、结构化输出、推理力度、模型白名单、缓存 token 用量和调用日志
- 代理侧批量关联租户接口 `PUT /api/v1/proxies/{id}/tenants`:一次调用设置关联该代理的租户全集——列表内租户改挂到本代理(含从其他代理改挂),列表外已关联的解除,事务生效并返回最新关联计数;含无效租户 ID 返回 400,代理不存在返回 404
- `/api/v1/about` 扩展运行时信息:进程启动时刻 `startedAt`、运行秒数 `uptimeSeconds` 与资源占用 `resources`——CPU 自启动均值(占单核百分比)、CPU 核数、goroutine 数、Go 堆内存与向系统申请内存、数据库引擎与磁盘占用(SQLite 统计主文件与 -wal / -shmMySQL / PostgreSQL 外部库返回 -1 表示不可度量)
- README 重构部署、配置、鉴权、反向代理、AI 网关与升级发布说明,并新增 Responses、Chat Completions、Embeddings、Anthropic Messages 与标准接口的字段兼容矩阵,明确直通、转换、降级、忽略和拒绝边界
### Changed
- 云端关键事件通知补充操作者、来源 IP、执行结果和说明信息;支持从 OCI Audit 与 IDCS 登录事件提取成功 / 失败状态、策略描述和登录失败原因,缺失字段统一显示占位符
### Removed
- **移除自定义审计告警规则**及其阈值、窗口和冷却匹配能力;`/api/v1/log-events/alert-rules``/api/v1/log-events/alert-rules/{ruleId}` 管理接口、`audit_alert` 通知模板不再提供,已有规则升级后不再执行
## [0.2.0] - 2026-07-12
### Added
- AI 网关模型黑名单:面板集中拉黑模型,全渠道剔除(模型列表 / 路由 / 探测候选不再出现,同步不入库),移除条目并重新同步后恢复;新表 `ai_model_blacklists`(启动自动迁移),管理接口 `/api/v1/ai-blacklist`
- AI 密钥模型白名单:密钥可限定放行的模型列表(空 = 不限),白名单外调用 404,模型列表只返回交集
- 推理力度(effort)直通:Responses `reasoning.effort` 与 Anthropic Messages `output_config.effort` 原样透传上游、不做档位校验;grok-4.3 五档(`none` 完全关闭思考)、gpt-oss `minimal`-`high`、multi-agent `xhigh`(控制并行 agent 数)实测生效
- xAI Grok 服务端工具:`/ai/v1/responses``tools` 支持 `{"type":"web_search"}` / `{"type":"x_search"}`(官方 Agent Tools 格式,可与 `function` 混用,仅非流式),响应保留 `web_search_call` 输出项与引用
### Changed
- **AI 网关对话链路整体切换到 OCI OpenAI 兼容面**`/actions/v1/responses`),typed chat 面全链路移除:`/ai/v1/responses` 请求体原样直通(仅强制 `store:false`),流式为上游原生事件流,gpt-oss 推理增量直达客户端;`/ai/v1/messages` 转换为 Responses 请求直通,流式经事件桥转回标准 Anthropic 事件序列,`stop_sequences` / `top_k` 等无兼容面对应物的字段改为静默忽略,思考输出不转 `thinking` 块;渠道探测改用兼容面最小请求;流式建立后绑定渠道,中断不重试不计熔断
- 后台任务「立即执行」改异步触发:接口立即返回 202,执行结果经任务日志轮询呈现;同一任务在途时重复触发返回 409「任务正在执行中」,与 cron 重叠触发静默跳过(此前同步阻塞数十秒且同一任务可并发重复执行)
- 渠道探测候选跨厂商分散(上限 8,voice 等不可对话形态排除):单一厂商在区域内全为微调基座时不再拖垮整个渠道的探测结论
- 网关调用遇「模型不可按需调用」类错误(微调基座 400 / 实体不存在 404)自动换渠道重试且不计入渠道熔断,错误信息带模型名便于加入黑名单;仅鉴权类错误才判定租户无配额
- 同名模型多条目去重优先保留非微调基座条目,降低缓存到不可调用 OCID 的概率
### Removed
- **`/ai/v1/chat/completions` 端点移除**(404):OpenAI 官方新能力已集中到 Responses 且本面板真实流量为零,请迁移 `/ai/v1/responses``/ai/v1/messages`
- google / cohere 对话模型不再供给(OCI OpenAI 兼容面不承接,上游 400cohere embed 向量模型不受影响),模型目录按 `xai.` / `meta.` / `openai.` 厂商前缀过滤
### Fixed
- 渠道「同步模型成功但探测报错不可用」:法兰克福等区域的微调基座模型占满探测候选导致的误报(探测状态与真实可用性不符的根因)
- 租户删除:告警命中清理改子查询,日志事件数万条时不再超出 SQL 绑定变量上限导致删除失败;无法解析的任务 payload 记警告跳过并保留原任务,不再永久阻断租户删除
## [0.1.0] - 2026-07-10 ## [0.1.0] - 2026-07-10
### Added ### Added
+1 -1
View File
@@ -1 +1 @@
v0.1.0 v0.3.0
+358 -83
View File
@@ -4,16 +4,21 @@
# OCI Portal # OCI Portal
**Oracle Cloud Infrastructure 多租户管理面板** **自托管的 OCI 多租户管理面板与 GenAI 兼容网关**
![Go](https://img.shields.io/badge/Go-1.26-00ADD8?logo=go&logoColor=white) ![Release](https://img.shields.io/github/v/release/wangdefaa/oci-portal?display_name=tag)
![License](https://img.shields.io/badge/License-MIT-blue) ![Go](https://img.shields.io/badge/Go-1.26.5-00ADD8?logo=go&logoColor=white)
![Docker](https://img.shields.io/badge/Docker-amd64%20%7C%20arm64-2496ED?logo=docker&logoColor=white) ![Docker](https://img.shields.io/badge/Docker-amd64%20%7C%20arm64-2496ED?logo=docker&logoColor=white)
![License](https://img.shields.io/badge/License-MIT-blue)
本仓库为后端;前端工程见 [oci-portal-dash](https://github.com/wangdefaa/oci-portal-dash)(构建产物嵌入本服务成单文件) [快速开始](#快速开始) · [核心能力](#核心能力) · [生产部署](#生产部署) · [AI 网关](#ai-网关) · [开发](#开发)
</div> </div>
OCI Portal 将多份 OCI API Key、云资源、自动化任务、审计事件和 OCI GenAI 渠道集中到一个管理界面。Vue 前端通过 `go:embed` 嵌入 Go 服务,Release 以单个二进制和多架构容器镜像交付。
本仓库为后端与发行仓库;前端源码位于 [oci-portal-dash](https://github.com/wangdefaa/oci-portal-dash)。
## 界面预览 ## 界面预览
| 总览 | 登录 | | 总览 | 登录 |
@@ -28,59 +33,118 @@
| --- | --- | | --- | --- |
| ![AI 网关](docs/assets/screenshot-ai-gateway.png) | ![通知设置](docs/assets/screenshot-settings-notify.png) | | ![AI 网关](docs/assets/screenshot-ai-gateway.png) | ![通知设置](docs/assets/screenshot-settings-notify.png) |
## 特性 ## 核心能力
- **租户管理**多份 OCI API Key 配置集中管理,私钥/口令 AES-256-GCM 加密落库;分组、批量测活、账户类型与订阅信息识别 - **租户与区域**:集中管理多份 OCI API Key,支持分组、批量测活、账户画像、订阅区域缓存与区域切换;私钥口令使用 AES-256-GCM 加密落库
- **资源管理**:实例(含创建/电源操作/更换公网 IP/IPv6/VNIC/引导卷)、VCN/子网/安全列表、块存储、限额成本查询,多区域/多区间支持 - **计算、网络与存储**:实例创建电源操作公网 IPIPv6VNIC、串行控制台连接,VCN / 子网 / 安全列表,引导卷与块存储挂载,限额成本查询
- **抢机任务**cron 周期尝试创建实例直到成功,支持熔断与通知 - **自动化任务**抢机、租户测活、成本同步、AI 渠道探测四类 cron 任务,提供执行日志、重叠执行防护、熔断与结果通知
- **网页控制台**实例串行控制台(xterm)与 VNC(noVNC),两跳 SSH 隧道 - **网页控制台**浏览器内使用 xterm 串行终端和 noVNC,通过 OCI 控制台连接建立两跳 SSH 隧道
- **AI 网关**OpenAI / Anthropic 兼容端点转发 OCI GenAI(对话/Responses/Embeddings),号池渠道加权负载均衡、熔断探测、密钥管理与用量日志 - **身份与审计**IAM 用户、MFA、API Key、密码策略、SAML 身份提供商、通知收件人和多 Identity Domain 管理;OCI Audit 事件可经 Service Connector Hub 与 Notifications 回传,关键事件按类别通过「云端事件」通知推送
- **日志回传**OCI Audit 事件经 Connector Hub → Notifications HTTPS 订阅回传入库,一键创建链路;自定义告警规则(事件类型/来源 IP 白名单/资源/频率阈值)命中即推送 - **通知与安全**Telegram、Webhook、ntfy、Bark、SMTP 五类渠道;JWT、bcrypt、TOTP、OIDC / GitHub 登录、登录锁定、IP 限速、会话撤销和系统操作审计
- **租户治理**IAM 用户/MFA/API Key 管理、密码策略、身份提供商(SAML)、通知收件人;多 Identity Domain 租户可按域切换管理 - **AI 网关**:提供 OpenAI Responses、Chat Completions、Embeddings 与 Anthropic Messages 兼容接口,支持渠道分组、加权路由、熔断探测、模型黑白名单、密钥管理和调用日志
- **安全**JWT + bcrypt(凭据变更旧令牌立即失效,可一键撤销全部会话)、TOTP 两步验证、OIDC/GitHub 外部登录、登录锁定、IP 限速、真实 IP 头可配、请求体/超时防护、系统操作审计
- **通知**:模板化推送,五渠道并存(Telegram / Webhook / ntfy / Bark / SMTP),Webhook 可对接飞书、钉钉、Slack、企业微信机器人 ## 运行形态
```text
浏览器 / API 客户端
┌──────────────────── OCI Portal 单个 Go 进程 ────────────────────┐
│ Vue 3 静态资源(go:embed
│ /api/v1 → Gin → Service → OCI SDK │
│ /ai/v1 → 兼容转换与路由 → OCI Generative AI │
│ GORM → SQLite(默认)/ MySQL / PostgreSQLexperimental
└─────────────────────────────────────────────────────────────────┘
```
默认推荐 SQLite 单实例部署。MySQL 和 PostgreSQL 适配仍属 experimental,不应据此推断服务支持多副本并发运行。
## 快速开始 ## 快速开始
### Docker Compose(推荐)
前置条件:Docker、Docker Compose v2、OpenSSL。
1. 克隆仓库并生成一份需要长期保存的 `.env`
```bash
git clone https://github.com/wangdefaa/oci-portal.git
cd oci-portal
umask 077
printf 'DATA_KEY=%s\nJWT_SECRET=%s\nADMIN_PASSWORD=%s\n' \
"$(openssl rand -hex 32)" \
"$(openssl rand -hex 32)" \
"$(openssl rand -base64 24)" > .env
chmod 600 .env
```
2. 准备数据目录并启动:
```bash
mkdir -p data
# Linux bind mount 需要让镜像内的 nonroot 用户(uid 65532)可写。
sudo chown 65532:65532 data
docker compose up -d
docker compose ps
```
3. 查看初始管理员密码并登录:
```bash
grep '^ADMIN_PASSWORD=' .env
```
访问 `http://127.0.0.1:18888`,默认用户名为 `admin`。
> `DATA_KEY` 用于解密数据库中的 OCI 私钥、口令和渠道凭据。它只能生成一次并持续复用;丢失或更换后,已有密文无法恢复。请将 `.env` 与 `data/oci-portal.db` 一起备份。
### 二进制运行 ### 二进制运行
Release 下载对应架构的二进制后 Release 提供 Linux amd64 / arm64 二进制。以下以 amd64 为例;arm64 主机将文件名中的 `amd64` 替换为 `arm64`
```bash ```bash
DATA_KEY=$(openssl rand -hex 32) JWT_SECRET=$(openssl rand -hex 32) ADMIN_PASSWORD=<初始密码> ./oci-portal-server curl -fLO https://github.com/wangdefaa/oci-portal/releases/latest/download/oci-portal-server-linux-amd64
chmod +x oci-portal-server-linux-amd64
# 复用上文生成并妥善保存的 .env。
set -a
. ./.env
set +a
./oci-portal-server-linux-amd64
``` ```
访问 `http://localhost:8080`,用 `admin` 与初始密码登录 默认访问地址为 `http://localhost:8080`。`ADMIN_PASSWORD` 只在数据库中没有用户时创建初始管理员,后续启动不会用它重置密码
### Docker Compose ### 源码构建
源码构建要求 Go 1.26.5、`curl`、`unzip` 和 `sha256sum`。仓库只保留前端占位页,编译完整单文件前应下载 `DASH_VERSION` 指定的前端产物并校验:
```bash ```bash
# 镜像以 distroless nonroot(uid 65532)运行,数据卷需可写,否则 SQLite 报 unable to open database file DASH_TAG="$(tr -d '\r\n' < DASH_VERSION)"
mkdir -p ./data && sudo chown 65532:65532 ./data DASH_BASE="https://github.com/wangdefaa/oci-portal-dash/releases/download/${DASH_TAG}"
DATA_KEY=$(openssl rand -hex 32) JWT_SECRET=$(openssl rand -hex 32) ADMIN_PASSWORD=<初始密码> docker compose up -d
curl -fL -o dist.zip "${DASH_BASE}/dist.zip"
curl -fL -o dist.zip.sha256 "${DASH_BASE}/dist.zip.sha256"
sha256sum -c dist.zip.sha256
rm -rf internal/webui/dist
mkdir -p internal/webui/dist
unzip -q dist.zip -d internal/webui/dist
CGO_ENABLED=0 go build -trimpath -o bin/oci-portal-server ./cmd/server
``` ```
默认只监听 `127.0.0.1:18888`;公网访问须置于 TLS 反向代理之后,见下文「反向代理」 本地构建显示 `dev` 版本;Release 工作流会注入正式版本和构建时间。`internal/webui/dist` 中的真实前端产物不应提交到仓库
### 源码构建(单文件,含前端) ## 生产部署
```bash 服务本身只提供 HTTP。除本机试用外,应保持服务仅监听回环地址或容器内部网络,并由 TLS 反向代理提供 HTTPS;否则管理员密码、JWT、AI 密钥和租户凭据会经明文连接传输。
# 1. 获取前端产物:从前端仓库 Release 下载 dist.zip(或本地 npm run build 后拷入)
curl -fL -o dist.zip https://github.com/wangdefaa/oci-portal-dash/releases/latest/download/dist.zip
rm -rf internal/webui/dist && mkdir -p internal/webui/dist && unzip -q dist.zip -d internal/webui/dist
# 2. 编译(免 CGO,可交叉编译;-X 两项注入「设置·关于」页的版本与构建时间,可省略,省略则显示 dev)
CGO_ENABLED=0 go build -trimpath \
-ldflags "-s -w -X oci-portal/internal/api.buildVersion=v0.0.1 -X oci-portal/internal/api.buildTime=$(date -u +%Y-%m-%dT%H:%M:%SZ)" \
-o bin/oci-portal-server ./cmd/server
```
> 注意:上述解压会覆盖仓库占位文件 `internal/webui/dist/index.html`,提交代码前勿把真实产物带入版本库。 ### Caddy
### 反向代理(公网部署必读)
面板自身只提供 HTTP,管理员口令、JWT 与租户 API 私钥都会经明文承载。除本机试用外,应让面板仅监听回环地址(compose 示例已默认 `127.0.0.1:18888`),由支持 TLS 的反向代理对外提供 HTTPS。
Caddy 最小示例(整站反代,自动申请并续期 Let's Encrypt 证书,WebSocket 自动透传):
```caddyfile ```caddyfile
portal.example.com { portal.example.com {
@@ -88,10 +152,18 @@ portal.example.com {
} }
``` ```
nginx 示例(证书自备;Web Console 的 WebSocket 升级与长超时必须显式配置,否则串行终端连不上或空闲即断) Caddy 会自动处理 WebSocket。使用 nginx、Traefik 或其他反向代理时,还需要满足
- 为串行终端和 VNC 转发 WebSocket Upgrade 头
- 为 AI 流式响应和控制台连接设置较长的读写超时,建议 `1h`
- 整站请求体上限至少为 `10MB`;后端会继续限制面板 API 为 `1MB`、AI 网关为 `10MB`
- 将面板内「设置 → 安全 → 真实 IP 请求头」配置为反向代理实际写入的头;系统审计、登录锁定和 IP 限速依赖它
- 配置 `PUBLIC_URL` 或面板地址,供 OAuth 回调和 OCI 日志回传链路使用
<details>
<summary>nginx 最小示例</summary>
```nginx ```nginx
# http 块内:按请求是否升级为 WebSocket 决定 Connection 头
map $http_upgrade $connection_upgrade { map $http_upgrade $connection_upgrade {
default upgrade; default upgrade;
"" close; "" close;
@@ -104,7 +176,6 @@ server {
ssl_certificate /etc/nginx/certs/portal.example.com.crt; ssl_certificate /etc/nginx/certs/portal.example.com.crt;
ssl_certificate_key /etc/nginx/certs/portal.example.com.key; ssl_certificate_key /etc/nginx/certs/portal.example.com.key;
# 放行到最大入口(AI 网关 10MB);各入口更细的上限由后端自身执行
client_max_body_size 10m; client_max_body_size 10m;
location / { location / {
@@ -114,69 +185,273 @@ server {
proxy_set_header Connection $connection_upgrade; proxy_set_header Connection $connection_upgrade;
proxy_set_header Host $host; proxy_set_header Host $host;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
# WebSocket 空闲与 AI 流式响应都需要长读超时(默认 60s 会断)
proxy_read_timeout 1h; proxy_read_timeout 1h;
proxy_send_timeout 1h; proxy_send_timeout 1h;
} }
} }
``` ```
Traefik 示例(已运行 Traefik 的 Docker 环境,面板容器加 label 接入即可,WebSocket 透明支持): </details>
```yaml 若前端静态文件与后端分离部署,`/api/*` 和 `/ai/*` 必须代理到后端,其余路径由 SPA 静态服务处理并回退到 `index.html`。Traefik 通过容器网络连接时应直达容器端口 `8080`,无需暴露宿主机端口。
services:
oci-portal:
image: ghcr.io/wangdefaa/oci-portal:latest
# ……环境变量与数据卷同上文 compose 示例;与 Traefik 同一 docker 网络,
# Traefik 经容器网络直达 8080,无需(也不应)对外映射 ports
labels:
- traefik.enable=true
- traefik.http.routers.oci-portal.rule=Host(`portal.example.com`)
- traefik.http.routers.oci-portal.entrypoints=websecure
- traefik.http.routers.oci-portal.tls.certresolver=le # 换成你的 certResolver 名
- traefik.http.services.oci-portal.loadbalancer.server.port=8080
```
若前端静态文件由反代直接伺服(不走内嵌页面),则按前缀代理,缺一不可: ## AI 网关
| 前缀 | 内容 | 何时需要 | AI 网关使用面板创建的独立密钥鉴权,支持 `Authorization: Bearer sk-...` 和 `x-api-key: sk-...`。密钥可绑定渠道分组和模型白名单;全局模型黑名单会从模型列表、路由和探测候选中同时排除目标模型。
| 端点 | 定位 | 流式 |
| --- | --- | --- | | --- | --- | --- |
| `/api/*` | 面板 REST、Web Console WebSocket、日志回传 webhook | 始终 | | `POST /ai/v1/responses` | OpenAI Responses,无状态主接口 | SSE;服务端工具除外 |
| `/ai/*` | AI 网关(OpenAI / Claude 兼容端点,独立密钥鉴权) | 启用 AI 网关时 | | `POST /ai/v1/chat/completions` | OpenAI Chat Completions,存量客户端兼容层 | SSE |
| 其余路径 | 前端 SPA 静态文件(404 回退 `index.html` | 静态分离形态 | | `POST /ai/v1/messages` | Anthropic Messages 转换层 | SSE |
| `POST /ai/v1/embeddings` | OpenAI Embeddings | 否 |
| `GET /ai/v1/models` | 当前密钥可见的模型列表 | 否 |
- 反代默认追加的 `X-Forwarded-For` 用于还原真实客户端 IP,系统日志留痕、登录锁定与 IP 限速都依赖它 兼容边界:
- 请求体上限建议与后端一致:`/api/*` 1MB、`/ai/*` 10MBCaddy 用 `request_body` 按前缀分层)
## 环境变量 - 对话请求统一转发 OCI OpenAI 兼容面,当前供给以 `xai.`、`meta.`、`openai.` 前缀模型为主;Cohere Embeddings 不受该对话模型范围影响
- 网关不保存会话历史,客户端需要携带完整上下文;Responses 拒绝非空 `previous_response_id`、非 `null` `conversation` 和 `background:true`
- Responses 支持 xAI Grok `web_search` / `x_search` 服务端工具,但仅限非流式请求
- Responses 的 `reasoning.effort`、Messages 的 `output_config.effort` 和 Chat Completions 的 `reasoning_effort` 会传给上游,实际档位和效果由模型决定
- Chat Completions 只承担协议转换与兼容修复;新能力优先在 Responses 和 Messages 提供
- 单次请求最多尝试三个渠道;可重试错误会切换渠道,流式响应建立后不会换渠道重试
### 已知上游限制:大 system 区流式断流
实测(2026-07-13OCI 兼容面对 `instructions` 与 `tools` 合计超约 64.5KB 的**流式**请求会在发出少量事件后静默断开连接(无任何错误事件;同请求非流式正常),与模型、字符集、消息正文大小均无关——消息正文(`input`)不计入该限制。Chat Completions 与 Messages 的 system/developer 提示会转换为 `instructions`,因此 Claude Code 等自带大体量系统提示与工具定义的客户端极易触发。
网关侧应对:
- Messages 与 Chat Completions 的流式请求在客户端尚未收到任何输出时遭遇上游断流,会自动降级为非流式重做,并按标准事件/chunk 序列一次推送;调用日志记 `retries=1` 与降级标记
- Responses 直通因初始事件已转发、协议上无法透明降级,调用日志记「上游流提前终止」,客户端需自行回退非流式
- 应急规避:将超长 system 内容移入首条 user 消息正文可绕过该限制(正文不计入),但语义有别,根治有待上游修复
这里提供的是兼容接口而非 OpenAI / Anthropic 协议的完整实现。OCI OpenAI 兼容面的部分行为来自实测,未见 Oracle 文档承诺,可能随上游调整。路由与鉴权定义以 [Swagger YAML](docs/swagger.yaml) 或运行时 Swagger UI 为准;无法由 OpenAPI 完整表达的兼容边界列于上方。
### 字段兼容矩阵
以下矩阵以 2026-07-13 的 [OpenAI Responses](https://developers.openai.com/api/reference/resources/responses/methods/create)、[Chat Completions](https://developers.openai.com/api/reference/resources/chat/subresources/completions/methods/create)、[Embeddings](https://developers.openai.com/api/reference/resources/embeddings/methods/create) 和 [Anthropic Messages](https://platform.claude.com/docs/en/api/messages/create) 为标准基线,并与当前实现逐项核对。
这是一份兼容性快照,不替代 Swagger。标准接口和 OCI 上游都可能变化,最终行为以当前版本代码与实测为准。
| 标记 | 含义 |
| :---: | --- |
| ✅ | 网关直接支持 |
| ➡️ | 网关原样透传;是否生效由 OCI 上游决定 |
| 🔄 | 网关进行字段或协议转换后支持 |
| ◐ | 部分支持、存在前置条件或语义降级 |
| ⚠️ | 请求可被接受,但字段会被忽略 |
| ❌ | 网关在请求到达上游前拒绝 |
<details>
<summary><code>POST /ai/v1/responses</code> 对比 OpenAI Responses</summary>
Responses 是“原始 JSON 直通 + 本地门禁”。除 `store` 外,网关不会重建请求体;未知顶层字段也会保留并送往 OCI。
| 标准字段 | 状态 | 网关行为 |
| --- | :---: | --- |
| `model` | ◐ | 必须非空,还要通过密钥模型白名单并匹配可用渠道;随后原样透传 |
| `input` | ➡️ | 支持标准的字符串或 item 数组,原始内容透传;网关不逐项保证 OCI 能处理所有 item 类型 |
| `instructions`、`max_output_tokens` | ➡️ | 类型可解析后原样透传,不做范围或模型能力校验 |
| `temperature`、`top_p`、`parallel_tool_calls` | ➡️ | 原样透传,不做取值范围校验 |
| `text` / `text.format` / `text.verbosity` | ➡️ | 整个原始对象透传;结构化输出是否可用由 OCI 模型决定 |
| `reasoning` | ➡️ | 整个原始对象透传;`effort` 不校验档位,`summary` 等字段不会被网关删除 |
| `tool_choice` | ➡️ | 任意 JSON 原样透传,本地不校验枚举或结构 |
| `context_management`、`include`、`max_tool_calls`、`metadata`、`moderation`、`prompt` | ➡️ | 网关未建模但会保留在原始请求中,由 OCI 决定是否接受 |
| `prompt_cache_key`、`prompt_cache_retention`、`safety_identifier`、`service_tier` | ➡️ | 原样透传,不代表 OCI 一定实现 OpenAI 的对应语义 |
| `stream_options`、`top_logprobs`、`truncation`、`user` | ➡️ | 原样透传,不做本地语义校验 |
| `store` | 🔄 | 无论客户端传什么,上游请求都强制改写为 `false` |
| `previous_response_id` | ❌ | 非空即返回 400;网关不保存历史响应 |
| `conversation` | ◐ | 省略或 `null` 可通过;任何非 `null` 值返回 400 |
| `background` | ◐ | `true` 返回 400`false`、`null` 或省略时继续透传 |
| `stream` | ◐ | 普通请求和 `function` 工具支持 SSE;含 `web_search` / `x_search` 时拒绝流式 |
| `tools[].type=function` | ➡️ | 非流式和流式均可,工具对象原样透传 |
| `tools[].type=web_search` / `x_search` | ◐ | 使用 xAI 服务端工具语义且仅支持非流式;`x_search` 不是 OpenAI 标准工具 |
| 其他标准工具类型 | ❌ | `web_search_preview`、`file_search`、`computer`、`code_interpreter`、`image_generation`、`mcp`、`shell`、`custom` 等返回 400 |
| 其他未知顶层字段 | ➡️ | 只要整个请求是合法 JSON,字段名、结构和数值都会保留 |
响应边界:
- 非流式成功响应不做转换,OCI JSON 原样返回;usage 解析只用于面板调用日志
- SSE 事件逐行原样转发,不补 `data: [DONE]`,推理增量也不会被网关过滤
- 流建立前最多切换三个渠道;流建立后中断不重试,客户端可能只收到部分事件
- 未知模型返回 404、无渠道返回 503;上游错误会套入 OpenAI 风格错误体,不保证与标准 OpenAI 错误字段完全相同
实现依据:[`airesponses.go`](internal/service/airesponses.go) · [`responses.go`](internal/aiwire/responses.go) · [`aigateway.go`](internal/api/aigateway.go)
</details>
<details>
<summary><code>POST /ai/v1/chat/completions</code> 对比 OpenAI Chat Completions</summary>
Chat Completions 会先转换为 Responses 请求,再把 OCI Responses 响应桥接回 Chat Completions 形态。
| 标准字段 | 状态 | 网关行为 |
| --- | :---: | --- |
| `model` | ◐ | 必填,受密钥白名单和可用渠道限制;模型名保留到上游请求 |
| `messages` | 🔄 | 必填并转换为 Responses `input` / `instructions` |
| `system` / `developer` 消息 | ◐ | 文本按出现顺序合并为 `instructions`;块数组中的非文本内容被忽略 |
| `user` / `assistant` 文本内容 | 🔄 | 字符串及 `text` 块分别转为 `input_text` / `output_text` |
| `image_url` 内容块 | ◐ | URL 或 data URI 转为 `input_image``image_url.detail` 被忽略 |
| `input_audio`、`file`、`refusal` 等内容块 | ❌ | 当前消息转换器不支持,返回 400 |
| assistant `tool_calls` / `role=tool` | 🔄 | 转为 `function_call` / `function_call_output`,保留调用 ID、函数名和参数 |
| 消息 `name`、`refusal`、`audio`、旧式 `function_call` | ⚠️ | 当前消息结构未建模,静默忽略 |
| `max_completion_tokens` | 🔄 | 转为 `max_output_tokens`,优先于 `max_tokens` |
| `max_tokens` | 🔄 | 未提供 `max_completion_tokens` 时转为 `max_output_tokens` |
| `temperature`、`top_p`、`parallel_tool_calls` | ◐ | 原值写入 Responses 请求,但不校验范围或模型能力 |
| `stream` | 🔄 | OCI Responses SSE 桥接为 `chat.completion.chunk`,末尾补 `data: [DONE]`;上游断流且尚无输出时自动降级非流式重做,结果按 chunk 序列一次推送 |
| `stream_options.include_usage` | 🔄 | 控制网关在终块后追加 `choices: []` 的 usage 块 |
| `stream_options.include_obfuscation` | ⚠️ | 未建模,静默忽略 |
| `tools[].type=function` | ◐ | `name`、`description`、`parameters` 支持;`function.strict` 被忽略 |
| `tools[].type=custom` 及其他工具 | ❌ | 只接受 `function`,其他类型返回 400 |
| `tool_choice` | ◐ | 支持 `auto` / `none` / `required` 和具名 function;非法或未知值被静默忽略 |
| `response_format` | ◐ | `json_object`、`json_schema` 转为 Responses `text.format`;未知类型交给 OCI 处理 |
| `reasoning_effort` | ◐ | 转小写后映射为 `reasoning.effort`,不校验模型或档位 |
| `store` | ⚠️ / 🔄 | 客户端字段被忽略,转换后的上游请求始终使用 `store:false` |
| `stop`、`seed`、`n`、`frequency_penalty`、`presence_penalty` | ⚠️ | 静默忽略;`n` 因此恒为单个 choice |
| `logprobs`、`top_logprobs`、`logit_bias`、`user` | ⚠️ | 静默忽略 |
| `audio`、`modalities`、`prediction`、`metadata`、`moderation` | ⚠️ | 静默忽略 |
| `prompt_cache_key`、`safety_identifier`、`service_tier`、`verbosity`、`web_search_options` | ⚠️ | 静默忽略 |
| 其他未知字段 | ⚠️ | JSON 解码器不会拒绝未知字段,但转换后的上游请求不包含它们 |
响应边界:
- 非流式固定生成一个 `choices[0]`;文本会合并,函数调用转为 `tool_calls`
- `finish_reason` 只生成 `stop`、`tool_calls`、`length`;其他上游终止原因不保留
- reasoning 输出、logprobs、refusal、annotations、audio、service tier 和 system fingerprint 不返回
- usage 保留 `prompt_tokens`、`completion_tokens`、`total_tokens` 和 `prompt_tokens_details.cached_tokens`
- 已建立的流中断或上游 `response.failed` 不会转换成标准 SSE 错误事件
实现依据:[`chatresponses.go`](internal/service/chatresponses.go) · [`openai.go`](internal/aiwire/openai.go) · [`aigateway.go`](internal/api/aigateway.go)
</details>
<details>
<summary><code>POST /ai/v1/embeddings</code> 对比 OpenAI Embeddings</summary>
| 标准字段 | 状态 | 网关行为 |
| --- | :---: | --- |
| `model` | ◐ | 必填,受密钥白名单限制,并且必须存在具有 `EMBEDDING` 能力的渠道 |
| `input` 为字符串 | ✅ | 包装为单个输入后调用 OCI |
| `input` 为字符串数组 | ✅ | 按原顺序调用 OCI |
| `input` 为 token ID 数组或二维 token ID 数组 | ❌ | 只能解码字符串或字符串数组,绑定阶段返回 400 |
| 空数组 / `null` | ❌ | 本地返回 400;空字符串不在本地拒绝,由 OCI 决定 |
| `dimensions` | ◐ | 映射为 OCI 输出维度,不做范围或模型能力校验 |
| `encoding_format=float` | ✅ | 返回 float 数组;省略时行为相同 |
| `encoding_format=base64` | ❌ | 本地返回 400,不提供 base64 响应 |
| `user` | ⚠️ | 能解析但不会传给 OCI |
| 其他未知字段 | ⚠️ | 静默忽略 |
响应使用标准的 `object:"list"`、`data[].object:"embedding"`、`index`、`model` 和可选 `usage` 外壳;向量为 `float32` 数组,不支持流式。
实现依据:[`embeddings.go`](internal/aiwire/embeddings.go) · [`aigateway_chat.go`](internal/service/aigateway_chat.go) · [`aigateway.go`](internal/api/aigateway.go)
</details>
<details>
<summary><code>POST /ai/v1/messages</code> 对比 Anthropic Messages</summary>
网关接受 `Authorization: Bearer` 或 `x-api-key`,但不会校验或使用标准 Anthropic `anthropic-version`、`anthropic-beta` 请求头。
| 标准字段 | 状态 | 网关行为 |
| --- | :---: | --- |
| `model` | ◐ | 用于模型与渠道选择,但空字符串不会在 handler 中按参数错误拒绝,通常最终返回模型不存在 |
| `max_tokens` | 🔄 | 可缺省(缺省或 ≤0 时按默认值 8192),转换为 `max_output_tokens` |
| `messages` | ◐ | 必须非空;角色、交替顺序和空内容不做完整校验 |
| `system` | ◐ | 支持字符串或 text 块数组;多个文本块直接拼接,`cache_control` 等附加字段被忽略 |
| `temperature`、`top_p` | ◐ | 写入 Responses 请求,不做取值范围或模型能力校验 |
| `top_k` | ⚠️ | 能解析但不会传给上游 |
| `stop_sequences` | ⚠️ | 能解析但不会传给上游;响应 `stop_sequence` 恒为 `null` |
| `stream` | 🔄 | Responses SSE 桥接为 Anthropic 事件序列;上游断流且尚无输出时自动降级非流式重做,结果按事件序列一次推送 |
| `tools` | ◐ | 每个工具都转换成 Responses `function`;自定义客户端工具可用,Anthropic 服务端工具类型不保留原语义 |
| `tool_choice` | ◐ | 支持 `auto`、`any`、`none`、具名 `tool``disable_parallel_tool_use` 等附加字段被忽略 |
| `metadata` | ⚠️ | 能解析但不会传给上游 |
| `thinking` | ⚠️ | 顶层 thinking 配置不会控制上游思考预算 |
| `output_config.effort` | 🔄 | 转小写后映射为 Responses `reasoning.effort` |
| `output_config` 其他子字段 | ⚠️ | 未建模,静默忽略 |
| `cache_control`、`container`、`inference_geo`、`service_tier` | ⚠️ | 标准 SDK 中存在,但当前请求结构未建模,静默忽略 |
| 其他未知顶层字段 | ⚠️ | JSON 解码器接受,但转换后的上游请求不包含它们 |
`messages[].content`
| 标准内容块 | 状态 | 网关行为 |
| --- | :---: | --- |
| 字符串 / `text` | 🔄 | user 转 `input_text`assistant 历史转 `output_text` |
| `image` | ◐ | 仅支持 `base64` 和 `url` source;缺字段或其他 source 类型返回 400 |
| `tool_use` | 🔄 | 转为 `function_call`,保留 ID、名称和输入 |
| `tool_result` | ◐ | 转为 `function_call_output`;块数组只拼接 text`is_error` 和非文本结果丢失 |
| `thinking` / `redacted_thinking` | ⚠️ | 历史思考块被删除,不进入上游上下文 |
| `document`、服务端工具结果及其他未知块 | ❌ | 返回 400 |
| `null` / 空块数组 | ◐ | 该消息可能从上游 input 中消失,不返回参数错误 |
响应边界:
- 非流式只把 Responses `output_text` 转成 `text`、`function_call` 转成 `tool_use`
- `stop_reason` 只生成 `end_turn`、`tool_use`、`max_tokens``stop_sequence` 恒为 `null`
- reasoning 不会生成 Anthropic `thinking` / `redacted_thinking` 块,也没有 signature
- usage 只保留 `input_tokens`、`output_tokens` 和 `cache_read_input_tokens`,不提供 `cache_creation_input_tokens`
- 流式输出标准事件骨架,但不生成 `thinking_delta` 和 `signature_delta`;上游错误事件与无终态断流会转成 Anthropic `error` 事件
实现依据:[`anthresponses.go`](internal/service/anthresponses.go) · [`anthropic.go`](internal/aiwire/anthropic.go) · [`aigateway.go`](internal/api/aigateway.go)
</details>
`GET /ai/v1/models` 使用 OpenAI Models 列表外壳(`object`、`data[].id/object/created/owned_by`),但只返回当前渠道目录中通过分组、全局黑名单和密钥白名单筛选后的模型;网关不提供标准的单模型检索端点。
## API 与配置
### API 文档
| 路径 | 鉴权 | 用途 |
| --- | --- | --- |
| `/api/v1/*` | 登录返回的 JWT Bearer Token | 面板管理 API |
| `/ai/v1/*` | AI 网关密钥 | OpenAI / Anthropic 兼容接口 |
| `/api/v1/webhooks/oci-logs/:secret` | URL 中的独立回传密钥 | OCI Notifications 日志回传 |
OpenAPI 文件随仓库维护:[`docs/swagger.yaml`](docs/swagger.yaml) · [`docs/swagger.json`](docs/swagger.json)。
运行进程时设置 `SWAGGER=1` 可开放 `/swagger/index.html`。Swagger 默认关闭,生产环境建议仅在受控网络内按需开启。接口注释变更后必须重新生成 OpenAPI。
### 环境变量
| 变量 | 必填 | 默认值 | 说明 | | 变量 | 必填 | 默认值 | 说明 |
| --- | --- | --- | --- | | --- | :---: | --- | --- |
| `DATA_KEY` | 是 | | 敏感字段加密主密钥(更换后已入库密文无法解密) | | `DATA_KEY` | 是 | | 敏感字段加密主密钥;必须持久保存,不能随意轮换 |
| `JWT_SECRET` | 是 | | 登录令牌签名密钥 | | `JWT_SECRET` | 是 | | JWT 签名密钥;更换会使已有登录令牌失效 |
| `ADMIN_USERNAME` | 否 | `admin` | 初始管理员用户名 | | `ADMIN_USERNAME` | 否 | `admin` | 初始管理员用户名 |
| `ADMIN_PASSWORD` | 首次启动 | | 仅在用户不存在时创建;已存在不重置 | | `ADMIN_PASSWORD` | 首次启动 | | 仅在数据库无用户时创建管理员,不会重置已有密码 |
| `ADDR` | 否 | `:8080` | HTTP 监听地址 | | `ADDR` | 否 | `:8080` | HTTP 监听地址 |
| `DB_DRIVER` | 否 | `sqlite` | `sqlite` / `mysql` / `postgres`后两者 experimental | | `DB_DRIVER` | 否 | `sqlite` | `sqlite` / `mysql` / `postgres`后两者 experimental |
| `DB_DSN` | 外部库时是 | 无 | MySQL`parseTime=True`PostgreSQL 标准 DSN | | `DB_PATH` | SQLite | `oci-portal.db` | SQLite 文件路径 |
| `DB_PATH` | 否 | `oci-portal.db` | SQLite 文件路径 | | `DB_DSN` | 外部数据库 | — | MySQL 需 `parseTime=True`;不要在日志或文档中暴露凭据 |
| `PUBLIC_URL` | 否 | | 面板公网基址,日志回传一键创建链路用 | | `PUBLIC_URL` | 否 | | 面板公网基址,作为 OAuth 回调和日志回传引导的回退值 |
| `HTTPS_PROXY` | 否 | 无 | 出站代理(如 Telegram 通知走代理) | | `TZ` | 否 | 系统时区 | cron 表达式的解释时区;容器示例使用 `Asia/Shanghai` |
| `TZ` | 否 | 系统 | cron 表达式解释时区(容器内建议显式设置,二进制已嵌 tzdata) | | `HTTP_PROXY` / `HTTPS_PROXY` / `NO_PROXY` | 否 | | Go 标准出站代理变量;面板内显式代理配置优先用于对应业务 |
| `SWAGGER` | 否 | 关 | `1` 时开放 `/swagger/index.html` API 文档(生产建议按需临时开启) | | `SWAGGER` | 否 | 关 | 设为 `1` 时开放 Swagger UI |
| `GIN_MODE` | 否 | `release` | `debug` / `release`,影响日志与调试输出 | | `GIN_MODE` | 否 | `release` | `debug` / `release` |
## 升级与备份
- 升级前同时备份 `.env` 和数据库;SQLite Compose 部署的数据文件为 `data/oci-portal.db`
- 保持 `DATA_KEY` 不变;只恢复数据库而没有原密钥,敏感字段无法解密
- 服务启动时会自动执行数据库迁移;跨版本升级前先阅读 [CHANGELOG](CHANGELOG.md)
- Compose 部署使用 `docker compose pull && docker compose up -d` 更新镜像
- OCI API Key 应遵循最小权限原则;生产环境保持 Swagger 关闭并限制管理面访问来源
## 开发 ## 开发
```bash ```bash
go test ./... # 全量测试 gofmt -l .
go vet ./... && gofmt -l . go vet ./...
go tool swag init -g cmd/server/main.go -o docs --parseInternal --parseDependency # 接口注释变更后重新生成 OpenAPI go test ./...
# Handler 注释变更后重新生成唯一的对外 API 文档。
go tool swag init -g cmd/server/main.go -o docs --parseInternal --parseDependency
``` ```
API 文档:全部接口带 swaggo 注释,`SWAGGER=1` 启动后访问 `/swagger/index.html`spec 文件在 `docs/swagger.json|yaml` - Go 后端规范与目录约定:[`AGENTS.md`](AGENTS.md) · [`.trellis/spec/backend/`](.trellis/spec/backend/)
- 前端开发与构建:[oci-portal-dash](https://github.com/wangdefaa/oci-portal-dash)
编码规范与项目约定见 [AGENTS.md](AGENTS.md) 与 [.trellis/spec/](.trellis/spec/)。 - 版本变更:[`CHANGELOG.md`](CHANGELOG.md)
## License ## License
+3 -1
View File
@@ -135,7 +135,7 @@ func run() error {
defer aiGateway.Wait() defer aiGateway.Wait()
defer stopCleanup() defer stopCleanup()
tasks := service.NewTaskService(db, ociConfigs, notifier, settings) tasks := service.NewTaskService(db, ociConfigs, notifier, settings)
ociConfigs.SetTenantCleanupDeps(tasks, logEvents) ociConfigs.SetTenantCleanupDeps(tasks)
tasks.AttachAiGateway(aiGateway) tasks.AttachAiGateway(aiGateway)
aiGateway.SetOnChannelsChanged(tasks.SyncAiProbeTask) aiGateway.SetOnChannelsChanged(tasks.SyncAiProbeTask)
if err := tasks.Start(); err != nil { if err := tasks.Start(); err != nil {
@@ -147,6 +147,8 @@ func run() error {
console := service.NewConsoleService(ociConfigs) console := service.NewConsoleService(ociConfigs)
proxies := service.NewProxyService(db, cipher) proxies := service.NewProxyService(db, cipher)
defer proxies.Wait() defer proxies.Wait()
// 「关于」页存储指标需知数据库形态
api.SetAboutRuntime(cfg.DBDriver, cfg.DBPath)
router := api.NewRouter(auth, oauth, ociConfigs, tasks, console, settings, notifier, systemLogs, logEvents, proxies, aiGateway) router := api.NewRouter(auth, oauth, ociConfigs, tasks, console, settings, notifier, systemLogs, logEvents, proxies, aiGateway)
return serveHTTP(cfg.Addr, router) return serveHTTP(cfg.Addr, router)
} }
+136 -159
View File
@@ -20,10 +20,10 @@ const docTemplate = `{
"tags": [ "tags": [
"AI 网关" "AI 网关"
], ],
"summary": "OpenAI 兼容对话补全", "summary": "OpenAI Chat Completions 兼容端点",
"parameters": [ "parameters": [
{ {
"description": "OpenAI chat/completions 请求体(支持 stream)", "description": "OpenAI chat/completions 请求体(支持 stream;经 Responses 转换直通)",
"name": "body", "name": "body",
"in": "body", "in": "body",
"required": true, "required": true,
@@ -34,7 +34,7 @@ const docTemplate = `{
], ],
"responses": { "responses": {
"200": { "200": {
"description": "OpenAI 兼容响应(流式为 SSE)", "description": "OpenAI 兼容响应(流式为 SSE,末尾 data: [DONE])",
"schema": { "schema": {
"type": "object", "type": "object",
"additionalProperties": true "additionalProperties": true
@@ -79,7 +79,7 @@ const docTemplate = `{
"summary": "Anthropic Messages 兼容端点", "summary": "Anthropic Messages 兼容端点",
"parameters": [ "parameters": [
{ {
"description": "Anthropic messages 请求体(支持 stream)", "description": "Anthropic messages 请求体(支持 stream;经 OCI OpenAI 兼容面直通;max_tokens 可缺省,默认 8192)",
"name": "body", "name": "body",
"in": "body", "in": "body",
"required": true, "required": true,
@@ -124,7 +124,7 @@ const docTemplate = `{
"summary": "OpenAI Responses 兼容端点", "summary": "OpenAI Responses 兼容端点",
"parameters": [ "parameters": [
{ {
"description": "OpenAI responses 请求体(支持 stream)", "description": "OpenAI responses 请求体(支持 stream;web_search/x_search 服务端工具仅非流式)",
"name": "body", "name": "body",
"in": "body", "in": "body",
"required": true, "required": true,
@@ -154,7 +154,7 @@ const docTemplate = `{
"tags": [ "tags": [
"设置" "设置"
], ],
"summary": "返回构建运行环境信息,「设置 · 关于」页展示", "summary": "返回构建运行时长与资源占用信息,「设置 · 关于」页展示",
"responses": { "responses": {
"200": { "200": {
"description": "OK", "description": "OK",
@@ -166,6 +166,86 @@ const docTemplate = `{
} }
} }
}, },
"/api/v1/ai-blacklist": {
"get": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"AI 管理"
],
"summary": "模型黑名单列表",
"responses": {
"200": {
"description": "OK",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
},
"post": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"AI 管理"
],
"summary": "添加模型黑名单",
"parameters": [
{
"description": "请求体:{name: 模型名}",
"name": "body",
"in": "body",
"required": true,
"schema": {
"type": "object"
}
}
],
"responses": {
"201": {
"description": "Created",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
}
},
"/api/v1/ai-blacklist/{id}": {
"delete": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"AI 管理"
],
"summary": "移除模型黑名单(重新同步后模型恢复入池)",
"parameters": [
{
"type": "integer",
"description": "黑名单条目 ID",
"name": "id",
"in": "path",
"required": true
}
],
"responses": {
"204": {
"description": "无内容"
}
}
}
},
"/api/v1/ai-channels": { "/api/v1/ai-channels": {
"get": { "get": {
"security": [ "security": [
@@ -1102,124 +1182,6 @@ const docTemplate = `{
} }
} }
}, },
"/api/v1/log-events/alert-rules": {
"get": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"任务与日志回传"
],
"summary": "告警规则列表",
"responses": {
"200": {
"description": "OK",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
},
"post": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"任务与日志回传"
],
"summary": "创建告警规则",
"parameters": [
{
"description": "请求体",
"name": "body",
"in": "body",
"required": true,
"schema": {
"$ref": "#/definitions/internal_api.alertRuleRequest"
}
}
],
"responses": {
"201": {
"description": "Created",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
}
},
"/api/v1/log-events/alert-rules/{ruleId}": {
"put": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"任务与日志回传"
],
"summary": "更新告警规则",
"parameters": [
{
"type": "integer",
"description": "规则 ID",
"name": "ruleId",
"in": "path",
"required": true
},
{
"description": "请求体",
"name": "body",
"in": "body",
"required": true,
"schema": {
"$ref": "#/definitions/internal_api.alertRuleRequest"
}
}
],
"responses": {
"200": {
"description": "OK",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
},
"delete": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"任务与日志回传"
],
"summary": "删除告警规则",
"parameters": [
{
"type": "integer",
"description": "规则 ID",
"name": "ruleId",
"in": "path",
"required": true
}
],
"responses": {
"204": {
"description": "无内容"
}
}
}
},
"/api/v1/oci-configs": { "/api/v1/oci-configs": {
"get": { "get": {
"security": [ "security": [
@@ -5039,6 +5001,46 @@ const docTemplate = `{
} }
} }
}, },
"/api/v1/proxies/{id}/tenants": {
"put": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"设置"
],
"summary": "设置代理关联的租户全集",
"parameters": [
{
"type": "integer",
"description": "代理 ID",
"name": "id",
"in": "path",
"required": true
},
{
"description": "请求体 {\\",
"name": "body",
"in": "body",
"required": true,
"schema": {
"type": "object"
}
}
],
"responses": {
"200": {
"description": "usedBy",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
}
},
"/api/v1/regions": { "/api/v1/regions": {
"get": { "get": {
"security": [ "security": [
@@ -5738,7 +5740,7 @@ const docTemplate = `{
"tags": [ "tags": [
"任务与日志回传" "任务与日志回传"
], ],
"summary": "立即执行任务", "summary": "立即执行任务(异步触发,结果经任务日志轮询获取)",
"parameters": [ "parameters": [
{ {
"type": "integer", "type": "integer",
@@ -5749,8 +5751,15 @@ const docTemplate = `{
} }
], ],
"responses": { "responses": {
"200": { "202": {
"description": "OK", "description": "Accepted",
"schema": {
"type": "object",
"additionalProperties": true
}
},
"409": {
"description": "任务正在执行中",
"schema": { "schema": {
"type": "object", "type": "object",
"additionalProperties": true "additionalProperties": true
@@ -5804,38 +5813,6 @@ const docTemplate = `{
} }
}, },
"definitions": { "definitions": {
"internal_api.alertRuleRequest": {
"type": "object",
"properties": {
"enabled": {
"type": "boolean"
},
"eventTypes": {
"type": "string"
},
"name": {
"type": "string"
},
"ociConfigId": {
"type": "integer"
},
"resourceMatch": {
"type": "string"
},
"sourceIpMode": {
"type": "string"
},
"sourceIps": {
"type": "string"
},
"threshold": {
"type": "integer"
},
"windowMinutes": {
"type": "integer"
}
}
},
"internal_api.attachBootVolumeRequest": { "internal_api.attachBootVolumeRequest": {
"type": "object", "type": "object",
"required": [ "required": [
+136 -159
View File
@@ -13,10 +13,10 @@
"tags": [ "tags": [
"AI 网关" "AI 网关"
], ],
"summary": "OpenAI 兼容对话补全", "summary": "OpenAI Chat Completions 兼容端点",
"parameters": [ "parameters": [
{ {
"description": "OpenAI chat/completions 请求体(支持 stream)", "description": "OpenAI chat/completions 请求体(支持 stream;经 Responses 转换直通)",
"name": "body", "name": "body",
"in": "body", "in": "body",
"required": true, "required": true,
@@ -27,7 +27,7 @@
], ],
"responses": { "responses": {
"200": { "200": {
"description": "OpenAI 兼容响应(流式为 SSE)", "description": "OpenAI 兼容响应(流式为 SSE,末尾 data: [DONE])",
"schema": { "schema": {
"type": "object", "type": "object",
"additionalProperties": true "additionalProperties": true
@@ -72,7 +72,7 @@
"summary": "Anthropic Messages 兼容端点", "summary": "Anthropic Messages 兼容端点",
"parameters": [ "parameters": [
{ {
"description": "Anthropic messages 请求体(支持 stream)", "description": "Anthropic messages 请求体(支持 stream;经 OCI OpenAI 兼容面直通;max_tokens 可缺省,默认 8192)",
"name": "body", "name": "body",
"in": "body", "in": "body",
"required": true, "required": true,
@@ -117,7 +117,7 @@
"summary": "OpenAI Responses 兼容端点", "summary": "OpenAI Responses 兼容端点",
"parameters": [ "parameters": [
{ {
"description": "OpenAI responses 请求体(支持 stream)", "description": "OpenAI responses 请求体(支持 stream;web_search/x_search 服务端工具仅非流式)",
"name": "body", "name": "body",
"in": "body", "in": "body",
"required": true, "required": true,
@@ -147,7 +147,7 @@
"tags": [ "tags": [
"设置" "设置"
], ],
"summary": "返回构建运行环境信息,「设置 · 关于」页展示", "summary": "返回构建运行时长与资源占用信息,「设置 · 关于」页展示",
"responses": { "responses": {
"200": { "200": {
"description": "OK", "description": "OK",
@@ -159,6 +159,86 @@
} }
} }
}, },
"/api/v1/ai-blacklist": {
"get": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"AI 管理"
],
"summary": "模型黑名单列表",
"responses": {
"200": {
"description": "OK",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
},
"post": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"AI 管理"
],
"summary": "添加模型黑名单",
"parameters": [
{
"description": "请求体:{name: 模型名}",
"name": "body",
"in": "body",
"required": true,
"schema": {
"type": "object"
}
}
],
"responses": {
"201": {
"description": "Created",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
}
},
"/api/v1/ai-blacklist/{id}": {
"delete": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"AI 管理"
],
"summary": "移除模型黑名单(重新同步后模型恢复入池)",
"parameters": [
{
"type": "integer",
"description": "黑名单条目 ID",
"name": "id",
"in": "path",
"required": true
}
],
"responses": {
"204": {
"description": "无内容"
}
}
}
},
"/api/v1/ai-channels": { "/api/v1/ai-channels": {
"get": { "get": {
"security": [ "security": [
@@ -1095,124 +1175,6 @@
} }
} }
}, },
"/api/v1/log-events/alert-rules": {
"get": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"任务与日志回传"
],
"summary": "告警规则列表",
"responses": {
"200": {
"description": "OK",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
},
"post": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"任务与日志回传"
],
"summary": "创建告警规则",
"parameters": [
{
"description": "请求体",
"name": "body",
"in": "body",
"required": true,
"schema": {
"$ref": "#/definitions/internal_api.alertRuleRequest"
}
}
],
"responses": {
"201": {
"description": "Created",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
}
},
"/api/v1/log-events/alert-rules/{ruleId}": {
"put": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"任务与日志回传"
],
"summary": "更新告警规则",
"parameters": [
{
"type": "integer",
"description": "规则 ID",
"name": "ruleId",
"in": "path",
"required": true
},
{
"description": "请求体",
"name": "body",
"in": "body",
"required": true,
"schema": {
"$ref": "#/definitions/internal_api.alertRuleRequest"
}
}
],
"responses": {
"200": {
"description": "OK",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
},
"delete": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"任务与日志回传"
],
"summary": "删除告警规则",
"parameters": [
{
"type": "integer",
"description": "规则 ID",
"name": "ruleId",
"in": "path",
"required": true
}
],
"responses": {
"204": {
"description": "无内容"
}
}
}
},
"/api/v1/oci-configs": { "/api/v1/oci-configs": {
"get": { "get": {
"security": [ "security": [
@@ -5032,6 +4994,46 @@
} }
} }
}, },
"/api/v1/proxies/{id}/tenants": {
"put": {
"security": [
{
"BearerAuth": []
}
],
"tags": [
"设置"
],
"summary": "设置代理关联的租户全集",
"parameters": [
{
"type": "integer",
"description": "代理 ID",
"name": "id",
"in": "path",
"required": true
},
{
"description": "请求体 {\\",
"name": "body",
"in": "body",
"required": true,
"schema": {
"type": "object"
}
}
],
"responses": {
"200": {
"description": "usedBy",
"schema": {
"type": "object",
"additionalProperties": true
}
}
}
}
},
"/api/v1/regions": { "/api/v1/regions": {
"get": { "get": {
"security": [ "security": [
@@ -5731,7 +5733,7 @@
"tags": [ "tags": [
"任务与日志回传" "任务与日志回传"
], ],
"summary": "立即执行任务", "summary": "立即执行任务(异步触发,结果经任务日志轮询获取)",
"parameters": [ "parameters": [
{ {
"type": "integer", "type": "integer",
@@ -5742,8 +5744,15 @@
} }
], ],
"responses": { "responses": {
"200": { "202": {
"description": "OK", "description": "Accepted",
"schema": {
"type": "object",
"additionalProperties": true
}
},
"409": {
"description": "任务正在执行中",
"schema": { "schema": {
"type": "object", "type": "object",
"additionalProperties": true "additionalProperties": true
@@ -5797,38 +5806,6 @@
} }
}, },
"definitions": { "definitions": {
"internal_api.alertRuleRequest": {
"type": "object",
"properties": {
"enabled": {
"type": "boolean"
},
"eventTypes": {
"type": "string"
},
"name": {
"type": "string"
},
"ociConfigId": {
"type": "integer"
},
"resourceMatch": {
"type": "string"
},
"sourceIpMode": {
"type": "string"
},
"sourceIps": {
"type": "string"
},
"threshold": {
"type": "integer"
},
"windowMinutes": {
"type": "integer"
}
}
},
"internal_api.attachBootVolumeRequest": { "internal_api.attachBootVolumeRequest": {
"type": "object", "type": "object",
"required": [ "required": [
+88 -102
View File
@@ -1,26 +1,5 @@
basePath: / basePath: /
definitions: definitions:
internal_api.alertRuleRequest:
properties:
enabled:
type: boolean
eventTypes:
type: string
name:
type: string
ociConfigId:
type: integer
resourceMatch:
type: string
sourceIpMode:
type: string
sourceIps:
type: string
threshold:
type: integer
windowMinutes:
type: integer
type: object
internal_api.attachBootVolumeRequest: internal_api.attachBootVolumeRequest:
properties: properties:
bootVolumeId: bootVolumeId:
@@ -602,7 +581,7 @@ paths:
/ai/v1/chat/completions: /ai/v1/chat/completions:
post: post:
parameters: parameters:
- description: OpenAI chat/completions 请求体(支持 stream) - description: OpenAI chat/completions 请求体(支持 stream;经 Responses 转换直通)
in: body in: body
name: body name: body
required: true required: true
@@ -610,11 +589,11 @@ paths:
type: object type: object
responses: responses:
"200": "200":
description: OpenAI 兼容响应(流式为 SSE) description: 'OpenAI 兼容响应(流式为 SSE,末尾 data: [DONE])'
schema: schema:
additionalProperties: true additionalProperties: true
type: object type: object
summary: OpenAI 兼容对话补全 summary: OpenAI Chat Completions 兼容端点
tags: tags:
- AI 网关 - AI 网关
/ai/v1/embeddings: /ai/v1/embeddings:
@@ -638,7 +617,8 @@ paths:
/ai/v1/messages: /ai/v1/messages:
post: post:
parameters: parameters:
- description: Anthropic messages 请求体(支持 stream) - description: Anthropic messages 请求体(支持 stream;经 OCI OpenAI 兼容面直通;max_tokens
可缺省,默认 8192)
in: body in: body
name: body name: body
required: true required: true
@@ -667,7 +647,7 @@ paths:
/ai/v1/responses: /ai/v1/responses:
post: post:
parameters: parameters:
- description: OpenAI responses 请求体(支持 stream) - description: OpenAI responses 请求体(支持 stream;web_search/x_search 服务端工具仅非流式)
in: body in: body
name: body name: body
required: true required: true
@@ -692,9 +672,57 @@ paths:
type: object type: object
security: security:
- BearerAuth: [] - BearerAuth: []
summary: 返回构建运行环境信息,「设置 · 关于」页展示 summary: 返回构建运行时长与资源占用信息,「设置 · 关于」页展示
tags: tags:
- 设置 - 设置
/api/v1/ai-blacklist:
get:
responses:
"200":
description: OK
schema:
additionalProperties: true
type: object
security:
- BearerAuth: []
summary: 模型黑名单列表
tags:
- AI 管理
post:
parameters:
- description: '请求体:{name: 模型名}'
in: body
name: body
required: true
schema:
type: object
responses:
"201":
description: Created
schema:
additionalProperties: true
type: object
security:
- BearerAuth: []
summary: 添加模型黑名单
tags:
- AI 管理
/api/v1/ai-blacklist/{id}:
delete:
parameters:
- description: 黑名单条目 ID
in: path
name: id
required: true
type: integer
responses:
"204":
description: 无内容
security:
- BearerAuth: []
summary: 移除模型黑名单(重新同步后模型恢复入池)
tags:
- AI 管理
/api/v1/ai-channels: /api/v1/ai-channels:
get: get:
responses: responses:
@@ -1267,78 +1295,6 @@ paths:
summary: 回传日志事件列表 summary: 回传日志事件列表
tags: tags:
- 任务与日志回传 - 任务与日志回传
/api/v1/log-events/alert-rules:
get:
responses:
"200":
description: OK
schema:
additionalProperties: true
type: object
security:
- BearerAuth: []
summary: 告警规则列表
tags:
- 任务与日志回传
post:
parameters:
- description: 请求体
in: body
name: body
required: true
schema:
$ref: '#/definitions/internal_api.alertRuleRequest'
responses:
"201":
description: Created
schema:
additionalProperties: true
type: object
security:
- BearerAuth: []
summary: 创建告警规则
tags:
- 任务与日志回传
/api/v1/log-events/alert-rules/{ruleId}:
delete:
parameters:
- description: 规则 ID
in: path
name: ruleId
required: true
type: integer
responses:
"204":
description: 无内容
security:
- BearerAuth: []
summary: 删除告警规则
tags:
- 任务与日志回传
put:
parameters:
- description: 规则 ID
in: path
name: ruleId
required: true
type: integer
- description: 请求体
in: body
name: body
required: true
schema:
$ref: '#/definitions/internal_api.alertRuleRequest'
responses:
"200":
description: OK
schema:
additionalProperties: true
type: object
security:
- BearerAuth: []
summary: 更新告警规则
tags:
- 任务与日志回传
/api/v1/oci-configs: /api/v1/oci-configs:
get: get:
responses: responses:
@@ -3713,6 +3669,31 @@ paths:
summary: 手动重测代理出口地区,同步返回最新视图 summary: 手动重测代理出口地区,同步返回最新视图
tags: tags:
- 设置 - 设置
/api/v1/proxies/{id}/tenants:
put:
parameters:
- description: 代理 ID
in: path
name: id
required: true
type: integer
- description: 请求体 {\
in: body
name: body
required: true
schema:
type: object
responses:
"200":
description: usedBy
schema:
additionalProperties: true
type: object
security:
- BearerAuth: []
summary: 设置代理关联的租户全集
tags:
- 设置
/api/v1/proxies/import: /api/v1/proxies/import:
post: post:
parameters: parameters:
@@ -4158,14 +4139,19 @@ paths:
required: true required: true
type: integer type: integer
responses: responses:
"200": "202":
description: OK description: Accepted
schema:
additionalProperties: true
type: object
"409":
description: 任务正在执行中
schema: schema:
additionalProperties: true additionalProperties: true
type: object type: object
security: security:
- BearerAuth: [] - BearerAuth: []
summary: 立即执行任务 summary: 立即执行任务(异步触发,结果经任务日志轮询获取)
tags: tags:
- 任务与日志回传 - 任务与日志回传
/api/v1/webhooks/oci-logs/{secret}: /api/v1/webhooks/oci-logs/{secret}:
+7
View File
@@ -20,6 +20,13 @@ type MessagesRequest struct {
ToolChoice json.RawMessage `json:"tool_choice,omitempty"` ToolChoice json.RawMessage `json:"tool_choice,omitempty"`
Metadata json.RawMessage `json:"metadata,omitempty"` Metadata json.RawMessage `json:"metadata,omitempty"`
Thinking json.RawMessage `json:"thinking,omitempty"` Thinking json.RawMessage `json:"thinking,omitempty"`
// OutputConfig 是 Anthropic 官方输出控制(effort 等);网关映射 effort 到 OCI reasoningEffort。
OutputConfig *AnthOutputConfig `json:"output_config,omitempty"`
}
// AnthOutputConfig 是 Anthropic Messages 顶层 output_config 对象。
type AnthOutputConfig struct {
Effort string `json:"effort,omitempty"`
} }
// SystemText 把顶层 system(string 或块数组)拍平为文本。 // SystemText 把顶层 system(string 或块数组)拍平为文本。
+94 -121
View File
@@ -1,6 +1,6 @@
// Package aiwire 定义 AI 网关的线格式类型:OpenAI Chat Completions(兼作网关内部 // Package aiwire 定义 AI 网关的线格式类型:OpenAI Responses / Chat Completions /
// 中间表示 IR,因 OCI GENERIC 格式即其镜像)与 Anthropic Messages。零内部依赖, // Anthropic Messages / Embeddings 与通用记账、错误、模型列表结构。上游统一为
// oci 层与 service 层共同引用 // OCI OpenAI 兼容面
package aiwire package aiwire
import ( import (
@@ -8,28 +8,6 @@ import (
"strings" "strings"
) )
// ChatRequest 是 OpenAI /v1/chat/completions 请求体;TopK 为 OCI/Anthropic 扩展字段。
type ChatRequest struct {
Model string `json:"model"`
Messages []ChatMessage `json:"messages"`
MaxTokens *int `json:"max_tokens,omitempty"`
MaxCompletionTokens *int `json:"max_completion_tokens,omitempty"`
Temperature *float64 `json:"temperature,omitempty"`
TopP *float64 `json:"top_p,omitempty"`
TopK *int `json:"top_k,omitempty"`
FrequencyPenalty *float64 `json:"frequency_penalty,omitempty"`
PresencePenalty *float64 `json:"presence_penalty,omitempty"`
Stop StringList `json:"stop,omitempty"`
Seed *int `json:"seed,omitempty"`
N *int `json:"n,omitempty"`
Stream bool `json:"stream,omitempty"`
StreamOptions *StreamOptions `json:"stream_options,omitempty"`
Tools []Tool `json:"tools,omitempty"`
ToolChoice json.RawMessage `json:"tool_choice,omitempty"`
ResponseFormat *ResponseFormat `json:"response_format,omitempty"`
ReasoningEffort string `json:"reasoning_effort,omitempty"`
}
// StringList 兼容 OpenAI stop 的 string 与 []string 两种线格式。 // StringList 兼容 OpenAI stop 的 string 与 []string 两种线格式。
type StringList []string type StringList []string
@@ -50,7 +28,75 @@ func (s *StringList) UnmarshalJSON(b []byte) error {
return nil return nil
} }
// StreamOptions 控制流式末尾是否附带 usage 事件 // Usage 是 token 用量
type Usage struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
// PromptTokensDetails 携带缓存命中细分;OCI 为隐式 prompt 缓存,只有读取计数
PromptTokensDetails *PromptTokensDetails `json:"prompt_tokens_details,omitempty"`
}
// PromptTokensDetails 是输入 token 细分(OpenAI prompt_tokens_details 同构)。
type PromptTokensDetails struct {
CachedTokens int `json:"cached_tokens"`
}
// CachedTokens 返回缓存命中读取的 token 数;无细分时为 0。
func (u *Usage) CachedTokens() int {
if u == nil || u.PromptTokensDetails == nil {
return 0
}
return u.PromptTokensDetails.CachedTokens
}
// Model 是 /v1/models 列表项。
type Model struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
OwnedBy string `json:"owned_by"`
}
// ModelList 是 /v1/models 响应体。
type ModelList struct {
Object string `json:"object"`
Data []Model `json:"data"`
}
// ErrorBody 是 OpenAI 格式错误响应。
type ErrorBody struct {
Error ErrorDetail `json:"error"`
}
// ErrorDetail 是错误明细。
type ErrorDetail struct {
Message string `json:"message"`
Type string `json:"type"`
Code string `json:"code,omitempty"`
}
// ---- Chat Completions 线格式(/ai/v1/chat/completions,Tier 2 兼容层)----
// ChatRequest 是 OpenAI /v1/chat/completions 请求体;只声明网关读取的字段,
// 无对应物的字段(stop / seed / n / penalty 等)由 JSON 解码天然忽略。
type ChatRequest struct {
Model string `json:"model"`
Messages []ChatMessage `json:"messages"`
MaxTokens *int `json:"max_tokens,omitempty"`
MaxCompletionTokens *int `json:"max_completion_tokens,omitempty"`
Temperature *float64 `json:"temperature,omitempty"`
TopP *float64 `json:"top_p,omitempty"`
ParallelToolCalls *bool `json:"parallel_tool_calls,omitempty"`
Stream bool `json:"stream,omitempty"`
StreamOptions *StreamOptions `json:"stream_options,omitempty"`
Tools []Tool `json:"tools,omitempty"`
ToolChoice json.RawMessage `json:"tool_choice,omitempty"`
ResponseFormat *ResponseFormat `json:"response_format,omitempty"`
ReasoningEffort string `json:"reasoning_effort,omitempty"`
}
// StreamOptions 控制流式末尾是否附带 usage 块。
type StreamOptions struct { type StreamOptions struct {
IncludeUsage bool `json:"include_usage"` IncludeUsage bool `json:"include_usage"`
} }
@@ -59,48 +105,31 @@ type StreamOptions struct {
type ChatMessage struct { type ChatMessage struct {
Role string `json:"role"` Role string `json:"role"`
Content Content `json:"content"` Content Content `json:"content"`
Name string `json:"name,omitempty"`
ToolCalls []ToolCall `json:"tool_calls,omitempty"` ToolCalls []ToolCall `json:"tool_calls,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"` ToolCallID string `json:"tool_call_id,omitempty"`
} }
// Content 兼容 string 与多模态块数组两种线格式;一期仅透传文本 // Content 兼容 string 与多模态块数组两种线格式,序列化时保形
type Content struct { type Content struct {
// Text 在原始值为字符串时填充
Text string Text string
// Parts 在原始值为数组时填充
Parts []ContentPart Parts []ContentPart
// isArray 记录原始形态,序列化时保形 IsArray bool
isArray bool
}
// ContentPart 是块数组中的一段;text 与 image_url 被网关处理,其余类型拒绝。
type ContentPart struct {
Type string `json:"type"`
Text string `json:"text,omitempty"`
ImageURL *ImageURL `json:"image_url,omitempty"`
}
// ImageURL 是 image_url 块负载;url 为 data URI(base64)或公网地址,与 OCI ImageUrl 同构。
type ImageURL struct {
URL string `json:"url"`
Detail string `json:"detail,omitempty"`
} }
func (c *Content) UnmarshalJSON(b []byte) error { func (c *Content) UnmarshalJSON(b []byte) error {
trimmed := strings.TrimSpace(string(b)) t := strings.TrimSpace(string(b))
if trimmed == "null" { if t == "null" {
return nil return nil
} }
if len(trimmed) > 0 && trimmed[0] == '[' { if strings.HasPrefix(t, "[") {
c.isArray = true c.IsArray = true
return json.Unmarshal(b, &c.Parts) return json.Unmarshal(b, &c.Parts)
} }
return json.Unmarshal(b, &c.Text) return json.Unmarshal(b, &c.Text)
} }
func (c Content) MarshalJSON() ([]byte, error) { func (c Content) MarshalJSON() ([]byte, error) {
if c.isArray { if c.IsArray {
return json.Marshal(c.Parts) return json.Marshal(c.Parts)
} }
return json.Marshal(c.Text) return json.Marshal(c.Text)
@@ -111,14 +140,9 @@ func NewTextContent(text string) Content {
return Content{Text: text} return Content{Text: text}
} }
// NewPartsContent 构造块数组形态的内容(多模态消息用)。 // JoinText 拍平内容为纯文本(数组形态拼接 text 块)。
func NewPartsContent(parts []ContentPart) Content {
return Content{Parts: parts, isArray: true}
}
// JoinText 返回全部文本内容(数组形态拼接 text 块)。
func (c Content) JoinText() string { func (c Content) JoinText() string {
if !c.isArray { if !c.IsArray {
return c.Text return c.Text
} }
var sb strings.Builder var sb strings.Builder
@@ -130,23 +154,20 @@ func (c Content) JoinText() string {
return sb.String() return sb.String()
} }
// HasUnsupported 报告是否含网关无法承接的内容块(text / image_url 之外的类型) // ContentPart 是块数组中的一段;text image_url 被网关承接,其余类型拒绝
func (c Content) HasUnsupported() bool { type ContentPart struct {
for _, p := range c.Parts { Type string `json:"type"`
switch p.Type { Text string `json:"text,omitempty"`
case "", "text": ImageURL *ImageURL `json:"image_url,omitempty"`
case "image_url":
if p.ImageURL == nil || p.ImageURL.URL == "" {
return true
}
default:
return true
}
}
return false
} }
// Tool 是 function 工具定义 // ImageURL 是 image_url 块负载;url 为 data URI(base64)或公网地址
type ImageURL struct {
URL string `json:"url"`
Detail string `json:"detail,omitempty"`
}
// Tool 是嵌套形态的 function 工具定义。
type Tool struct { type Tool struct {
Type string `json:"type"` Type string `json:"type"`
Function FunctionDef `json:"function"` Function FunctionDef `json:"function"`
@@ -178,7 +199,7 @@ type ResponseFormat struct {
JSONSchema json.RawMessage `json:"json_schema,omitempty"` JSONSchema json.RawMessage `json:"json_schema,omitempty"`
} }
// ChatResponse 是非流式响应体。 // ChatResponse 是非流式响应体(chat.completion)
type ChatResponse struct { type ChatResponse struct {
ID string `json:"id"` ID string `json:"id"`
Object string `json:"object"` Object string `json:"object"`
@@ -195,28 +216,6 @@ type Choice struct {
FinishReason string `json:"finish_reason"` FinishReason string `json:"finish_reason"`
} }
// Usage 是 token 用量。
type Usage struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
// PromptTokensDetails 携带缓存命中细分;OCI 为隐式 prompt 缓存,只有读取计数
PromptTokensDetails *PromptTokensDetails `json:"prompt_tokens_details,omitempty"`
}
// PromptTokensDetails 是输入 token 细分(OpenAI prompt_tokens_details 同构)。
type PromptTokensDetails struct {
CachedTokens int `json:"cached_tokens"`
}
// CachedTokens 返回缓存命中读取的 token 数;无细分时为 0。
func (u *Usage) CachedTokens() int {
if u == nil || u.PromptTokensDetails == nil {
return 0
}
return u.PromptTokensDetails.CachedTokens
}
// ChatChunk 是流式增量事件体(chat.completion.chunk)。 // ChatChunk 是流式增量事件体(chat.completion.chunk)。
type ChatChunk struct { type ChatChunk struct {
ID string `json:"id"` ID string `json:"id"`
@@ -254,29 +253,3 @@ type FunctionCallDelta struct {
Name string `json:"name,omitempty"` Name string `json:"name,omitempty"`
Arguments string `json:"arguments,omitempty"` Arguments string `json:"arguments,omitempty"`
} }
// Model 是 /v1/models 列表项。
type Model struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
OwnedBy string `json:"owned_by"`
}
// ModelList 是 /v1/models 响应体。
type ModelList struct {
Object string `json:"object"`
Data []Model `json:"data"`
}
// ErrorBody 是 OpenAI 格式错误响应。
type ErrorBody struct {
Error ErrorDetail `json:"error"`
}
// ErrorDetail 是错误明细。
type ErrorDetail struct {
Message string `json:"message"`
Type string `json:"type"`
Code string `json:"code,omitempty"`
}
-33
View File
@@ -156,39 +156,6 @@ type RespReasoning struct {
Effort string `json:"effort,omitempty"` Effort string `json:"effort,omitempty"`
} }
// Response 是 /responses 响应对象。
type Response struct {
ID string `json:"id"`
Object string `json:"object"`
CreatedAt int64 `json:"created_at"`
Status string `json:"status"`
Model string `json:"model"`
Output []RespOutItem `json:"output"`
Usage *RespUsage `json:"usage,omitempty"`
Error *RespError `json:"error"`
IncompleteDetails *RespIncomplete `json:"incomplete_details"`
Store bool `json:"store"`
}
// RespOutItem 是 output 数组项:message 或 function_call。
type RespOutItem struct {
Type string `json:"type"`
ID string `json:"id"`
Status string `json:"status,omitempty"`
Role string `json:"role,omitempty"`
Content []RespOutPart `json:"content,omitempty"`
CallID string `json:"call_id,omitempty"`
Name string `json:"name,omitempty"`
Arguments string `json:"arguments,omitempty"`
}
// RespOutPart 是 message item 的内容部件(output_text)。
type RespOutPart struct {
Type string `json:"type"`
Text string `json:"text"`
Annotations []any `json:"annotations"`
}
// RespUsage 是 Responses 口径的用量(input/output 命名)。 // RespUsage 是 Responses 口径的用量(input/output 命名)。
type RespUsage struct { type RespUsage struct {
InputTokens int `json:"input_tokens"` InputTokens int `json:"input_tokens"`
+51 -1
View File
@@ -1,8 +1,11 @@
package api package api
import ( import (
"math"
"net/http" "net/http"
"os"
"runtime" "runtime"
"time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@@ -12,11 +15,20 @@ import (
var ( var (
buildVersion = "dev" buildVersion = "dev"
buildTime = "" buildTime = ""
// startedAt 为进程启动时刻,运行时长与 CPU 均值的计算基准
startedAt = time.Now()
// aboutDB 由 main 启动时注入,存储指标据此定位数据库文件
aboutDB struct{ driver, path string }
) )
// SetAboutRuntime 注入数据库形态,「关于」页存储指标展示用;启动时调用一次。
func SetAboutRuntime(dbDriver, dbPath string) {
aboutDB.driver, aboutDB.path = dbDriver, dbPath
}
// about 返回构建与运行环境信息,「设置 · 关于」页展示。 // about 返回构建与运行环境信息,「设置 · 关于」页展示。
// //
// @Summary 返回构建运行环境信息,「设置 · 关于」页展示 // @Summary 返回构建运行时长与资源占用信息,「设置 · 关于」页展示
// @Tags 设置 // @Tags 设置
// @Success 200 {object} map[string]any // @Success 200 {object} map[string]any
// @Security BearerAuth // @Security BearerAuth
@@ -27,5 +39,43 @@ func about(c *gin.Context) {
"buildTime": buildTime, "buildTime": buildTime,
"goVersion": runtime.Version(), "goVersion": runtime.Version(),
"platform": runtime.GOOS + "/" + runtime.GOARCH, "platform": runtime.GOOS + "/" + runtime.GOARCH,
"startedAt": startedAt.UTC().Format(time.RFC3339),
"uptimeSeconds": int64(time.Since(startedAt).Seconds()),
"resources": resourcesSnapshot(),
}) })
} }
// resourcesSnapshot 采集进程资源占用:CPU 自启动均值、内存、goroutine 与存储。
func resourcesSnapshot() gin.H {
var ms runtime.MemStats
runtime.ReadMemStats(&ms)
pct := 0.0
if up := time.Since(startedAt).Seconds(); up > 0 {
// 多核并行时均值可超 100%,口径为「占单核百分比」
pct = processCPUTime().Seconds() / up * 100
}
return gin.H{
"cpuAvgPercent": math.Round(pct*100) / 100,
"numCpu": runtime.NumCPU(),
"goroutines": runtime.NumGoroutine(),
"memHeapBytes": ms.HeapAlloc,
"memSysBytes": ms.Sys,
"dbEngine": aboutDB.driver,
"dbBytes": dbFileBytes(aboutDB.driver, aboutDB.path),
}
}
// dbFileBytes 统计数据库磁盘占用:sqlite 累加主文件与 -wal / -shm,
// 外部数据库(mysql / postgres)不可从本机度量,返回 -1 表示不可用。
func dbFileBytes(driver, path string) int64 {
if driver != "sqlite" || path == "" {
return -1
}
var total int64
for _, p := range []string{path, path + "-wal", path + "-shm"} {
if fi, err := os.Stat(p); err == nil {
total += fi.Size()
}
}
return total
}
+8
View File
@@ -0,0 +1,8 @@
//go:build !unix
package api
import "time"
// processCPUTime 非 unix 平台无 getrusage,CPU 均值退化为 0。
func processCPUTime() time.Duration { return 0 }
+57
View File
@@ -0,0 +1,57 @@
package api
import (
"os"
"path/filepath"
"testing"
)
func TestDBFileBytes(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "app.db")
mustWrite := func(p string, n int) {
t.Helper()
if err := os.WriteFile(p, make([]byte, n), 0o600); err != nil {
t.Fatalf("write %s: %v", p, err)
}
}
mustWrite(dbPath, 100)
mustWrite(dbPath+"-wal", 30)
cases := []struct {
name string
driver string
path string
want int64
}{
{name: "sqlite 主文件加 wal", driver: "sqlite", path: dbPath, want: 130},
{name: "sqlite 文件不存在计 0", driver: "sqlite", path: filepath.Join(dir, "none.db"), want: 0},
{name: "外部数据库不可度量", driver: "postgres", path: "", want: -1},
{name: "路径为空不可度量", driver: "sqlite", path: "", want: -1},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if got := dbFileBytes(tc.driver, tc.path); got != tc.want {
t.Fatalf("dbFileBytes(%q, %q) = %d, want %d", tc.driver, tc.path, got, tc.want)
}
})
}
}
func TestResourcesSnapshot(t *testing.T) {
SetAboutRuntime("sqlite", filepath.Join(t.TempDir(), "absent.db"))
got := resourcesSnapshot()
if got["numCpu"].(int) < 1 || got["goroutines"].(int) < 1 {
t.Fatalf("numCpu / goroutines 应为正数: %+v", got)
}
if got["memHeapBytes"].(uint64) == 0 || got["memSysBytes"].(uint64) == 0 {
t.Fatalf("内存指标应为正数: %+v", got)
}
if got["cpuAvgPercent"].(float64) < 0 {
t.Fatalf("cpuAvgPercent 不应为负: %+v", got)
}
if got["dbEngine"].(string) != "sqlite" || got["dbBytes"].(int64) != 0 {
t.Fatalf("存储指标不符: %+v", got)
}
}
+18
View File
@@ -0,0 +1,18 @@
//go:build unix
package api
import (
"syscall"
"time"
)
// processCPUTime 返回进程自启动累计的 CPU 时间(用户态 + 内核态);
// 取不到时返回 0,「关于」页 CPU 均值显示为 0 而非报错。
func processCPUTime() time.Duration {
var ru syscall.Rusage
if err := syscall.Getrusage(syscall.RUSAGE_SELF, &ru); err != nil {
return 0
}
return time.Duration(ru.Utime.Nano() + ru.Stime.Nano())
}
+62 -2
View File
@@ -43,12 +43,13 @@ func (h *aiAdminHandler) createKey(c *gin.Context) {
Name string `json:"name" binding:"required"` Name string `json:"name" binding:"required"`
Value string `json:"value"` Value string `json:"value"`
Group string `json:"group"` Group string `json:"group"`
Models []string `json:"models"`
} }
if err := c.ShouldBindJSON(&req); err != nil { if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return return
} }
plaintext, key, err := h.gw.CreateKey(c.Request.Context(), req.Name, req.Value, req.Group) plaintext, key, err := h.gw.CreateKey(c.Request.Context(), req.Name, req.Value, req.Group, req.Models)
if err != nil { if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return return
@@ -72,12 +73,13 @@ func (h *aiAdminHandler) updateKey(c *gin.Context) {
Name string `json:"name"` Name string `json:"name"`
Enabled *bool `json:"enabled"` Enabled *bool `json:"enabled"`
Group *string `json:"group"` Group *string `json:"group"`
Models *[]string `json:"models"`
} }
if err := c.ShouldBindJSON(&req); err != nil { if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return return
} }
if err := h.gw.UpdateKey(c.Request.Context(), id, req.Name, req.Enabled, req.Group); err != nil { if err := h.gw.UpdateKey(c.Request.Context(), id, req.Name, req.Enabled, req.Group, req.Models); err != nil {
respondError(c, err) respondError(c, err)
return return
} }
@@ -285,6 +287,64 @@ func (h *aiAdminHandler) gatewayModels(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"items": list.Data}) c.JSON(http.StatusOK, gin.H{"items": list.Data})
} }
// ---- 模型黑名单 ----
// @Summary 模型黑名单列表
// @Tags AI 管理
// @Success 200 {object} map[string]any
// @Security BearerAuth
// @Router /api/v1/ai-blacklist [get]
func (h *aiAdminHandler) listBlacklist(c *gin.Context) {
items, err := h.gw.Blacklist(c.Request.Context())
if err != nil {
respondError(c, err)
return
}
c.JSON(http.StatusOK, gin.H{"items": items})
}
// addBlacklist 拉黑模型:删除全部渠道缓存中的同名条目,后续同步 / 探测均过滤。
//
// @Summary 添加模型黑名单
// @Tags AI 管理
// @Param body body object true "请求体:{name: 模型名}"
// @Success 201 {object} map[string]any
// @Security BearerAuth
// @Router /api/v1/ai-blacklist [post]
func (h *aiAdminHandler) addBlacklist(c *gin.Context) {
var req struct {
Name string `json:"name" binding:"required"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
item, err := h.gw.AddBlacklist(c.Request.Context(), req.Name)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusCreated, gin.H{"item": item})
}
// @Summary 移除模型黑名单(重新同步后模型恢复入池)
// @Tags AI 管理
// @Param id path int true "黑名单条目 ID"
// @Success 204 "无内容"
// @Security BearerAuth
// @Router /api/v1/ai-blacklist/{id} [delete]
func (h *aiAdminHandler) removeBlacklist(c *gin.Context) {
id, ok := aiPathID(c)
if !ok {
return
}
if err := h.gw.RemoveBlacklist(c.Request.Context(), id); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.Status(http.StatusNoContent)
}
// @Summary AI 调用日志列表 // @Summary AI 调用日志列表
// @Tags AI 管理 // @Tags AI 管理
// @Success 200 {object} map[string]any // @Success 200 {object} map[string]any
+355 -178
View File
@@ -1,13 +1,15 @@
package api package api
import ( import (
"bufio"
"bytes"
"crypto/rand" "crypto/rand"
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt"
"io" "io"
"net/http" "net/http"
"slices"
"strings" "strings"
"time" "time"
@@ -86,6 +88,25 @@ func keyGroup(c *gin.Context) string {
return "" return ""
} }
// keyModels 取当前请求密钥的模型白名单(空 = 不限模型)。
func keyModels(c *gin.Context) []string {
if key, ok := c.Get(aiKeyCtx); ok {
return key.(*model.AiKey).Models
}
return nil
}
// checkKeyModel 校验请求模型是否在密钥白名单内;拒绝时按端点协议返回
// 404 model_not_found(与未知模型同口径,不泄露密钥配置细节)并返回 false。
func checkKeyModel(c *gin.Context, modelName string) bool {
allowed := keyModels(c)
if len(allowed) == 0 || slices.Contains(allowed, modelName) {
return true
}
aiError(c, http.StatusNotFound, "model_not_found", "模型不存在或无权访问: "+modelName)
return false
}
// logEntry 组装调用日志骨架;usage / status 由调用方补齐。 // logEntry 组装调用日志骨架;usage / status 由调用方补齐。
func (h *aiGatewayHandler) logEntry(c *gin.Context, endpoint, modelName string, stream bool, meta service.ChatMeta, start time.Time) model.AiCallLog { func (h *aiGatewayHandler) logEntry(c *gin.Context, endpoint, modelName string, stream bool, meta service.ChatMeta, start time.Time) model.AiCallLog {
entry := model.AiCallLog{ entry := model.AiCallLog{
@@ -102,21 +123,6 @@ func (h *aiGatewayHandler) logEntry(c *gin.Context, endpoint, modelName string,
} }
// validateIR 做协议无关的入参检查(模型名、消息、不支持的内容块)。 // validateIR 做协议无关的入参检查(模型名、消息、不支持的内容块)。
func validateIR(ir aiwire.ChatRequest) error {
if strings.TrimSpace(ir.Model) == "" {
return fmt.Errorf("model 不能为空")
}
if len(ir.Messages) == 0 {
return fmt.Errorf("messages 不能为空")
}
for _, m := range ir.Messages {
if m.Content.HasUnsupported() {
return service.ErrAiUnsupportedBlock
}
}
return nil
}
// maybeLogContent 在调用密钥显式开启内容日志且未过期时写正文(红线例外); // maybeLogContent 在调用密钥显式开启内容日志且未过期时写正文(红线例外);
// respBody 为 nil 时只记请求(流式响应与向量结果不记录);callLogID 关联同次调用日志。 // respBody 为 nil 时只记请求(流式响应与向量结果不记录);callLogID 关联同次调用日志。
func (h *aiGatewayHandler) maybeLogContent(c *gin.Context, callLogID uint, endpoint, modelName string, stream bool, reqBody, respBody any) { func (h *aiGatewayHandler) maybeLogContent(c *gin.Context, callLogID uint, endpoint, modelName string, stream bool, reqBody, respBody any) {
@@ -147,47 +153,6 @@ func (h *aiGatewayHandler) logFailure(c *gin.Context, entry model.AiCallLog, req
h.maybeLogContent(c, callID, entry.Endpoint, entry.Model, entry.Stream, reqBody, nil) h.maybeLogContent(c, callID, entry.Endpoint, entry.Model, entry.Stream, reqBody, nil)
} }
// chatCompletions 是 OpenAI /ai/v1/chat/completions 端点。
//
// @Summary OpenAI 兼容对话补全
// @Tags AI 网关
// @Param body body object true "OpenAI chat/completions 请求体(支持 stream)"
// @Success 200 {object} map[string]any "OpenAI 兼容响应(流式为 SSE)"
// @Router /ai/v1/chat/completions [post]
func (h *aiGatewayHandler) chatCompletions(c *gin.Context) {
var ir aiwire.ChatRequest
if err := c.ShouldBindJSON(&ir); err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if err := validateIR(ir); err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if ir.Stream {
h.streamOpenAI(c, ir)
return
}
start := time.Now()
resp, meta, err := h.gw.Chat(c.Request.Context(), ir, keyGroup(c))
entry := h.logEntry(c, "openai", ir.Model, false, meta, start)
if err != nil {
upstreamError(c, err)
entry.ErrMsg = err.Error()
h.logFailure(c, entry, ir)
return
}
resp.ID = aiRandID("chatcmpl-")
if resp.Created == 0 {
resp.Created = time.Now().Unix()
}
entry.Status = http.StatusOK
fillUsage(&entry, resp.Usage)
callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "openai", ir.Model, false, ir, resp)
c.JSON(http.StatusOK, resp)
}
func fillUsage(entry *model.AiCallLog, u *aiwire.Usage) { func fillUsage(entry *model.AiCallLog, u *aiwire.Usage) {
if u == nil { if u == nil {
return return
@@ -217,6 +182,9 @@ func (h *aiGatewayHandler) embeddings(c *gin.Context) {
aiError(c, http.StatusBadRequest, "invalid_request_error", "仅支持 encoding_format=float") aiError(c, http.StatusBadRequest, "invalid_request_error", "仅支持 encoding_format=float")
return return
} }
if !checkKeyModel(c, req.Model) {
return
}
start := time.Now() start := time.Now()
resp, meta, err := h.gw.Embeddings(c.Request.Context(), req, keyGroup(c)) resp, meta, err := h.gw.Embeddings(c.Request.Context(), req, keyGroup(c))
entry := h.logEntry(c, "embeddings", req.Model, false, meta, start) entry := h.logEntry(c, "embeddings", req.Model, false, meta, start)
@@ -235,44 +203,6 @@ func (h *aiGatewayHandler) embeddings(c *gin.Context) {
c.JSON(http.StatusOK, resp) c.JSON(http.StatusOK, resp)
} }
// streamOpenAI 以 OpenAI SSE 直通流式响应,末尾发 [DONE]。
func (h *aiGatewayHandler) streamOpenAI(c *gin.Context, ir aiwire.ChatRequest) {
start := time.Now()
stream, meta, err := h.gw.OpenStream(c.Request.Context(), ir, keyGroup(c))
entry := h.logEntry(c, "openai", ir.Model, true, meta, start)
if err != nil {
upstreamError(c, err)
entry.ErrMsg = err.Error()
h.logFailure(c, entry, ir)
return
}
defer stream.Close()
sseHeaders(c)
id, created := aiRandID("chatcmpl-"), time.Now().Unix()
var usage *aiwire.Usage
for {
chunk, err := stream.Next()
if err != nil {
if !errors.Is(err, io.EOF) {
entry.ErrMsg = err.Error()
}
break
}
chunk.ID, chunk.Created = id, created
if chunk.Usage != nil {
usage = chunk.Usage
}
writeSSEData(c, chunk)
}
c.Writer.WriteString("data: [DONE]\n\n")
c.Writer.Flush()
entry.Status = http.StatusOK
entry.LatencyMs = time.Since(start).Milliseconds()
fillUsage(&entry, usage)
callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "openai", ir.Model, true, ir, nil)
}
func sseHeaders(c *gin.Context) { func sseHeaders(c *gin.Context) {
c.Header("Content-Type", "text/event-stream") c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache") c.Header("Cache-Control", "no-cache")
@@ -280,22 +210,15 @@ func sseHeaders(c *gin.Context) {
c.Writer.Flush() c.Writer.Flush()
} }
func writeSSEData(c *gin.Context, v any) { // anthDefaultMaxTokens 是 max_tokens 缺省时的默认输出上限:Anthropic 协议
b, err := json.Marshal(v) // 该字段必填,但部分客户端当选填不传,按默认值放行而非 400 拒绝。
if err != nil { const anthDefaultMaxTokens = 8192
return
}
c.Writer.WriteString("data: ")
c.Writer.Write(b)
c.Writer.WriteString("\n\n")
c.Writer.Flush()
}
// messages 是 Anthropic /ai/v1/messages 端点。 // messages 是 Anthropic /ai/v1/messages 端点。
// //
// @Summary Anthropic Messages 兼容端点 // @Summary Anthropic Messages 兼容端点
// @Tags AI 网关 // @Tags AI 网关
// @Param body body object true "Anthropic messages 请求体(支持 stream)" // @Param body body object true "Anthropic messages 请求体(支持 stream;经 OCI OpenAI 兼容面直通;max_tokens 可缺省,默认 8192)"
// @Success 200 {object} map[string]any "Anthropic 兼容响应(流式为 SSE)" // @Success 200 {object} map[string]any "Anthropic 兼容响应(流式为 SSE)"
// @Router /ai/v1/messages [post] // @Router /ai/v1/messages [post]
func (h *aiGatewayHandler) messages(c *gin.Context) { func (h *aiGatewayHandler) messages(c *gin.Context) {
@@ -305,72 +228,139 @@ func (h *aiGatewayHandler) messages(c *gin.Context) {
return return
} }
if req.MaxTokens <= 0 { if req.MaxTokens <= 0 {
aiError(c, http.StatusBadRequest, "invalid_request_error", "max_tokens 必填且需大于 0") req.MaxTokens = anthDefaultMaxTokens
}
if len(req.Messages) == 0 {
aiError(c, http.StatusBadRequest, "invalid_request_error", "messages 不能为空")
return return
} }
ir, err := service.AnthropicToIR(req) if !checkKeyModel(c, req.Model) {
return
}
body, err := service.AnthropicToResponsesBody(req)
if err != nil { if err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return return
} }
if err := validateIR(ir); err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if req.Stream { if req.Stream {
h.streamAnthropic(c, ir) h.streamAnthropic(c, body, req)
return return
} }
start := time.Now() start := time.Now()
resp, meta, err := h.gw.Chat(c.Request.Context(), ir, keyGroup(c)) payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
entry := h.logEntry(c, "anthropic", ir.Model, false, meta, start) entry := h.logEntry(c, "anthropic", req.Model, false, meta, start)
if err != nil { if err != nil {
upstreamError(c, err) upstreamError(c, err)
entry.ErrMsg = err.Error() entry.ErrMsg = err.Error()
h.logFailure(c, entry, req) h.logFailure(c, entry, req)
return return
} }
out, err := service.ResponsesToAnthropic(payload, aiRandID("msg_"))
if err != nil {
aiError(c, http.StatusBadGateway, "api_error", err.Error())
entry.ErrMsg = err.Error()
h.logFailure(c, entry, req)
return
}
entry.Status = http.StatusOK entry.Status = http.StatusOK
fillUsage(&entry, resp.Usage) fillUsage(&entry, service.RespPassthroughUsage(payload))
callID := h.gw.LogCall(entry) callID := h.gw.LogCall(entry)
out := service.IRRespToAnthropic(resp, aiRandID("msg_")) h.maybeLogContent(c, callID, "anthropic", req.Model, false, req, out)
h.maybeLogContent(c, callID, "anthropic", ir.Model, false, req, out)
c.JSON(http.StatusOK, out) c.JSON(http.StatusOK, out)
} }
// streamAnthropic 经状态机把 IR chunk 流聚合为 Anthropic SSE 事件流。 // streamAnthropic 流式直通并桥接:上游 Responses SSE 事件经 AnthRespBridge
func (h *aiGatewayHandler) streamAnthropic(c *gin.Context, ir aiwire.ChatRequest) { // 转为 Anthropic 事件序列逐块写出;usage 从 completed 事件记账。
func (h *aiGatewayHandler) streamAnthropic(c *gin.Context, body []byte, req aiwire.MessagesRequest) {
start := time.Now() start := time.Now()
stream, meta, err := h.gw.OpenStream(c.Request.Context(), ir, keyGroup(c)) upstream, meta, err := h.gw.RespPassthroughStream(c.Request.Context(), body, req.Model, keyGroup(c))
entry := h.logEntry(c, "anthropic", ir.Model, true, meta, start) entry := h.logEntry(c, "anthropic", req.Model, true, meta, start)
if err != nil { if err != nil {
upstreamError(c, err) upstreamError(c, err)
entry.ErrMsg = err.Error() entry.ErrMsg = err.Error()
h.logFailure(c, entry, ir) h.logFailure(c, entry, req)
return return
} }
defer stream.Close() defer upstream.Close()
sseHeaders(c) sseHeaders(c)
st := service.NewAnthStream(aiRandID("msg_"), ir.Model) bridge := service.NewAnthRespBridge(aiRandID("msg_"), req.Model)
for { emitted := 0
chunk, err := stream.Next() if err := forwardSSEData(upstream, func(data []byte) {
if err != nil { evs := bridge.Feed(data)
if !errors.Is(err, io.EOF) { emitted += len(evs)
writeAnthEvents(c, evs)
}); err != nil {
entry.ErrMsg = err.Error() entry.ErrMsg = err.Error()
} }
break // 上游断流且客户端尚未收到任何事件(OCI 兼容面对大请求 + 推理模型的
} // 流式通道会在 reasoning 阶段掐断):降级非流式重做,结果按事件序列推送
writeAnthEvents(c, st.Feed(chunk)) if emitted == 0 && !bridge.SawTerminal() && c.Request.Context().Err() == nil {
} if h.anthFallback(c, &entry, req) {
writeAnthEvents(c, st.Finish())
entry.Status = http.StatusOK entry.Status = http.StatusOK
entry.LatencyMs = time.Since(start).Milliseconds() entry.LatencyMs = time.Since(start).Milliseconds()
usage := st.Usage() callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "anthropic", req.Model, true, req, nil)
return
}
}
writeAnthEvents(c, bridge.Finish())
// 上游错误事件 / 提前终止:SSE 头已发出维持 200,错误落日志可查
if msg := bridge.Err(); msg != "" && entry.ErrMsg == "" {
entry.ErrMsg = msg
}
entry.Status = http.StatusOK
entry.LatencyMs = time.Since(start).Milliseconds()
usage := bridge.Usage()
entry.PromptTokens, entry.CompletionTokens = usage.InputTokens, usage.OutputTokens entry.PromptTokens, entry.CompletionTokens = usage.InputTokens, usage.OutputTokens
entry.TotalTokens = usage.InputTokens + usage.OutputTokens entry.TotalTokens = usage.InputTokens + usage.OutputTokens
entry.CachedTokens = usage.CacheReadInputTokens entry.CachedTokens = usage.CacheReadInputTokens
callID := h.gw.LogCall(entry) callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "anthropic", ir.Model, true, ir, nil) h.maybeLogContent(c, callID, "anthropic", req.Model, true, req, nil)
}
// anthFallback 用非流式重做同一请求并把完整结果按事件序列推送;
// 成功返回 true 并把渠道 / 用量 / 降级标注写入日志条目。
func (h *aiGatewayHandler) anthFallback(c *gin.Context, entry *model.AiCallLog, req aiwire.MessagesRequest) bool {
req.Stream = false
body, err := service.AnthropicToResponsesBody(req)
if err != nil {
return false
}
payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
if err != nil {
return false
}
out, err := service.ResponsesToAnthropic(payload, aiRandID("msg_"))
if err != nil {
return false
}
writeAnthEvents(c, service.AnthMessageEvents(out))
entry.ChannelID, entry.ChannelName = meta.ChannelID, meta.ChannelName
entry.Retries++
entry.ErrMsg = "流式上游断流,已降级非流式完成"
fillUsage(entry, service.RespPassthroughUsage(payload))
return true
}
// forwardSSEData 逐行读上游 SSE,把 data 行交给 emit 即时转换写出;
// Anthropic 与 Chat Completions 两条流式桥共用。
func forwardSSEData(upstream io.Reader, emit func([]byte)) error {
reader := bufio.NewReader(upstream)
for {
line, err := reader.ReadBytes('\n')
if len(line) > 0 {
trimmed := bytes.TrimSpace(line)
if data, ok := bytes.CutPrefix(trimmed, []byte("data: ")); ok {
emit(data)
}
}
if err != nil {
if errors.Is(err, io.EOF) {
return nil
}
return err
}
}
} }
func writeAnthEvents(c *gin.Context, events []service.AnthEvent) { func writeAnthEvents(c *gin.Context, events []service.AnthEvent) {
@@ -386,6 +376,154 @@ func writeAnthEvents(c *gin.Context, events []service.AnthEvent) {
c.Writer.Flush() c.Writer.Flush()
} }
// chatCompletions 是 OpenAI /ai/v1/chat/completions 端点(Tier 2 兼容层,
// 承接只会说 Chat Completions 的存量客户端)。
//
// @Summary OpenAI Chat Completions 兼容端点
// @Tags AI 网关
// @Param body body object true "OpenAI chat/completions 请求体(支持 stream;经 Responses 转换直通)"
// @Success 200 {object} map[string]any "OpenAI 兼容响应(流式为 SSE,末尾 data: [DONE])"
// @Router /ai/v1/chat/completions [post]
func (h *aiGatewayHandler) chatCompletions(c *gin.Context) {
var req aiwire.ChatRequest
if err := c.ShouldBindJSON(&req); err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if strings.TrimSpace(req.Model) == "" || len(req.Messages) == 0 {
aiError(c, http.StatusBadRequest, "invalid_request_error", "model 与 messages 不能为空")
return
}
if !checkKeyModel(c, req.Model) {
return
}
body, err := service.ChatToResponsesBody(req)
if err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if req.Stream {
h.streamChat(c, body, req)
return
}
h.chatOnce(c, body, req)
}
// chatOnce 处理非流式:直通上游 → 转回 chat.completion → 记账。
func (h *aiGatewayHandler) chatOnce(c *gin.Context, body []byte, req aiwire.ChatRequest) {
start := time.Now()
payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
entry := h.logEntry(c, "openai", req.Model, false, meta, start)
if err != nil {
upstreamError(c, err)
entry.ErrMsg = err.Error()
h.logFailure(c, entry, req)
return
}
out, err := service.ResponsesToChat(payload, aiRandID("chatcmpl-"), time.Now().Unix())
if err != nil {
aiError(c, http.StatusBadGateway, "api_error", err.Error())
entry.ErrMsg = err.Error()
h.logFailure(c, entry, req)
return
}
entry.Status = http.StatusOK
fillUsage(&entry, out.Usage)
callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "openai", req.Model, false, req, out)
c.JSON(http.StatusOK, out)
}
// streamChat 流式直通并桥接:上游 Responses SSE 经 ChatRespBridge 转为
// chat.completion.chunk 序列逐块写出,末尾发 [DONE];usage 从 completed 事件记账。
func (h *aiGatewayHandler) streamChat(c *gin.Context, body []byte, req aiwire.ChatRequest) {
start := time.Now()
upstream, meta, err := h.gw.RespPassthroughStream(c.Request.Context(), body, req.Model, keyGroup(c))
entry := h.logEntry(c, "openai", req.Model, true, meta, start)
if err != nil {
upstreamError(c, err)
entry.ErrMsg = err.Error()
h.logFailure(c, entry, req)
return
}
defer upstream.Close()
sseHeaders(c)
includeUsage := req.StreamOptions != nil && req.StreamOptions.IncludeUsage
bridge := service.NewChatRespBridge(aiRandID("chatcmpl-"), req.Model, time.Now().Unix(), includeUsage)
emitted := 0
if err := forwardSSEData(upstream, func(data []byte) {
chunks := bridge.Feed(data)
emitted += len(chunks)
writeChatChunks(c, chunks)
}); err != nil {
entry.ErrMsg = err.Error()
}
// 上游断流且客户端尚未收到任何块:降级非流式重做(成因同 messages 端点)
if emitted == 0 && bridge.Usage() == nil && c.Request.Context().Err() == nil {
if h.chatFallback(c, &entry, req, includeUsage) {
entry.Status = http.StatusOK
entry.LatencyMs = time.Since(start).Milliseconds()
callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "openai", req.Model, true, req, nil)
return
}
}
writeChatChunks(c, bridge.Finish())
c.Writer.WriteString("data: [DONE]\n\n")
c.Writer.Flush()
// 未见 completed 即结束:usage 缺失说明上游流异常提前终止,落日志可查
if bridge.Usage() == nil && entry.ErrMsg == "" {
entry.ErrMsg = "上游流提前终止,未返回终态事件"
}
entry.Status = http.StatusOK
entry.LatencyMs = time.Since(start).Milliseconds()
fillUsage(&entry, bridge.Usage())
callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "openai", req.Model, true, req, nil)
}
// chatFallback 用非流式重做同一请求并把完整结果按 chunk 序列推送;
// 成功返回 true 并把渠道 / 用量 / 降级标注写入日志条目。
func (h *aiGatewayHandler) chatFallback(c *gin.Context, entry *model.AiCallLog, req aiwire.ChatRequest, includeUsage bool) bool {
req.Stream = false
body, err := service.ChatToResponsesBody(req)
if err != nil {
return false
}
payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
if err != nil {
return false
}
out, err := service.ResponsesToChat(payload, aiRandID("chatcmpl-"), time.Now().Unix())
if err != nil {
return false
}
writeChatChunks(c, service.ChatResponseChunks(out, includeUsage))
c.Writer.WriteString("data: [DONE]\n\n")
c.Writer.Flush()
entry.ChannelID, entry.ChannelName = meta.ChannelID, meta.ChannelName
entry.Retries++
entry.ErrMsg = "流式上游断流,已降级非流式完成"
fillUsage(entry, service.RespPassthroughUsage(payload))
return true
}
// writeChatChunks 逐块写出 chunk 的 SSE data 行并 flush。
func writeChatChunks(c *gin.Context, chunks []aiwire.ChatChunk) {
for _, ch := range chunks {
b, err := json.Marshal(ch)
if err != nil {
continue
}
c.Writer.WriteString("data: ")
c.Writer.Write(b)
c.Writer.WriteString("\n\n")
}
if len(chunks) > 0 {
c.Writer.Flush()
}
}
// listModels 是 /ai/v1/models 端点(OpenAI 格式,从启用渠道的模型缓存聚合)。 // listModels 是 /ai/v1/models 端点(OpenAI 格式,从启用渠道的模型缓存聚合)。
// //
// @Summary 可用模型列表 // @Summary 可用模型列表
@@ -398,6 +536,15 @@ func (h *aiGatewayHandler) listModels(c *gin.Context) {
aiError(c, http.StatusInternalServerError, "api_error", "查询模型列表失败") aiError(c, http.StatusInternalServerError, "api_error", "查询模型列表失败")
return return
} }
if allowed := keyModels(c); len(allowed) > 0 {
kept := make([]aiwire.Model, 0, len(list.Data))
for _, m := range list.Data {
if slices.Contains(allowed, m.ID) {
kept = append(kept, m)
}
}
list.Data = kept
}
c.JSON(http.StatusOK, list) c.JSON(http.StatusOK, list)
} }
@@ -405,31 +552,45 @@ func (h *aiGatewayHandler) listModels(c *gin.Context) {
// //
// @Summary OpenAI Responses 兼容端点 // @Summary OpenAI Responses 兼容端点
// @Tags AI 网关 // @Tags AI 网关
// @Param body body object true "OpenAI responses 请求体(支持 stream)" // @Param body body object true "OpenAI responses 请求体(支持 stream;web_search/x_search 服务端工具仅非流式)"
// @Success 200 {object} map[string]any "OpenAI 兼容响应(流式为 SSE)" // @Success 200 {object} map[string]any "OpenAI 兼容响应(流式为 SSE)"
// @Router /ai/v1/responses [post] // @Router /ai/v1/responses [post]
func (h *aiGatewayHandler) responses(c *gin.Context) { func (h *aiGatewayHandler) responses(c *gin.Context) {
var req aiwire.RespRequest raw, err := c.GetRawData()
if err := c.ShouldBindJSON(&req); err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
ir, err := service.ResponsesToIR(req)
if err != nil { if err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return return
} }
if err := validateIR(ir); err != nil { var req aiwire.RespRequest
if err := json.Unmarshal(raw, &req); err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if !checkKeyModel(c, req.Model) {
return
}
h.responsesPassthrough(c, raw, req)
}
// responsesPassthrough 直通链路(唯一上游):原始 body 改写后直达
// OCI /actions/v1/responses,响应原样透传。
func (h *aiGatewayHandler) responsesPassthrough(c *gin.Context, raw []byte, req aiwire.RespRequest) {
if err := service.RespPassthroughValidate(req); err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
body, err := service.RespPassthroughBody(raw)
if err != nil {
aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error()) aiError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return return
} }
if req.Stream { if req.Stream {
h.streamResponses(c, ir) h.responsesPassthroughStream(c, body, req)
return return
} }
start := time.Now() start := time.Now()
resp, meta, err := h.gw.Chat(c.Request.Context(), ir, keyGroup(c)) payload, meta, err := h.gw.RespPassthrough(c.Request.Context(), body, req.Model, keyGroup(c))
entry := h.logEntry(c, "responses", ir.Model, false, meta, start) entry := h.logEntry(c, "responses", req.Model, false, meta, start)
if err != nil { if err != nil {
upstreamError(c, err) upstreamError(c, err)
entry.ErrMsg = err.Error() entry.ErrMsg = err.Error()
@@ -437,54 +598,70 @@ func (h *aiGatewayHandler) responses(c *gin.Context) {
return return
} }
entry.Status = http.StatusOK entry.Status = http.StatusOK
fillUsage(&entry, resp.Usage) fillUsage(&entry, service.RespPassthroughUsage(payload))
callID := h.gw.LogCall(entry) callID := h.gw.LogCall(entry)
out := service.IRRespToResponses(resp, aiRandID("resp_"), time.Now().Unix()) h.maybeLogContent(c, callID, "responses", req.Model, false, req, json.RawMessage(payload))
h.maybeLogContent(c, callID, "responses", ir.Model, false, req, out) c.Data(http.StatusOK, "application/json; charset=utf-8", payload)
c.JSON(http.StatusOK, out)
} }
// streamResponses 经状态机把 IR chunk 流聚合为 Responses 语义事件流。 // responsesPassthroughStream 流式直通:SSE 事件原样转发(推理增量等直达客户端),
func (h *aiGatewayHandler) streamResponses(c *gin.Context, ir aiwire.ChatRequest) { // 逐行扫描 completed 事件提取 usage 记账。
func (h *aiGatewayHandler) responsesPassthroughStream(c *gin.Context, body []byte, req aiwire.RespRequest) {
start := time.Now() start := time.Now()
stream, meta, err := h.gw.OpenStream(c.Request.Context(), ir, keyGroup(c)) upstream, meta, err := h.gw.RespPassthroughStream(c.Request.Context(), body, req.Model, keyGroup(c))
entry := h.logEntry(c, "responses", ir.Model, true, meta, start) entry := h.logEntry(c, "responses", req.Model, true, meta, start)
if err != nil { if err != nil {
upstreamError(c, err) upstreamError(c, err)
entry.ErrMsg = err.Error() entry.ErrMsg = err.Error()
h.logFailure(c, entry, ir) h.logFailure(c, entry, req)
return return
} }
defer stream.Close() defer upstream.Close()
sseHeaders(c) sseHeaders(c)
st := service.NewRespStream(aiRandID("resp_"), ir.Model, time.Now().Unix()) usage, upErr, err := forwardSSE(c, upstream)
for {
chunk, err := stream.Next()
if err != nil { if err != nil {
if !errors.Is(err, io.EOF) {
entry.ErrMsg = err.Error() entry.ErrMsg = err.Error()
} else if upErr != "" {
// 上游错误事件已原样转发给客户端,这里落日志可查
entry.ErrMsg = upErr
} else if usage == nil {
entry.ErrMsg = "上游流提前终止,未返回终态事件"
} }
break
}
writeRespEvents(c, st.Feed(chunk))
}
writeRespEvents(c, st.Finish())
entry.Status = http.StatusOK entry.Status = http.StatusOK
entry.LatencyMs = time.Since(start).Milliseconds() entry.LatencyMs = time.Since(start).Milliseconds()
fillUsage(&entry, st.Usage()) fillUsage(&entry, usage)
callID := h.gw.LogCall(entry) callID := h.gw.LogCall(entry)
h.maybeLogContent(c, callID, "responses", ir.Model, true, ir, nil) h.maybeLogContent(c, callID, "responses", req.Model, true, req, nil)
} }
func writeRespEvents(c *gin.Context, events []service.RespEvent) { // forwardSSE 把上游 SSE 逐行转发给客户端,空行(事件边界)即 flush;
for _, ev := range events { // 顺带从 data 行提取 response.completed 的 usage 与错误事件消息。
b, err := json.Marshal(ev.Data) func forwardSSE(c *gin.Context, upstream io.Reader) (*aiwire.Usage, string, error) {
if err != nil { reader := bufio.NewReader(upstream)
continue var usage *aiwire.Usage
} var upErr string
c.Writer.WriteString("event: " + ev.Event + "\ndata: ") for {
c.Writer.Write(b) line, err := reader.ReadBytes('\n')
c.Writer.WriteString("\n\n") if len(line) > 0 {
} c.Writer.Write(line)
trimmed := bytes.TrimSpace(line)
if len(trimmed) == 0 {
c.Writer.Flush() c.Writer.Flush()
} else if data, ok := bytes.CutPrefix(trimmed, []byte("data: ")); ok {
if u := service.RespStreamCompletedUsage(data); u != nil {
usage = u
}
if m := service.RespStreamErrorMsg(data); m != "" && upErr == "" {
upErr = m
}
}
}
if err != nil {
c.Writer.Flush()
if errors.Is(err, io.EOF) {
return usage, upErr, nil
}
return usage, upErr, err
}
}
} }
+218
View File
@@ -0,0 +1,218 @@
package api
import (
"context"
"encoding/json"
"net/http"
"strconv"
"strings"
"testing"
"time"
"gorm.io/gorm"
"oci-portal/internal/model"
)
// seedGatewayModel 直插启用渠道与模型缓存,绕过云同步。
func seedGatewayModel(t *testing.T, db *gorm.DB, name string) {
t.Helper()
ch := &model.AiChannel{Name: "t-" + name, OciConfigID: 1, Region: "r-" + name, Enabled: true, Priority: 1, Weight: 1}
if err := db.Create(ch).Error; err != nil {
t.Fatalf("seed channel: %v", err)
}
cache := &model.AiModelCache{ChannelID: ch.ID, ModelOcid: "ocid1.." + name, Name: name, Vendor: "v", SyncedAt: time.Now()}
if err := db.Create(cache).Error; err != nil {
t.Fatalf("seed model cache: %v", err)
}
}
func TestAiGatewayKeyModelRestrict(t *testing.T) {
r, auth, _, db := newTestRouterDB(t)
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
if err != nil {
t.Fatalf("login: %v", err)
}
// 受限密钥:白名单含脏输入,入库应规范化为 2 项(cohere + ghost)
w := doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token,
`{"name":"limited","value":"limited-key-1234","models":[" cohere.command-a-03-2025 ","cohere.command-a-03-2025","","ghost-model"]}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建受限密钥 status = %d, body %s", w.Code, w.Body.String())
}
var created struct {
Item model.AiKey `json:"item"`
}
if err := json.Unmarshal(w.Body.Bytes(), &created); err != nil {
t.Fatalf("decode: %v", err)
}
if len(created.Item.Models) != 2 || created.Item.Models[0] != "cohere.command-a-03-2025" {
t.Fatalf("创建后 models 未规范化: %v", created.Item.Models)
}
// 不限密钥对照
w = doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token, `{"name":"open","value":"open-key-12345"}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建不限密钥 status = %d, body %s", w.Code, w.Body.String())
}
seedGatewayModel(t, db, "cohere.command-a-03-2025")
seedGatewayModel(t, db, "meta.llama-3.3-70b-instruct")
deny := []string{"model_not_found", "无权访问"}
pass := []string{"未知模型"} // 穿过白名单闸门,由编排层以「未知模型」拒绝(池内无 ghost-model)
tests := []struct {
name string
key string
path string
body string
wantCode int
wantSub []string
}{
{"responses 流式同样拦截", "limited-key-1234", "/ai/v1/responses",
`{"model":"meta.llama-3.3-70b-instruct","stream":true,"input":"hi"}`, 404, deny},
{"responses 白名单内穿透", "limited-key-1234", "/ai/v1/responses",
`{"model":"ghost-model","input":"hi"}`, 404, pass},
{"embeddings 白名单外拦截", "limited-key-1234", "/ai/v1/embeddings",
`{"model":"meta.llama-3.3-70b-instruct","input":["hi"]}`, 404, deny},
{"messages 白名单外拦截为 Anthropic 体", "limited-key-1234", "/ai/v1/messages",
`{"model":"meta.llama-3.3-70b-instruct","max_tokens":16,"messages":[{"role":"user","content":"hi"}]}`, 404,
[]string{`"type":"error"`, "model_not_found", "无权访问"}},
{"responses 白名单外拦截", "limited-key-1234", "/ai/v1/responses",
`{"model":"meta.llama-3.3-70b-instruct","input":"hi"}`, 404, deny},
{"不限密钥穿透闸门", "open-key-12345", "/ai/v1/responses",
`{"model":"ghost-model","input":"hi"}`, 404, pass},
{"chat completions 白名单外拦截", "limited-key-1234", "/ai/v1/chat/completions",
`{"model":"meta.llama-3.3-70b-instruct","messages":[{"role":"user","content":"hi"}]}`, 404, deny},
{"chat completions 白名单内穿透", "limited-key-1234", "/ai/v1/chat/completions",
`{"model":"ghost-model","messages":[{"role":"user","content":"hi"}]}`, 404, pass},
{"messages 缺 max_tokens 按默认值放行", "open-key-12345", "/ai/v1/messages",
`{"model":"ghost-model","messages":[{"role":"user","content":"hi"}]}`, 404, pass},
{"chat completions 缺 messages 拒绝", "open-key-12345", "/ai/v1/chat/completions",
`{"model":"ghost-model"}`, 400, []string{"invalid_request_error"}},
{"chat completions 非 function 工具拒绝", "open-key-12345", "/ai/v1/chat/completions",
`{"model":"ghost-model","messages":[{"role":"user","content":"hi"}],"tools":[{"type":"web_search","function":{}}]}`,
400, []string{"仅支持 function"}},
{"chat completions 不支持内容块拒绝", "open-key-12345", "/ai/v1/chat/completions",
`{"model":"ghost-model","messages":[{"role":"user","content":[{"type":"input_audio"}]}]}`,
400, []string{"invalid_request_error"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := doRequest(t, r, http.MethodPost, tt.path, tt.key, tt.body)
if w.Code != tt.wantCode {
t.Fatalf("status = %d, want %d, body %s", w.Code, tt.wantCode, w.Body.String())
}
for _, sub := range tt.wantSub {
if !strings.Contains(w.Body.String(), sub) {
t.Errorf("body 缺少 %q: %s", sub, w.Body.String())
}
}
})
}
}
func TestAiGatewayModelsFilteredByKey(t *testing.T) {
r, auth, _, db := newTestRouterDB(t)
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
if err != nil {
t.Fatalf("login: %v", err)
}
seedGatewayModel(t, db, "cohere.command-a-03-2025")
seedGatewayModel(t, db, "meta.llama-3.3-70b-instruct")
w := doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token,
`{"name":"limited","value":"limited-key-1234","models":["cohere.command-a-03-2025"]}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建密钥 status = %d", w.Code)
}
w = doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token, `{"name":"open","value":"open-key-12345"}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建密钥 status = %d", w.Code)
}
tests := []struct {
name string
key string
want []string
notWant []string
}{
{"受限密钥仅见交集", "limited-key-1234", []string{"cohere.command-a-03-2025"}, []string{"meta.llama-3.3-70b-instruct"}},
{"不限密钥全量可见", "open-key-12345", []string{"cohere.command-a-03-2025", "meta.llama-3.3-70b-instruct"}, nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := doRequest(t, r, http.MethodGet, "/ai/v1/models", tt.key, "")
if w.Code != http.StatusOK {
t.Fatalf("status = %d, body %s", w.Code, w.Body.String())
}
for _, sub := range tt.want {
if !strings.Contains(w.Body.String(), sub) {
t.Errorf("应包含 %q: %s", sub, w.Body.String())
}
}
for _, sub := range tt.notWant {
if strings.Contains(w.Body.String(), sub) {
t.Errorf("不应包含 %q: %s", sub, w.Body.String())
}
}
})
}
}
func TestAiKeyModelsUpdateRoundTrip(t *testing.T) {
r, auth, _, _ := newTestRouterDB(t)
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
if err != nil {
t.Fatalf("login: %v", err)
}
w := doRequest(t, r, http.MethodPost, "/api/v1/ai-keys", token, `{"name":"k","value":"round-key-1234"}`)
if w.Code != http.StatusCreated {
t.Fatalf("创建 status = %d", w.Code)
}
var created struct {
Item model.AiKey `json:"item"`
}
_ = json.Unmarshal(w.Body.Bytes(), &created)
id := created.Item.ID
listModels := func() []string {
w := doRequest(t, r, http.MethodGet, "/api/v1/ai-keys", token, "")
var resp struct {
Items []model.AiKey `json:"items"`
}
_ = json.Unmarshal(w.Body.Bytes(), &resp)
for _, k := range resp.Items {
if k.ID == id {
return k.Models
}
}
t.Fatalf("密钥 %d 不在列表中", id)
return nil
}
tests := []struct {
name string
body string
want []string
}{
{"设置白名单", `{"models":[" a-model ","a-model","b-model"]}`, []string{"a-model", "b-model"}},
{"不传 models 保持不变", `{"name":"k2"}`, []string{"a-model", "b-model"}},
{"空数组清空恢复不限", `{"models":[]}`, nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := doRequest(t, r, http.MethodPut, "/api/v1/ai-keys/"+strconv.Itoa(int(id)), token, tt.body)
if w.Code != http.StatusNoContent {
t.Fatalf("update status = %d, body %s", w.Code, w.Body.String())
}
got := listModels()
if len(got) != len(tt.want) {
t.Fatalf("models = %v, want %v", got, tt.want)
}
for i := range tt.want {
if got[i] != tt.want[i] {
t.Fatalf("models = %v, want %v", got, tt.want)
}
}
})
}
}
-122
View File
@@ -7,7 +7,6 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"oci-portal/internal/model"
"oci-portal/internal/service" "oci-portal/internal/service"
) )
@@ -107,127 +106,6 @@ func (h *logEventHandler) list(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"items": items, "total": total}) c.JSON(http.StatusOK, gin.H{"items": items, "total": total})
} }
// alertRuleRequest 是创建/更新告警规则的请求体。
type alertRuleRequest struct {
Name string `json:"name"`
Enabled bool `json:"enabled"`
OciConfigID uint `json:"ociConfigId"`
EventTypes string `json:"eventTypes"`
SourceIPs string `json:"sourceIps"`
SourceIPMode string `json:"sourceIpMode"`
ResourceMatch string `json:"resourceMatch"`
Threshold int `json:"threshold"`
WindowMinutes int `json:"windowMinutes"`
}
// toModel 转为规则模型(默认阈值 1)。
func (r alertRuleRequest) toModel() model.AlertRule {
if r.Threshold == 0 {
r.Threshold = 1
}
return model.AlertRule{
Name: r.Name, Enabled: r.Enabled, OciConfigID: r.OciConfigID,
EventTypes: r.EventTypes, SourceIPs: r.SourceIPs, SourceIPMode: r.SourceIPMode,
ResourceMatch: r.ResourceMatch, Threshold: r.Threshold, WindowMinutes: r.WindowMinutes,
}
}
// listAlertRules 返回全部告警规则。
//
// @Summary 告警规则列表
// @Tags 任务与日志回传
// @Success 200 {object} map[string]any
// @Security BearerAuth
// @Router /api/v1/log-events/alert-rules [get]
func (h *logEventHandler) listAlertRules(c *gin.Context) {
rules, err := h.svc.ListAlertRules(c.Request.Context())
if err != nil {
respondError(c, err)
return
}
c.JSON(http.StatusOK, gin.H{"items": rules})
}
// createAlertRule 创建告警规则;字段非法返回 400。
//
// @Summary 创建告警规则
// @Tags 任务与日志回传
// @Param body body alertRuleRequest true "请求体"
// @Success 201 {object} map[string]any
// @Security BearerAuth
// @Router /api/v1/log-events/alert-rules [post]
func (h *logEventHandler) createAlertRule(c *gin.Context) {
var req alertRuleRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
rule, err := h.svc.CreateAlertRule(c.Request.Context(), req.toModel())
if err != nil {
respondAlertRuleError(c, err)
return
}
c.JSON(http.StatusCreated, rule)
}
// updateAlertRule 整体覆盖告警规则(含启停)。
//
// @Summary 更新告警规则
// @Tags 任务与日志回传
// @Param ruleId path int true "规则 ID"
// @Param body body alertRuleRequest true "请求体"
// @Success 200 {object} map[string]any
// @Security BearerAuth
// @Router /api/v1/log-events/alert-rules/{ruleId} [put]
func (h *logEventHandler) updateAlertRule(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("ruleId"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "规则 ID 非法"})
return
}
var req alertRuleRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
rule, err := h.svc.UpdateAlertRule(c.Request.Context(), uint(id), req.toModel())
if err != nil {
respondAlertRuleError(c, err)
return
}
c.JSON(http.StatusOK, rule)
}
// deleteAlertRule 删除告警规则及其命中记录。
//
// @Summary 删除告警规则
// @Tags 任务与日志回传
// @Param ruleId path int true "规则 ID"
// @Success 204 "无内容"
// @Security BearerAuth
// @Router /api/v1/log-events/alert-rules/{ruleId} [delete]
func (h *logEventHandler) deleteAlertRule(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("ruleId"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "规则 ID 非法"})
return
}
if err := h.svc.DeleteAlertRule(c.Request.Context(), uint(id)); err != nil {
respondError(c, err)
return
}
c.Status(http.StatusNoContent)
}
// respondAlertRuleError 把字段校验错误映射为 400。
func respondAlertRuleError(c *gin.Context, err error) {
if errors.Is(err, service.ErrInvalidAlertRule) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
respondError(c, err)
}
// getRelay 查询 OCI 侧链路状态与关键事件清单。 // getRelay 查询 OCI 侧链路状态与关键事件清单。
// //
// @Summary 查询 OCI 侧链路状态与关键事件清单 // @Summary 查询 OCI 侧链路状态与关键事件清单
+29
View File
@@ -144,6 +144,35 @@ func (h *proxyHandler) probe(c *gin.Context) {
c.JSON(http.StatusOK, view) c.JSON(http.StatusOK, view)
} }
// setTenants 设置代理关联的租户全集(代理侧批量关联 / 解除)。
//
// @Summary 设置代理关联的租户全集
// @Tags 设置
// @Param id path int true "代理 ID"
// @Param body body object true "请求体 {\"ociConfigIds\": [1,2]}"
// @Success 200 {object} map[string]any "usedBy"
// @Security BearerAuth
// @Router /api/v1/proxies/{id}/tenants [put]
func (h *proxyHandler) setTenants(c *gin.Context) {
id, ok := pathID(c)
if !ok {
return
}
var req struct {
OciConfigIDs []uint `json:"ociConfigIds"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
used, err := h.svc.SetTenants(c.Request.Context(), id, req.OciConfigIDs)
if err != nil {
h.respond(c, err)
return
}
c.JSON(http.StatusOK, gin.H{"usedBy": used})
}
// respond 把代理业务错误映射到语义状态码。 // respond 把代理业务错误映射到语义状态码。
func (h *proxyHandler) respond(c *gin.Context, err error) { func (h *proxyHandler) respond(c *gin.Context, err error) {
switch { switch {
+8 -1
View File
@@ -57,6 +57,13 @@ func (nullClient) GetSubscription(context.Context, oci.Credentials, string) (oci
} }
func newTestRouter(t *testing.T) (*gin.Engine, *service.AuthService, *service.SystemLogService) { func newTestRouter(t *testing.T) (*gin.Engine, *service.AuthService, *service.SystemLogService) {
t.Helper()
r, auth, systemLogs, _ := newTestRouterDB(t)
return r, auth, systemLogs
}
// newTestRouterDB 额外回传 db 句柄,供需要直插数据(如网关模型缓存)的测试使用。
func newTestRouterDB(t *testing.T) (*gin.Engine, *service.AuthService, *service.SystemLogService, *gorm.DB) {
t.Helper() t.Helper()
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
@@ -90,7 +97,7 @@ func newTestRouter(t *testing.T) (*gin.Engine, *service.AuthService, *service.Sy
logEvents := service.NewLogEventService(db) logEvents := service.NewLogEventService(db)
oauth := service.NewOAuthService(db, settings, auth) oauth := service.NewOAuthService(db, settings, auth)
r := NewRouter(auth, oauth, ociConfigs, tasks, service.NewConsoleService(ociConfigs), settings, notifier, systemLogs, logEvents, service.NewProxyService(db, cipher), service.NewAiGatewayService(db, ociConfigs, nullClient{})) r := NewRouter(auth, oauth, ociConfigs, tasks, service.NewConsoleService(ociConfigs), settings, notifier, systemLogs, logEvents, service.NewProxyService(db, cipher), service.NewAiGatewayService(db, ociConfigs, nullClient{}))
return r, auth, systemLogs return r, auth, systemLogs, db
} }
func doRequest(t *testing.T, r *gin.Engine, method, path, token, body string) *httptest.ResponseRecorder { func doRequest(t *testing.T, r *gin.Engine, method, path, token, body string) *httptest.ResponseRecorder {
+3
View File
@@ -33,6 +33,9 @@ func registerAiAdmin(secured *gin.RouterGroup, aiGateway *service.AiGatewayServi
secured.POST("/ai-channels/:id/probe", aiadmin.probeChannel) secured.POST("/ai-channels/:id/probe", aiadmin.probeChannel)
secured.POST("/ai-channels/:id/sync-models", aiadmin.syncChannelModels) secured.POST("/ai-channels/:id/sync-models", aiadmin.syncChannelModels)
secured.GET("/ai-models", aiadmin.gatewayModels) secured.GET("/ai-models", aiadmin.gatewayModels)
secured.GET("/ai-blacklist", aiadmin.listBlacklist)
secured.POST("/ai-blacklist", aiadmin.addBlacklist)
secured.DELETE("/ai-blacklist/:id", aiadmin.removeBlacklist)
secured.GET("/ai-logs", aiadmin.listLogs) secured.GET("/ai-logs", aiadmin.listLogs)
secured.GET("/ai-content-logs", aiadmin.listContentLogs) secured.GET("/ai-content-logs", aiadmin.listContentLogs)
} }
+1
View File
@@ -35,4 +35,5 @@ func registerSettings(secured *gin.RouterGroup, settings *service.SettingService
secured.PUT("/proxies/:id", px.update) secured.PUT("/proxies/:id", px.update)
secured.DELETE("/proxies/:id", px.remove) secured.DELETE("/proxies/:id", px.remove)
secured.POST("/proxies/:id/probe", px.probe) secured.POST("/proxies/:id/probe", px.probe)
secured.PUT("/proxies/:id/tenants", px.setTenants)
} }
-4
View File
@@ -19,10 +19,6 @@ func registerTasksAndLogs(secured *gin.RouterGroup, tasks *service.TaskService,
le := &logEventHandler{svc: logEvents} le := &logEventHandler{svc: logEvents}
secured.GET("/log-events", le.list) secured.GET("/log-events", le.list)
secured.GET("/log-events/alert-rules", le.listAlertRules)
secured.POST("/log-events/alert-rules", le.createAlertRule)
secured.PUT("/log-events/alert-rules/:ruleId", le.updateAlertRule)
secured.DELETE("/log-events/alert-rules/:ruleId", le.deleteAlertRule)
secured.GET("/oci-configs/:id/log-webhook", le.getWebhook) secured.GET("/oci-configs/:id/log-webhook", le.getWebhook)
secured.POST("/oci-configs/:id/log-webhook", le.ensureWebhook) secured.POST("/oci-configs/:id/log-webhook", le.ensureWebhook)
secured.DELETE("/oci-configs/:id/log-webhook", le.revokeWebhook) secured.DELETE("/oci-configs/:id/log-webhook", le.revokeWebhook)
+10 -5
View File
@@ -2,6 +2,7 @@ package api
import ( import (
"encoding/json" "encoding/json"
"errors"
"net/http" "net/http"
"strconv" "strconv"
@@ -155,10 +156,11 @@ func (h *taskHandler) logs(c *gin.Context) {
c.JSON(http.StatusOK, logs) c.JSON(http.StatusOK, logs)
} }
// @Summary 立即执行任务 // @Summary 立即执行任务(异步触发,结果经任务日志轮询获取)
// @Tags 任务与日志回传 // @Tags 任务与日志回传
// @Param id path int true "配置 ID" // @Param id path int true "配置 ID"
// @Success 200 {object} map[string]any // @Success 202 {object} map[string]any
// @Failure 409 {object} map[string]any "任务正在执行中"
// @Security BearerAuth // @Security BearerAuth
// @Router /api/v1/tasks/{id}/run [post] // @Router /api/v1/tasks/{id}/run [post]
func (h *taskHandler) run(c *gin.Context) { func (h *taskHandler) run(c *gin.Context) {
@@ -166,10 +168,13 @@ func (h *taskHandler) run(c *gin.Context) {
if !ok { if !ok {
return return
} }
entry, err := h.svc.RunTaskNow(c.Request.Context(), id) if err := h.svc.TriggerTask(c.Request.Context(), id); err != nil {
if err != nil { if errors.Is(err, service.ErrTaskRunning) {
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
return
}
respondError(c, err) respondError(c, err)
return return
} }
c.JSON(http.StatusOK, entry) c.JSON(http.StatusAccepted, gin.H{"triggered": true})
} }
+2 -2
View File
@@ -67,8 +67,8 @@ func autoMigrate(db *gorm.DB) error {
&model.User{}, &model.UserIdentity{}, &model.OciConfig{}, &model.Task{}, &model.TaskLog{}, &model.User{}, &model.UserIdentity{}, &model.OciConfig{}, &model.Task{}, &model.TaskLog{},
&model.CheckSnapshot{}, &model.CostSnapshot{}, &model.CheckSnapshot{}, &model.CostSnapshot{},
&model.RegionCache{}, &model.CompartmentCache{}, &model.Setting{}, &model.RegionCache{}, &model.CompartmentCache{}, &model.Setting{},
&model.SystemLog{}, &model.LogEvent{}, &model.AlertRule{}, &model.AlertRuleHit{}, &model.Proxy{}, &model.SystemLog{}, &model.LogEvent{}, &model.Proxy{},
&model.AiKey{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}, &model.AiKey{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiModelBlacklist{}, &model.AiCallLog{},
&model.AiContentLog{}, &model.AiContentLog{},
) )
} }
+10 -25
View File
@@ -211,31 +211,6 @@ type LogEvent struct {
ReceivedAt time.Time `json:"receivedAt"` ReceivedAt time.Time `json:"receivedAt"`
} }
// AlertRule 是回传事件的自定义告警规则;条件间 AND 关系,空条件视为任意。
type AlertRule struct {
ID uint `gorm:"primaryKey" json:"id"`
Name string `gorm:"size:64" json:"name"`
Enabled bool `json:"enabled"`
OciConfigID uint `json:"ociConfigId"` // 0=全部租户
EventTypes string `gorm:"size:512" json:"eventTypes"` // 逗号分隔事件短名,空=全部
SourceIPs string `gorm:"size:512" json:"sourceIps"` // 逗号分隔 IP/CIDR,空=任意
// in:来源命中列表才告警(默认);notin:不在列表才告警(白名单场景)
SourceIPMode string `gorm:"size:8" json:"sourceIpMode"`
ResourceMatch string `gorm:"size:128" json:"resourceMatch"` // 资源名子串,空=任意
Threshold int `json:"threshold"` // 触发阈值,默认 1(即时)
WindowMinutes int `json:"windowMinutes"` // 聚合窗口,Threshold>1 时必填
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
}
// AlertRuleHit 记录规则命中,供阈值窗口计数;随周期清理删除过期行。
type AlertRuleHit struct {
ID uint `gorm:"primaryKey" json:"id"`
RuleID uint `gorm:"index" json:"ruleId"`
LogEventID uint `json:"logEventId"`
HitAt time.Time `gorm:"index" json:"hitAt"`
}
// RegionCache 是开启多区域支持的配置缓存的订阅区域(每配置多行,整组覆盖)。 // RegionCache 是开启多区域支持的配置缓存的订阅区域(每配置多行,整组覆盖)。
// Status 非 READY(新订阅进行中)时读取接口会实时刷新,直到全部 READY。 // Status 非 READY(新订阅进行中)时读取接口会实时刷新,直到全部 READY。
type RegionCache struct { type RegionCache struct {
@@ -266,6 +241,8 @@ type AiKey struct {
KeyHash string `gorm:"uniqueIndex;size:64" json:"-"` KeyHash string `gorm:"uniqueIndex;size:64" json:"-"`
Tail string `gorm:"size:4" json:"tail"` Tail string `gorm:"size:4" json:"tail"`
Group string `gorm:"column:key_group;size:32" json:"group"` Group string `gorm:"column:key_group;size:32" json:"group"`
// Models 是模型白名单(精确匹配模型名);空 = 不限,非空时网关仅放行列表内模型
Models []string `gorm:"serializer:json;type:text" json:"models"`
Enabled bool `json:"enabled"` Enabled bool `json:"enabled"`
LastUsedAt *time.Time `json:"lastUsedAt"` LastUsedAt *time.Time `json:"lastUsedAt"`
// ContentLogUntil 是内容日志开启截止时间;nil = 永关(默认,红线), // ContentLogUntil 是内容日志开启截止时间;nil = 永关(默认,红线),
@@ -311,6 +288,14 @@ type AiModelCache struct {
RetiredAt *time.Time `json:"retiredAt"` RetiredAt *time.Time `json:"retiredAt"`
} }
// AiModelBlacklist 是全局模型黑名单(按模型名):加入即从全部渠道缓存删除该模型,
// 同步与探测过滤不再入库;适用于微调基座等实际不可按需调用的模型,由用户手动维护。
type AiModelBlacklist struct {
ID uint `gorm:"primaryKey" json:"id"`
Name string `gorm:"size:96;uniqueIndex" json:"name"`
CreatedAt time.Time `json:"createdAt"`
}
// AiCallLog 是网关调用日志:仅元数据与 token 用量,绝不记录 prompt / 响应正文。 // AiCallLog 是网关调用日志:仅元数据与 token 用量,绝不记录 prompt / 响应正文。
type AiCallLog struct { type AiCallLog struct {
ID uint `gorm:"primaryKey" json:"id"` ID uint `gorm:"primaryKey" json:"id"`
+5 -2
View File
@@ -6,6 +6,7 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"io"
"strings" "strings"
"sync" "sync"
"time" "time"
@@ -99,10 +100,12 @@ type Client interface {
AddVnicIpv6(ctx context.Context, cred Credentials, region, vnicID, address string) (string, error) AddVnicIpv6(ctx context.Context, cred Credentials, region, vnicID, address string) (string, error)
// GenAI 网关:区域模型列表、聊天(非流式/流式)、配额探测、文本向量化。 // GenAI 网关:区域模型列表、聊天(非流式/流式)、配额探测、文本向量化。
ListGenAiModels(ctx context.Context, cred Credentials, region string) ([]GenAiModel, error) ListGenAiModels(ctx context.Context, cred Credentials, region string) ([]GenAiModel, error)
GenAiChat(ctx context.Context, cred Credentials, region, modelOcid string, ir aiwire.ChatRequest) (*aiwire.ChatResponse, error)
GenAiChatStream(ctx context.Context, cred Credentials, region, modelOcid string, ir aiwire.ChatRequest) (GenAiStream, error)
GenAiProbeChat(ctx context.Context, cred Credentials, region, modelOcid, modelName string) (int, error) GenAiProbeChat(ctx context.Context, cred Credentials, region, modelOcid, modelName string) (int, error)
GenAiEmbed(ctx context.Context, cred Credentials, region, modelOcid string, inputs []string, dimensions *int) ([][]float32, *aiwire.Usage, error) GenAiEmbed(ctx context.Context, cred Credentials, region, modelOcid string, inputs []string, dimensions *int) ([][]float32, *aiwire.Usage, error)
// GenAiCompatResponses 直通 OpenAI Responses 请求体到 /actions/v1/responses(xAI 服务端工具通路)。
GenAiCompatResponses(ctx context.Context, cred Credentials, region string, body []byte) ([]byte, error)
// GenAiCompatResponsesStream 流式直通 /actions/v1/responses,建立成功返回 SSE body。
GenAiCompatResponsesStream(ctx context.Context, cred Credentials, region string, body []byte) (io.ReadCloser, error)
// 控制台连接:创建(VNC/串口连接串)、列出、删除。 // 控制台连接:创建(VNC/串口连接串)、列出、删除。
CreateConsoleConnection(ctx context.Context, cred Credentials, region, instanceID, sshPublicKey string) (ConsoleConnection, error) CreateConsoleConnection(ctx context.Context, cred Credentials, region, instanceID, sshPublicKey string) (ConsoleConnection, error)
ListConsoleConnections(ctx context.Context, cred Credentials, region, instanceID string) ([]ConsoleConnection, error) ListConsoleConnections(ctx context.Context, cred Credentials, region, instanceID string) ([]ConsoleConnection, error)
+31
View File
@@ -76,3 +76,34 @@ func ServiceStatus(err error) (int, bool) {
} }
return 0, false return 0, false
} }
// IsOnDemandUnsupported 识别「模型在该区域仅为微调基座、不支持按需调用」的 400:
// OCI 消息形如 "Not allowed to call finetune base model …, use Endpoint: false"。
// ListModels 无字段可事先区分,只能在调用报错时识别并换渠道。
func IsOnDemandUnsupported(err error) bool {
var svcErr common.ServiceError
if !errors.As(err, &svcErr) {
return false
}
return svcErr.GetHTTPStatusCode() == 400 &&
strings.Contains(strings.ToLower(svcErr.GetMessage()), "finetune base model")
}
// IsEntityNotFound 识别「实体不存在」404(消息 "Entity with key … not found"):
// GenAI 对区域内无按需供给的模型 OCID 返回此类 404,属模型级错误;
// 鉴权失败的 404 是 NotAuthorizedOrNotFound(消息为 Authorization failed…),不在此列。
func IsEntityNotFound(err error) bool {
var svcErr common.ServiceError
if !errors.As(err, &svcErr) {
return false
}
msg := strings.ToLower(svcErr.GetMessage())
return svcErr.GetHTTPStatusCode() == 404 &&
strings.Contains(msg, "entity with key") && strings.Contains(msg, "not found")
}
// IsModelUnavailable 归并「模型×区域不可按需调用」两类报错:微调基座 400 与实体不存在 404;
// 命中即应把该 (渠道, 模型) 从池中剔除并换候选/换渠道,而非定论租户配额问题。
func IsModelUnavailable(err error) bool {
return IsOnDemandUnsupported(err) || IsEntityNotFound(err)
}
+49
View File
@@ -109,3 +109,52 @@ func TestErrorHint(t *testing.T) {
}) })
} }
} }
func TestIsOnDemandUnsupported(t *testing.T) {
ft := fakeServiceError{status: 400, code: "InvalidParameter",
message: "Not allowed to call finetune base model ocid1.generativeaimodel.oc1.eu-frankfurt-1.tpel5q, use Endpoint: false"}
tests := []struct {
name string
err error
want bool
}{
{"微调基座 400 命中(含包装链)", fmt.Errorf("genai chat: %w", ft), true},
{"同消息但非 400 不命中", fakeServiceError{status: 500, code: "InternalError", message: ft.message}, false},
{"普通 400 不命中", fakeServiceError{status: 400, code: "InvalidParameter", message: "bad request"}, false},
{"非服务端错误不命中", errors.New("finetune base model"), false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := IsOnDemandUnsupported(tt.err); got != tt.want {
t.Errorf("IsOnDemandUnsupported() = %v, want %v", got, tt.want)
}
})
}
}
func TestIsModelUnavailable(t *testing.T) {
entity404 := fakeServiceError{status: 404, code: "NotFound",
message: "Entity with key ocid1.generativeaimodel.oc1.eu-frankfurt-1.2flsfq not found"}
auth404 := fakeServiceError{status: 404, code: "NotAuthorizedOrNotFound",
message: "Authorization failed or requested resource not found."}
ft400 := fakeServiceError{status: 400, code: "InvalidParameter",
message: "Not allowed to call finetune base model ocid1.generativeaimodel.oc1..x, use Endpoint: false"}
tests := []struct {
name string
err error
want bool
}{
{"实体不存在 404 命中(含包装链)", fmt.Errorf("genai chat: %w", entity404), true},
{"鉴权类 404 不命中(仍属租户级)", auth404, false},
{"微调基座 400 命中", ft400, true},
{"其他 404 消息不命中", fakeServiceError{status: 404, code: "NotFound", message: "route not found"}, false},
{"非服务端错误不命中", errors.New("entity with key x not found"), false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := IsModelUnavailable(tt.err); got != tt.want {
t.Errorf("IsModelUnavailable() = %v, want %v", got, tt.want)
}
})
}
}
+55 -128
View File
@@ -2,13 +2,11 @@ package oci
import ( import (
"context" "context"
"encoding/json"
"fmt" "fmt"
"io"
"net/http" "net/http"
"strings"
"time" "time"
"github.com/oracle/oci-go-sdk/v65/common"
"github.com/oracle/oci-go-sdk/v65/generativeai" "github.com/oracle/oci-go-sdk/v65/generativeai"
"github.com/oracle/oci-go-sdk/v65/generativeaiinference" "github.com/oracle/oci-go-sdk/v65/generativeaiinference"
@@ -20,7 +18,6 @@ type GenAiModel struct {
Ocid string `json:"ocid"` Ocid string `json:"ocid"`
Name string `json:"name"` Name string `json:"name"`
Vendor string `json:"vendor"` Vendor string `json:"vendor"`
ChatOnly bool `json:"-"`
Caps []string `json:"capabilities"` Caps []string `json:"capabilities"`
// Capability 是网关侧归一能力:CHAT / EMBEDDING(兼具时算 CHAT) // Capability 是网关侧归一能力:CHAT / EMBEDDING(兼具时算 CHAT)
Capability string `json:"capability"` Capability string `json:"capability"`
@@ -31,12 +28,6 @@ type GenAiModel struct {
Retired *time.Time `json:"retired"` Retired *time.Time `json:"retired"`
} }
// GenAiStream 逐事件读取流式聊天;Next 在流结束时返回 io.EOF。
type GenAiStream interface {
Next() (aiwire.ChatChunk, error)
Close() error
}
func (c *RealClient) genAiClient(cred Credentials, region string) (generativeai.GenerativeAiClient, error) { func (c *RealClient) genAiClient(cred Credentials, region string) (generativeai.GenerativeAiClient, error) {
gc, err := generativeai.NewGenerativeAiClientWithConfigurationProvider(provider(cred)) gc, err := generativeai.NewGenerativeAiClientWithConfigurationProvider(provider(cred))
if err != nil { if err != nil {
@@ -71,18 +62,39 @@ func (c *RealClient) ListGenAiModels(ctx context.Context, cred Credentials, regi
if err != nil { if err != nil {
return nil, fmt.Errorf("list genai models: %w", err) return nil, fmt.Errorf("list genai models: %w", err)
} }
seen := map[string]bool{} return dedupGenAiModels(resp.Items, time.Now()), nil
now := time.Now() }
// dedupGenAiModels 压平并按名称去重;同名多条目时优先保留不含 FINE_TUNE 能力的条目
// (微调基座条目在部分区域不支持按需调用,缓存其 OCID 会导致调用 400)。
func dedupGenAiModels(items []generativeai.ModelSummary, now time.Time) []GenAiModel {
seen := map[string]int{}
var out []GenAiModel var out []GenAiModel
for _, m := range resp.Items { for _, m := range items {
gm, ok := toGenAiModel(m, now) gm, ok := toGenAiModel(m, now)
if !ok || seen[gm.Name] { if !ok {
continue continue
} }
seen[gm.Name] = true if i, dup := seen[gm.Name]; dup {
if hasFineTune(out[i].Caps) && !hasFineTune(gm.Caps) {
out[i] = gm
}
continue
}
seen[gm.Name] = len(out)
out = append(out, gm) out = append(out, gm)
} }
return out, nil return out
}
// hasFineTune 判断能力列表是否含 FINE_TUNE(微调基座条目)。
func hasFineTune(caps []string) bool {
for _, c := range caps {
if c == string(generativeai.ModelCapabilityFineTune) {
return true
}
}
return false
} }
// toGenAiModel 压平模型摘要;无归一能力、无名称或按需推理已退役 // toGenAiModel 压平模型摘要;无归一能力、无名称或按需推理已退役
@@ -132,121 +144,15 @@ func capStrings(caps []generativeai.ModelCapabilityEnum) []string {
return out return out
} }
// GenAiChat 实现 Client:非流式聊天;modelOcid 走 on-demand serving, // GenAiProbeChat 实现 Client:经 OpenAI 兼容面(直通同链路)发一次极小请求探测渠道;配额探测专用的最小聊天(maxTokens=1),返回 HTTP 状态码;
// cohere.* 模型走 COHERE 请求格式,其余走 GENERIC。
func (c *RealClient) GenAiChat(ctx context.Context, cred Credentials, region, modelOcid string, ir aiwire.ChatRequest) (*aiwire.ChatResponse, error) {
ic, err := c.genAiInferenceClient(cred, region)
if err != nil {
return nil, err
}
req, err := buildChatRequest(ir, false)
if err != nil {
return nil, err
}
resp, err := ic.Chat(ctx, chatRequest(cred, modelOcid, req))
if err != nil {
return nil, fmt.Errorf("genai chat: %w", err)
}
return chatResponseToIR(resp.ChatResult.ChatResponse, ir.Model)
}
// buildChatRequest 按模型 vendor 组装底层请求体。
func buildChatRequest(ir aiwire.ChatRequest, stream bool) (generativeaiinference.BaseChatRequest, error) {
if isCohereModel(ir.Model) {
return irToCohereSDK(ir, stream)
}
return irToSDK(ir, stream), nil
}
// chatResponseToIR 按响应实际形态(GENERIC / COHERE)转回 IR。
func chatResponseToIR(resp generativeaiinference.BaseChatResponse, model string) (*aiwire.ChatResponse, error) {
switch r := resp.(type) {
case generativeaiinference.GenericChatResponse:
return sdkToIR(r, model), nil
case generativeaiinference.CohereChatResponse:
return cohereSDKToIR(r, model), nil
default:
return nil, fmt.Errorf("genai chat: unexpected response format %T", resp)
}
}
func chatRequest(cred Credentials, modelOcid string, req generativeaiinference.BaseChatRequest) generativeaiinference.ChatRequest {
return generativeaiinference.ChatRequest{
ChatDetails: generativeaiinference.ChatDetails{
CompartmentId: &cred.TenancyOCID,
ServingMode: generativeaiinference.OnDemandServingMode{ModelId: &modelOcid},
ChatRequest: req,
},
}
}
// GenAiChatStream 实现 Client:流式聊天,返回逐事件读取器。
// SDK 对 text/event-stream 跳过 unmarshal,原始流经 RawResponse 交由 SSEReader 消费。
func (c *RealClient) GenAiChatStream(ctx context.Context, cred Credentials, region, modelOcid string, ir aiwire.ChatRequest) (GenAiStream, error) {
ic, err := c.genAiInferenceClient(cred, region)
if err != nil {
return nil, err
}
req, err := buildChatRequest(ir, true)
if err != nil {
return nil, err
}
resp, err := ic.Chat(ctx, chatRequest(cred, modelOcid, req))
if err != nil {
return nil, fmt.Errorf("genai chat stream: %w", err)
}
reader, err := common.NewSSEReader(resp.RawResponse)
if err != nil {
return nil, fmt.Errorf("genai chat stream: sse reader: %w", err)
}
return &genAiSSEStream{reader: reader, body: resp.RawResponse.Body, model: ir.Model}, nil
}
// genAiSSEStream 把 OCI SSE 事件流适配为 IR chunk 流。
type genAiSSEStream struct {
reader *common.SseReader
body io.ReadCloser
model string
}
// Next 读取下一条有效事件;空事件与 [DONE] 哨兵跳过,流尽返回 io.EOF。
func (s *genAiSSEStream) Next() (aiwire.ChatChunk, error) {
for {
data, err := s.reader.ReadNextEvent()
if err != nil {
return aiwire.ChatChunk{}, err
}
text := strings.TrimSpace(string(data))
if text == "" {
continue
}
if text == "[DONE]" {
return aiwire.ChatChunk{}, io.EOF
}
if chunk, ok := parseGenAiEvent([]byte(text), s.model); ok {
return chunk, nil
}
}
}
func (s *genAiSSEStream) Close() error {
if s.body != nil {
return s.body.Close()
}
return nil
}
// GenAiProbeChat 实现 Client:配额探测专用的最小聊天(maxTokens=1),返回 HTTP 状态码;
// modelName 决定请求格式(cohere.* 走 COHERE)。 // modelName 决定请求格式(cohere.* 走 COHERE)。
func (c *RealClient) GenAiProbeChat(ctx context.Context, cred Credentials, region, modelOcid, modelName string) (int, error) { func (c *RealClient) GenAiProbeChat(ctx context.Context, cred Credentials, region, modelOcid, modelName string) (int, error) {
one := 1 body, err := json.Marshal(map[string]any{"model": modelName, "input": "hi",
ir := aiwire.ChatRequest{ "max_output_tokens": 1, "store": false})
Model: modelName, if err != nil {
Messages: []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent("hi")}}, return 0, err
MaxTokens: &one,
} }
_, err := c.GenAiChat(ctx, cred, region, modelOcid, ir) if _, err = c.GenAiCompatResponses(ctx, cred, region, body); err == nil {
if err == nil {
return http.StatusOK, nil return http.StatusOK, nil
} }
if status, ok := ServiceStatus(err); ok { if status, ok := ServiceStatus(err); ok {
@@ -255,6 +161,27 @@ func (c *RealClient) GenAiProbeChat(ctx context.Context, cred Credentials, regio
return 0, err return 0, err
} }
// sdkUsageToIR 把 SDK 用量转为内部记账结构(缓存命中挂 details,仅命中时出现)。
func sdkUsageToIR(u *generativeaiinference.Usage) *aiwire.Usage {
if u == nil {
return nil
}
out := &aiwire.Usage{}
if u.PromptTokens != nil {
out.PromptTokens = *u.PromptTokens
}
if u.CompletionTokens != nil {
out.CompletionTokens = *u.CompletionTokens
}
if u.TotalTokens != nil {
out.TotalTokens = *u.TotalTokens
}
if u.PromptTokensDetails != nil && u.PromptTokensDetails.CachedTokens != nil && *u.PromptTokensDetails.CachedTokens > 0 {
out.PromptTokensDetails = &aiwire.PromptTokensDetails{CachedTokens: *u.PromptTokensDetails.CachedTokens}
}
return out
}
// GenAiEmbed 实现 Client:文本向量化(on-demand serving);dimensions 透传 OutputDimensions。 // GenAiEmbed 实现 Client:文本向量化(on-demand serving);dimensions 透传 OutputDimensions。
func (c *RealClient) GenAiEmbed(ctx context.Context, cred Credentials, region, modelOcid string, inputs []string, dimensions *int) ([][]float32, *aiwire.Usage, error) { func (c *RealClient) GenAiEmbed(ctx context.Context, cred Credentials, region, modelOcid string, inputs []string, dimensions *int) ([][]float32, *aiwire.Usage, error) {
ic, err := c.genAiInferenceClient(cred, region) ic, err := c.genAiInferenceClient(cred, region)
-266
View File
@@ -1,266 +0,0 @@
package oci
import (
"encoding/json"
"fmt"
"strings"
"github.com/oracle/oci-go-sdk/v65/generativeaiinference"
"oci-portal/internal/aiwire"
)
// isCohereModel 依据模型名前缀判定 COHERE 系(走专属请求格式)。
func isCohereModel(name string) bool {
return strings.HasPrefix(strings.ToLower(name), "cohere.")
}
// irToCohereSDK 把 IR 转为 COHERE 聊天请求;工具与结构化输出做有损降级
// (参数仅取 JSON Schema 顶层 properties,tool_choice / json_schema 的 name、strict 无对应)。
func irToCohereSDK(ir aiwire.ChatRequest, stream bool) (generativeaiinference.CohereChatRequest, error) {
var req generativeaiinference.CohereChatRequest
if err := cohereRejectUnsupported(ir); err != nil {
return req, err
}
p, err := cohereSplitMessages(ir.Messages)
if err != nil {
return req, err
}
req = generativeaiinference.CohereChatRequest{
Message: &p.message,
ChatHistory: p.history,
ToolResults: p.toolResults,
Tools: cohereTools(ir.Tools),
ResponseFormat: cohereResponseFormat(ir.ResponseFormat),
MaxTokens: ir.MaxTokens,
Temperature: ir.Temperature,
TopP: ir.TopP,
TopK: ir.TopK,
FrequencyPenalty: ir.FrequencyPenalty,
PresencePenalty: ir.PresencePenalty,
Seed: ir.Seed,
}
if p.preamble != "" {
req.PreambleOverride = &p.preamble
}
if len(ir.Stop) > 0 {
req.StopSequences = ir.Stop
}
if stream {
req.IsStream = &stream
includeUsage := ir.StreamOptions == nil || ir.StreamOptions.IncludeUsage
req.StreamOptions = &generativeaiinference.StreamOptions{IsIncludeUsage: &includeUsage}
}
return req, nil
}
// cohereRejectUnsupported 拒绝 COHERE 无法承接的能力(多模态图片)。
func cohereRejectUnsupported(ir aiwire.ChatRequest) error {
for _, m := range ir.Messages {
for _, p := range m.Content.Parts {
if p.Type == "image_url" {
return fmt.Errorf("cohere 系模型暂不支持图片输入,请改用 meta/google 等多模态模型")
}
}
}
return nil
}
// cohereParts 是 IR 消息拆装为 COHERE 请求的中间结果。
type cohereParts struct {
message string
history []generativeaiinference.CohereMessage
preamble string
toolResults []generativeaiinference.CohereToolResult
}
// cohereSplitMessages 拆装消息:system 拼 preamble,末尾 user 抽为 message,
// assistant(含 toolCalls)与其余 user 进 history,tool 结果转顶层 toolResults
// (工具结果在末尾时 message 留空,COHERE 以 toolResults 续跑)。
func cohereSplitMessages(msgs []aiwire.ChatMessage) (cohereParts, error) {
lastUser, lastTool := -1, -1
for i, m := range msgs {
switch m.Role {
case "user":
lastUser = i
case "tool":
lastTool = i
}
}
if lastUser == -1 && lastTool == -1 {
return cohereParts{}, fmt.Errorf("cohere 系模型至少需要一条 user 消息")
}
var p cohereParts
var preamble []string
calls := map[string]generativeaiinference.CohereToolCall{}
for i, m := range msgs {
text := m.Content.JoinText()
switch m.Role {
case "system", "developer":
preamble = append(preamble, text)
case "assistant":
p.history = append(p.history, cohereBotMessage(text, m.ToolCalls, calls))
case "tool":
p.toolResults = append(p.toolResults, cohereToolResult(m, calls))
default: // user
if i == lastUser && lastUser > lastTool {
p.message = text
continue
}
p.history = append(p.history, generativeaiinference.CohereUserMessage{Message: &text})
}
}
p.preamble = strings.Join(preamble, "\n")
return p, nil
}
// cohereBotMessage 转 assistant 消息;toolCalls 同步登记进 id→call 映射供 tool 结果回查。
func cohereBotMessage(text string, tcs []aiwire.ToolCall, calls map[string]generativeaiinference.CohereToolCall) generativeaiinference.CohereChatBotMessage {
msg := generativeaiinference.CohereChatBotMessage{}
if text != "" {
msg.Message = &text
}
for _, tc := range tcs {
name := tc.Function.Name
var params interface{}
if json.Unmarshal([]byte(tc.Function.Arguments), &params) != nil || params == nil {
params = map[string]any{}
}
call := generativeaiinference.CohereToolCall{Name: &name, Parameters: &params}
calls[tc.ID] = call
msg.ToolCalls = append(msg.ToolCalls, call)
}
return msg
}
// cohereToolResult 把 IR tool 消息转顶层工具结果;COHERE 工具调用无 id,
// 靠 assistant 历史登记的映射回查,缺失时以 tool_call_id 名义调用兜底。
func cohereToolResult(m aiwire.ChatMessage, calls map[string]generativeaiinference.CohereToolCall) generativeaiinference.CohereToolResult {
call, ok := calls[m.ToolCallID]
if !ok {
name := m.ToolCallID
var params interface{} = map[string]any{}
call = generativeaiinference.CohereToolCall{Name: &name, Parameters: &params}
}
text := m.Content.JoinText()
var output interface{}
if json.Unmarshal([]byte(text), &output) != nil || output == nil {
output = map[string]any{"output": text}
}
if arr, isArr := output.([]interface{}); isArr {
return generativeaiinference.CohereToolResult{Call: &call, Outputs: arr}
}
if _, isMap := output.(map[string]interface{}); !isMap {
output = map[string]any{"output": output}
}
return generativeaiinference.CohereToolResult{Call: &call, Outputs: []interface{}{output}}
}
// cohereTools 把 JSON Schema 工具定义降级为 COHERE 扁平参数表(嵌套结构有损:仅取顶层)。
func cohereTools(tools []aiwire.Tool) []generativeaiinference.CohereTool {
if len(tools) == 0 {
return nil
}
out := make([]generativeaiinference.CohereTool, 0, len(tools))
for _, t := range tools {
name, desc := t.Function.Name, t.Function.Description
if desc == "" {
desc = name // Description 为 COHERE 必填
}
out = append(out, generativeaiinference.CohereTool{
Name: &name, Description: &desc,
ParameterDefinitions: cohereParams(t.Function.Parameters),
})
}
return out
}
// cohereParams 取 JSON Schema 顶层 properties 转扁平参数定义;嵌套 schema 只保留类型名。
func cohereParams(schema json.RawMessage) map[string]generativeaiinference.CohereParameterDefinition {
var s struct {
Properties map[string]struct {
Type string `json:"type"`
Description string `json:"description"`
} `json:"properties"`
Required []string `json:"required"`
}
if json.Unmarshal(schema, &s) != nil || len(s.Properties) == 0 {
return nil
}
req := map[string]bool{}
for _, r := range s.Required {
req[r] = true
}
out := make(map[string]generativeaiinference.CohereParameterDefinition, len(s.Properties))
for k, v := range s.Properties {
typ, desc := v.Type, v.Description
if typ == "" {
typ = "string"
}
pd := generativeaiinference.CohereParameterDefinition{Type: &typ}
if desc != "" {
pd.Description = &desc
}
if req[k] {
t := true
pd.IsRequired = &t
}
out[k] = pd
}
return out
}
// cohereResponseFormat 映射 json_object / json_schema(name、strict 无对应,有损降级)。
func cohereResponseFormat(rf *aiwire.ResponseFormat) generativeaiinference.CohereResponseFormat {
if rf == nil {
return nil
}
switch rf.Type {
case "json_object", "json_schema":
out := generativeaiinference.CohereResponseJsonFormat{}
var in struct {
Schema json.RawMessage `json:"schema"`
}
if json.Unmarshal(rf.JSONSchema, &in) == nil && len(in.Schema) > 0 {
var schema interface{}
if json.Unmarshal(in.Schema, &schema) == nil {
out.Schema = &schema
}
}
return out
}
return nil
}
// cohereSDKToIR 把 COHERE 非流式响应转回 IR;工具调用补派生 ID,存在时结束原因归 tool_calls。
func cohereSDKToIR(resp generativeaiinference.CohereChatResponse, model string) *aiwire.ChatResponse {
msg := aiwire.ChatMessage{Role: "assistant", Content: aiwire.NewTextContent(deref(resp.Text))}
for i, tc := range resp.ToolCalls {
msg.ToolCalls = append(msg.ToolCalls, cohereCallToIR(deref(tc.Name), tc.Parameters, i))
}
finish := mapFinishReason(string(resp.FinishReason))
if len(msg.ToolCalls) > 0 {
finish = "tool_calls"
}
return &aiwire.ChatResponse{
Object: "chat.completion",
Model: model,
Choices: []aiwire.Choice{{Message: msg, FinishReason: finish}},
Usage: sdkUsageToIR(resp.Usage),
}
}
// cohereCallToIR 生成 IR 工具调用;COHERE 无调用 id,按名称+序号派生(回传时按此回查)。
func cohereCallToIR(name string, params *interface{}, idx int) aiwire.ToolCall {
args := "{}"
if params != nil && *params != nil {
if b, err := json.Marshal(*params); err == nil {
args = string(b)
}
}
return aiwire.ToolCall{
ID: fmt.Sprintf("call_%s_%d", name, idx),
Type: "function",
Function: aiwire.FunctionCall{Name: name, Arguments: args},
}
}
-367
View File
@@ -1,367 +0,0 @@
package oci
import (
"encoding/json"
"strings"
"github.com/oracle/oci-go-sdk/v65/generativeaiinference"
"oci-portal/internal/aiwire"
)
// irToSDK 把 IR(OpenAI 线格式)转为 OCI GENERIC 聊天请求;字段近一一镜像。
func irToSDK(ir aiwire.ChatRequest, stream bool) generativeaiinference.GenericChatRequest {
req := generativeaiinference.GenericChatRequest{
Messages: irMessages(ir.Messages),
MaxTokens: ir.MaxTokens,
MaxCompletionTokens: ir.MaxCompletionTokens,
Temperature: ir.Temperature,
TopP: ir.TopP,
TopK: ir.TopK,
FrequencyPenalty: ir.FrequencyPenalty,
PresencePenalty: ir.PresencePenalty,
Seed: ir.Seed,
NumGenerations: ir.N,
Tools: irTools(ir.Tools),
ToolChoice: irToolChoice(ir.ToolChoice),
ResponseFormat: irResponseFormat(ir.ResponseFormat),
}
if len(ir.Stop) > 0 {
req.Stop = ir.Stop
}
if ir.ReasoningEffort != "" {
req.ReasoningEffort = generativeaiinference.GenericChatRequestReasoningEffortEnum(strings.ToUpper(ir.ReasoningEffort))
}
if stream {
req.IsStream = &stream
includeUsage := ir.StreamOptions == nil || ir.StreamOptions.IncludeUsage
req.StreamOptions = &generativeaiinference.StreamOptions{IsIncludeUsage: &includeUsage}
}
return req
}
// irMessages 按 role 拆装消息;user 消息保留图文块,其余角色只取文本。
func irMessages(msgs []aiwire.ChatMessage) []generativeaiinference.Message {
out := make([]generativeaiinference.Message, 0, len(msgs))
for _, m := range msgs {
switch m.Role {
case "system", "developer":
out = append(out, generativeaiinference.SystemMessage{Content: textContents(m.Content.JoinText())})
case "assistant":
out = append(out, generativeaiinference.AssistantMessage{
Content: textContents(m.Content.JoinText()),
ToolCalls: irToolCalls(m.ToolCalls),
})
case "tool":
tm := generativeaiinference.ToolMessage{Content: textContents(m.Content.JoinText())}
if m.ToolCallID != "" {
id := m.ToolCallID
tm.ToolCallId = &id
}
out = append(out, tm)
default: // user
out = append(out, generativeaiinference.UserMessage{Content: chatContents(m.Content)})
}
}
return out
}
// chatContents 把 IR 内容转 SDK 内容块;文本直通,image_url 转 IMAGE 块(data URI / 公网地址)。
func chatContents(c aiwire.Content) []generativeaiinference.ChatContent {
if len(c.Parts) == 0 {
return textContents(c.Text)
}
var out []generativeaiinference.ChatContent
for _, p := range c.Parts {
switch {
case p.Type == "image_url" && p.ImageURL != nil:
img := generativeaiinference.ImageUrl{Url: &p.ImageURL.URL}
if d, ok := generativeaiinference.GetMappingImageUrlDetailEnum(p.ImageURL.Detail); ok {
img.Detail = d
}
out = append(out, generativeaiinference.ImageContent{ImageUrl: &img})
case p.Text != "":
t := p.Text
out = append(out, generativeaiinference.TextContent{Text: &t})
}
}
return out
}
// textContents 把纯文本包装为单元素 TextContent 数组;空文本返回 nil。
func textContents(text string) []generativeaiinference.ChatContent {
if text == "" {
return nil
}
return []generativeaiinference.ChatContent{generativeaiinference.TextContent{Text: &text}}
}
func irToolCalls(calls []aiwire.ToolCall) []generativeaiinference.ToolCall {
if len(calls) == 0 {
return nil
}
out := make([]generativeaiinference.ToolCall, 0, len(calls))
for _, c := range calls {
id, name, args := c.ID, c.Function.Name, c.Function.Arguments
out = append(out, generativeaiinference.FunctionCall{Id: &id, Name: &name, Arguments: &args})
}
return out
}
func irTools(tools []aiwire.Tool) []generativeaiinference.ToolDefinition {
if len(tools) == 0 {
return nil
}
out := make([]generativeaiinference.ToolDefinition, 0, len(tools))
for _, t := range tools {
name, desc := t.Function.Name, t.Function.Description
fd := generativeaiinference.FunctionDefinition{Name: &name}
if desc != "" {
fd.Description = &desc
}
if len(t.Function.Parameters) > 0 {
var params interface{}
if json.Unmarshal(t.Function.Parameters, &params) == nil {
fd.Parameters = &params
}
}
out = append(out, fd)
}
return out
}
// irToolChoice 解析 "auto"/"none"/"required" 或 {type:function,function:{name}}。
func irToolChoice(raw json.RawMessage) generativeaiinference.ToolChoice {
if len(raw) == 0 {
return nil
}
var s string
if json.Unmarshal(raw, &s) == nil {
switch s {
case "none":
return generativeaiinference.ToolChoiceNone{}
case "required":
return generativeaiinference.ToolChoiceRequired{}
case "auto":
return generativeaiinference.ToolChoiceAuto{}
}
return nil
}
var obj struct {
Function struct {
Name string `json:"name"`
} `json:"function"`
}
if json.Unmarshal(raw, &obj) == nil && obj.Function.Name != "" {
return generativeaiinference.ToolChoiceFunction{Name: &obj.Function.Name}
}
return nil
}
func irResponseFormat(rf *aiwire.ResponseFormat) generativeaiinference.ResponseFormat {
if rf == nil {
return nil
}
switch rf.Type {
case "json_object":
return generativeaiinference.JsonObjectResponseFormat{}
case "json_schema":
return irJSONSchemaFormat(rf.JSONSchema)
default:
return nil
}
}
// irJSONSchemaFormat 映射 OpenAI json_schema{name,description,schema,strict}。
func irJSONSchemaFormat(raw json.RawMessage) generativeaiinference.ResponseFormat {
var in struct {
Name string `json:"name"`
Description string `json:"description"`
Schema json.RawMessage `json:"schema"`
Strict *bool `json:"strict"`
}
if json.Unmarshal(raw, &in) != nil || in.Name == "" {
return generativeaiinference.JsonObjectResponseFormat{}
}
js := generativeaiinference.ResponseJsonSchema{Name: &in.Name, IsStrict: in.Strict}
if in.Description != "" {
js.Description = &in.Description
}
if len(in.Schema) > 0 {
var schema interface{}
if json.Unmarshal(in.Schema, &schema) == nil {
js.Schema = &schema
}
}
return generativeaiinference.JsonSchemaResponseFormat{JsonSchema: &js}
}
// sdkToIR 把非流式 GenericChatResponse 转回 IR 响应;id/created 由调用方补。
func sdkToIR(resp generativeaiinference.GenericChatResponse, model string) *aiwire.ChatResponse {
out := &aiwire.ChatResponse{Object: "chat.completion", Model: model}
if resp.TimeCreated != nil {
out.Created = resp.TimeCreated.Unix()
}
for i, ch := range resp.Choices {
idx := i
if ch.Index != nil {
idx = *ch.Index
}
out.Choices = append(out.Choices, aiwire.Choice{
Index: idx,
Message: sdkMessageToIR(ch.Message),
FinishReason: mapFinishReason(deref(ch.FinishReason)),
})
}
out.Usage = sdkUsageToIR(resp.Usage)
return out
}
func sdkUsageToIR(u *generativeaiinference.Usage) *aiwire.Usage {
if u == nil {
return nil
}
out := &aiwire.Usage{}
if u.PromptTokens != nil {
out.PromptTokens = *u.PromptTokens
}
if u.CompletionTokens != nil {
out.CompletionTokens = *u.CompletionTokens
}
if u.TotalTokens != nil {
out.TotalTokens = *u.TotalTokens
}
if d := u.PromptTokensDetails; d != nil && d.CachedTokens != nil && *d.CachedTokens > 0 {
out.PromptTokensDetails = &aiwire.PromptTokensDetails{CachedTokens: *d.CachedTokens}
}
return out
}
// sdkMessageToIR 提取助手消息文本与工具调用。
func sdkMessageToIR(m generativeaiinference.Message) aiwire.ChatMessage {
out := aiwire.ChatMessage{Role: "assistant"}
am, ok := m.(generativeaiinference.AssistantMessage)
if !ok {
return out
}
var sb strings.Builder
for _, c := range am.Content {
if tc, ok := c.(generativeaiinference.TextContent); ok && tc.Text != nil {
sb.WriteString(*tc.Text)
}
}
out.Content = aiwire.NewTextContent(sb.String())
for _, call := range am.ToolCalls {
if fc, ok := call.(generativeaiinference.FunctionCall); ok {
out.ToolCalls = append(out.ToolCalls, aiwire.ToolCall{
ID: deref(fc.Id),
Type: "function",
Function: aiwire.FunctionCall{
Name: deref(fc.Name),
Arguments: deref(fc.Arguments),
},
})
}
}
return out
}
// mapFinishReason 归一化结束原因;OCI 与 OpenAI 取值同名,做小写与别名兜底。
func mapFinishReason(reason string) string {
switch strings.ToLower(reason) {
case "", "null":
return ""
case "stop", "completed", "end_turn":
return "stop"
case "length", "max_tokens":
return "length"
case "tool_calls", "tool_call", "tool_use":
return "tool_calls"
case "complete":
return "stop"
case "error_toxic":
return "content_filter"
default:
return strings.ToLower(reason)
}
}
// genAiStreamEvent 是 GENERIC 流式事件的宽容解析结构(字段按增量 choice 形状)。
type genAiStreamEvent struct {
Index *int `json:"index"`
Message struct {
Role string `json:"role"`
Content []struct {
Type string `json:"type"`
Text string `json:"text"`
} `json:"content"`
ToolCalls []struct {
ID string `json:"id"`
Name string `json:"name"`
Arguments string `json:"arguments"`
} `json:"toolCalls"`
} `json:"message"`
FinishReason *string `json:"finishReason"`
Usage *generativeaiinference.Usage `json:"usage"`
Choices []json.RawMessage `json:"choices"` // 兼容整包 choices 形态
Text string `json:"text"` // COHERE 流事件的顶层增量文本
// CohereCalls 是 COHERE 流事件顶层的工具调用(与 GENERIC 的 message.toolCalls 不同层级)
CohereCalls []struct {
Name string `json:"name"`
Parameters interface{} `json:"parameters"`
} `json:"toolCalls"`
}
// parseGenAiEvent 把一条 SSE 事件 JSON 解析为 IR chunk;无有效负载时 ok=false。
func parseGenAiEvent(data []byte, model string) (aiwire.ChatChunk, bool) {
var ev genAiStreamEvent
if err := json.Unmarshal(data, &ev); err != nil {
return aiwire.ChatChunk{}, false
}
if len(ev.Choices) > 0 { // choices 包裹形态:取首个展开重解析
var inner genAiStreamEvent
if json.Unmarshal(ev.Choices[0], &inner) == nil {
inner.Usage = ev.Usage
ev = inner
}
}
chunk := aiwire.ChatChunk{Object: "chat.completion.chunk", Model: model}
choice := aiwire.ChunkChoice{}
if ev.Index != nil {
choice.Index = *ev.Index
}
choice.Delta.Role = strings.ToLower(ev.Message.Role)
var sb strings.Builder
for _, c := range ev.Message.Content {
sb.WriteString(c.Text)
}
if sb.Len() == 0 && ev.Text != "" { // COHERE 事件形态
sb.WriteString(ev.Text)
}
choice.Delta.Content = sb.String()
for i, tc := range ev.Message.ToolCalls {
d := aiwire.ToolCallDelta{Index: i, ID: tc.ID}
if tc.ID != "" {
d.Type = "function"
}
d.Function.Name = tc.Name
d.Function.Arguments = tc.Arguments
choice.Delta.ToolCalls = append(choice.Delta.ToolCalls, d)
}
for i, tc := range ev.CohereCalls { // COHERE 顶层工具调用:整包出现,补派生 ID
call := cohereCallToIR(tc.Name, &tc.Parameters, i)
d := aiwire.ToolCallDelta{Index: len(choice.Delta.ToolCalls) + i, ID: call.ID, Type: "function"}
d.Function.Name = call.Function.Name
d.Function.Arguments = call.Function.Arguments
choice.Delta.ToolCalls = append(choice.Delta.ToolCalls, d)
}
if r := mapFinishReason(deref(ev.FinishReason)); r != "" {
choice.FinishReason = &r
}
hasPayload := choice.Delta.Content != "" || len(choice.Delta.ToolCalls) > 0 || choice.FinishReason != nil
if hasPayload {
chunk.Choices = []aiwire.ChunkChoice{choice}
}
chunk.Usage = sdkUsageToIR(ev.Usage)
return chunk, hasPayload || chunk.Usage != nil
}
@@ -0,0 +1,496 @@
package oci
// 真实 OCI GenAI effort 探针(重建版)。
//
// 前身 genai_cache_integration_test.go 承载缓存/亲和/服务端工具等历轮调研模式,
// 调研结论均已归档(.trellis/tasks/archive/2026-07/*/research/),该文件已删除;
// 本文件重建共享基础设施,保留 effort 支持矩阵探测,并新增 responses 端点探测
// (multi-agent 模型仅允许走 Responses 面)。
//
// 运行(默认跳过,不触网):
//
// OCI_GENAI_CACHE_PROBE=effort-matrix OCI_CACHE_DB=<db 绝对路径> DATA_KEY=<主密钥> \
// OCI_CACHE_RUN_ID=<唯一标识> go test ./internal/oci/ -run TestRealGenAiEffortMatrix -v
//
// 可选:OCI_CACHE_CONFIG_ID(默认 6)、OCI_CACHE_REGION(默认 us-chicago-1)、
// OCI_EFFORT_CASES 自定义用例("model=effort" 逗号分隔;effort 后缀 "@responses"
// 表示走 /actions/v1/responses 端点,如 "xai.grok-4.20-multi-agent=low@responses")。
//
// 纪律:数据库只读打开;OCI_GO_SDK_DEBUG 必须为空;日志不落原始请求体与凭据,
// 错误消息经 OCID 遮掩后截断。
import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"os"
"regexp"
"strings"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/oracle/oci-go-sdk/v65/common"
"github.com/oracle/oci-go-sdk/v65/generativeaiinference"
"gorm.io/gorm"
"gorm.io/gorm/logger"
projectcrypto "oci-portal/internal/crypto"
"oci-portal/internal/model"
)
const (
effortProbeRegion = "us-chicago-1"
effortProbeLimit = int64(4 << 20)
)
var effortRunIDPattern = regexp.MustCompile(`^[A-Za-z0-9._-]{1,64}$`)
type effortProbeCase struct {
model string
effort string
responses bool // true 走 /actions/v1/responses,否则 typed /actions/chat
stream bool // 仅 responses 面:请求 SSE 流,观测事件类型序列
}
type effortProbeResources struct {
cred Credentials
region string
modelOCID map[string]string
inference generativeaiinference.GenerativeAiInferenceClient
}
type effortProbeResult struct {
status int
requestRef string
duration time.Duration
errorCode string
errorMessage string
completionTokens int
reasoningTokens int
answerLen int
outputTypes []string
}
// effortMatrixCases 默认批为 grok-4.3 五档基线;全模型矩阵结论见任务归档,
// 复测用 OCI_EFFORT_CASES 自定义。
func effortMatrixCases() []effortProbeCase {
if raw := os.Getenv("OCI_EFFORT_CASES"); raw != "" {
return parseEffortCases(raw)
}
return []effortProbeCase{
{model: "xai.grok-4.3", effort: "none", responses: true},
{model: "xai.grok-4.3", effort: "low", responses: true},
{model: "xai.grok-4.3", effort: "high", responses: true},
}
}
// parseEffortCases 解析 "model=effort[@responses|@responses-stream]" 逗号分隔用例。
func parseEffortCases(raw string) []effortProbeCase {
var out []effortProbeCase
for _, item := range strings.Split(raw, ",") {
parts := strings.SplitN(strings.TrimSpace(item), "=", 2)
if len(parts) != 2 || parts[0] == "" || parts[1] == "" {
continue
}
effort, streaming := strings.CutSuffix(parts[1], "@responses-stream")
if streaming {
out = append(out, effortProbeCase{model: parts[0], effort: effort, responses: true, stream: true})
continue
}
effort, _ = strings.CutSuffix(effort, "@responses") // 后缀兼容保留;所有用例均走 responses 面
out = append(out, effortProbeCase{model: parts[0], effort: effort, responses: true})
}
return out
}
func TestRealGenAiEffortMatrix(t *testing.T) {
requireEffortProbeMode(t, "effort-matrix")
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Minute)
defer cancel()
resources := newEffortProbeResources(t, ctx)
for _, tc := range effortMatrixCases() {
runEffortCase(t, ctx, resources, tc)
}
}
func requireEffortProbeMode(t *testing.T, wanted string) {
t.Helper()
if os.Getenv("OCI_GENAI_CACHE_PROBE") != wanted {
t.Skipf("set OCI_GENAI_CACHE_PROBE=%s to run", wanted)
}
if os.Getenv("OCI_GO_SDK_DEBUG") != "" {
t.Fatal("OCI_GO_SDK_DEBUG must be unset to protect signed requests")
}
if !effortRunIDPattern.MatchString(os.Getenv("OCI_CACHE_RUN_ID")) {
t.Fatal("set a unique OCI_CACHE_RUN_ID (1-64 safe characters)")
}
}
func runEffortCase(t *testing.T, ctx context.Context, resources effortProbeResources, tc effortProbeCase) {
t.Helper()
ocid, ok := resources.modelOCID[tc.model]
if !ok {
t.Logf("model=%s effort=%s resolve-failed: model not in on-demand catalog", tc.model, tc.effort)
return
}
_ = ocid
body, path, err := effortProbeBody(tc)
if err != nil {
t.Fatalf("build effort body: %v", err)
}
result := executeEffortProbe(ctx, resources, path, body)
logEffortResult(t, tc, result)
}
// effortProbeBody 构造 OpenAI Responses 直通请求体(typed chat 面已随网关剔除,
// 探针仅保留 responses 端点;历史 typed 结论见任务归档)。
func effortProbeBody(tc effortProbeCase) ([]byte, string, error) {
question := "How many positive divisors does 360 have? Answer with the number only."
body := map[string]interface{}{"model": tc.model, "input": question,
"max_output_tokens": 2048, "store": false}
if tc.effort != "" {
body["reasoning"] = map[string]string{"effort": tc.effort}
}
if tc.stream {
body["stream"] = true
}
payload, err := json.Marshal(body)
return payload, "/actions/v1/responses", err
}
// executeEffortProbe 以 IAM 签名裸 POST 指定路径;responses 面按生产实现补 compartment 头。
func executeEffortProbe(ctx context.Context, resources effortProbeResources, path string, body []byte) effortProbeResult {
client := resources.inference.BaseClient
client.Configuration.CircuitBreaker = nil
common.UpdateEndpointTemplateForOptions(&client)
common.SetMissingTemplateParams(&client)
request, err := http.NewRequestWithContext(ctx, http.MethodPost, path, strings.NewReader(string(body)))
if err != nil {
return effortProbeResult{errorCode: "BuildRequest", errorMessage: err.Error()}
}
request.Header.Set("Content-Type", "application/json")
if path != "/actions/chat" {
request.Header.Set("CompartmentId", resources.cred.TenancyOCID)
request.Header.Set("opc-compartment-id", resources.cred.TenancyOCID)
}
return effortCallProbe(ctx, client, request)
}
func effortCallProbe(ctx context.Context, client common.BaseClient, request *http.Request) effortProbeResult {
started := time.Now()
response, err := client.Call(ctx, request)
result := effortProbeResult{duration: time.Since(started)}
if response != nil {
result.status = response.StatusCode
result.requestRef = shortEffortRef(response.Header.Get("opc-request-id"))
defer response.Body.Close()
}
if err != nil {
result.errorCode, result.errorMessage = effortProbeError(err)
return result
}
payload, readErr := io.ReadAll(io.LimitReader(response.Body, effortProbeLimit))
if readErr != nil {
result.errorCode, result.errorMessage = "ReadResponse", readErr.Error()
return result
}
parseEffortPayload(&result, payload)
return result
}
func effortProbeError(err error) (string, string) {
if serviceErr, ok := common.IsServiceError(err); ok {
return fmt.Sprintf("%d", serviceErr.GetHTTPStatusCode()), sanitizeEffortText(serviceErr.GetMessage())
}
return "CallError", sanitizeEffortText(err.Error())
}
// parseEffortPayload 兼容三种 usage 形态:typed 驼峰(chatResponse.usage.completionTokens)、
// chat 兼容面下划线(usage.completion_tokens)、responses 面(usage.output_tokens)。
func parseEffortPayload(result *effortProbeResult, payload []byte) {
if text := strings.TrimSpace(string(payload)); strings.HasPrefix(text, "event:") || strings.HasPrefix(text, "data:") {
parseEffortSSE(result, text)
return
}
var root map[string]interface{}
if json.Unmarshal(payload, &root) != nil {
return
}
result.answerLen = len(parseEffortAnswer(root))
result.outputTypes = parseEffortOutputTypes(root)
if typed, ok := root["chatResponse"].(map[string]interface{}); ok {
root = typed
}
usage, _ := root["usage"].(map[string]interface{})
for _, key := range []string{"completionTokens", "completion_tokens", "output_tokens"} {
if v, ok := effortInt(usage, key); ok {
result.completionTokens = v
break
}
}
for _, detailsKey := range []string{"completionTokensDetails", "completion_tokens_details", "output_tokens_details"} {
details, _ := usage[detailsKey].(map[string]interface{})
for _, key := range []string{"reasoningTokens", "reasoning_tokens"} {
if v, ok := effortInt(details, key); ok {
result.reasoningTokens = v
}
}
}
}
// parseEffortAnswer 提取正文文本:typed 取 chatResponse.choices[].message.content[].text,
// responses 取 output[] 里 message 项的 content[].text。仅在内存中量长度,不落日志。
func parseEffortAnswer(root map[string]interface{}) string {
var sb strings.Builder
if typed, ok := root["chatResponse"].(map[string]interface{}); ok {
choices, _ := typed["choices"].([]interface{})
for _, c := range choices {
cm, _ := c.(map[string]interface{})
msg, _ := cm["message"].(map[string]interface{})
collectEffortText(&sb, msg["content"])
}
return sb.String()
}
output, _ := root["output"].([]interface{})
for _, item := range output {
im, _ := item.(map[string]interface{})
if im["type"] == "message" {
collectEffortText(&sb, im["content"])
}
}
return sb.String()
}
func collectEffortText(sb *strings.Builder, content interface{}) {
parts, _ := content.([]interface{})
for _, p := range parts {
pm, _ := p.(map[string]interface{})
if text, ok := pm["text"].(string); ok {
sb.WriteString(text)
}
}
}
// parseEffortSSE 汇总 SSE 流:事件类型去重序列进 outputTypes,正文增量长度进 answerLen,
// 末尾 completed 事件里的 usage 进 tokens 字段。
func parseEffortSSE(result *effortProbeResult, text string) {
seen := map[string]bool{}
for _, line := range strings.Split(text, "\n") {
if name, ok := strings.CutPrefix(line, "event: "); ok {
if name = strings.TrimSpace(name); !seen[name] {
seen[name] = true
result.outputTypes = append(result.outputTypes, name)
}
continue
}
data, ok := strings.CutPrefix(line, "data: ")
if !ok {
continue
}
var ev map[string]interface{}
if json.Unmarshal([]byte(data), &ev) != nil {
continue
}
if name, ok := ev["type"].(string); ok && !seen[name] { // 纯 data: 行流(无 event: 行)从事件体取类型
seen[name] = true
result.outputTypes = append(result.outputTypes, name)
}
if delta, ok := ev["delta"].(string); ok && strings.Contains(fmt.Sprint(ev["type"]), "output_text") {
result.answerLen += len(delta)
}
if resp, ok := ev["response"].(map[string]interface{}); ok {
usage, _ := resp["usage"].(map[string]interface{})
if v, ok := effortInt(usage, "output_tokens"); ok {
result.completionTokens = v
}
}
}
}
// parseEffortOutputTypes 收集 responses 面 output 项类型序列(multi-agent 行为观测)。
func parseEffortOutputTypes(root map[string]interface{}) []string {
output, _ := root["output"].([]interface{})
var types []string
for _, item := range output {
im, _ := item.(map[string]interface{})
if s, ok := im["type"].(string); ok {
types = append(types, s)
}
}
return types
}
func effortInt(root map[string]interface{}, key string) (int, bool) {
if root == nil {
return 0, false
}
if value, ok := root[key].(float64); ok {
return int(value), true
}
return 0, false
}
func logEffortResult(t *testing.T, tc effortProbeCase, result effortProbeResult) {
t.Helper()
endpoint := "typed-chat"
if tc.responses {
endpoint = "compat-responses"
}
t.Logf("endpoint=%s model=%s effort=%s status=%d completion=%d reasoning=%d answer_len=%d output_types=%v request=%s duration_ms=%d error=%s message=%q",
endpoint, tc.model, tc.effort, result.status, result.completionTokens, result.reasoningTokens,
result.answerLen, result.outputTypes, result.requestRef, result.duration.Milliseconds(),
result.errorCode, result.errorMessage)
}
// sanitizeEffortText 遮掩 OCID 并截断,避免错误消息携带租户可定位信息。
func sanitizeEffortText(text string) string {
text = regexp.MustCompile(`ocid1\.[a-z0-9._-]+`).ReplaceAllString(text, "ocid1.***")
if len(text) > 300 {
text = text[:300] + "..."
}
return text
}
func shortEffortRef(ref string) string {
if len(ref) > 12 {
return ref[:12]
}
return ref
}
func effortEnvOr(key, fallback string) string {
if value := os.Getenv(key); value != "" {
return value
}
return fallback
}
// newEffortProbeResources 装配凭据、区域、inference 客户端与按需模型目录。
func newEffortProbeResources(t *testing.T, ctx context.Context) effortProbeResources {
t.Helper()
cred := effortProbeCredentials(t)
region := effortEnvOr("OCI_CACHE_REGION", effortProbeRegion)
cred.Region = region
real := &RealClient{}
ic, err := real.genAiInferenceClient(cred, region)
if err != nil {
t.Fatalf("new inference client: %v", err)
}
models, err := real.ListGenAiModels(ctx, cred, region)
if err != nil {
t.Fatalf("list models: %s", sanitizeEffortText(err.Error()))
}
index := make(map[string]string, len(models))
for _, m := range models {
index[m.Name] = m.Ocid
}
return effortProbeResources{cred: cred, region: region, modelOCID: index, inference: ic}
}
// effortProbeCredentials 优先从面板数据库(只读)取渠道凭据,否则退回本地 ini 测试凭据。
func effortProbeCredentials(t *testing.T) Credentials {
t.Helper()
if dbPath := os.Getenv("OCI_CACHE_DB"); dbPath != "" {
return effortDBCredentials(t, dbPath)
}
return loadTestCredentials(t, effortEnvOr("OCI_TEST_KEY", "试用期"))
}
func effortDBCredentials(t *testing.T, dbPath string) Credentials {
t.Helper()
dsn := fmt.Sprintf("file:%s?mode=ro", dbPath)
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatalf("open read-only probe database: %v", err)
}
cipher, err := projectcrypto.NewCipher(os.Getenv("DATA_KEY"))
if err != nil {
t.Fatalf("create data cipher: %v", err)
}
var config model.OciConfig
if err := db.First(&config, effortEnvOr("OCI_CACHE_CONFIG_ID", "6")).Error; err != nil {
t.Fatalf("load OCI config: %v", err)
}
return decryptEffortCredentials(t, db, cipher, config)
}
func decryptEffortCredentials(t *testing.T, db *gorm.DB, cipher *projectcrypto.Cipher, config model.OciConfig) Credentials {
t.Helper()
privateKey, err := cipher.DecryptString(config.PrivateKeyEnc)
if err != nil {
t.Fatalf("decrypt OCI config %d private key: %v", config.ID, err)
}
passphrase := ""
if config.PassphraseEnc != "" {
if passphrase, err = cipher.DecryptString(config.PassphraseEnc); err != nil {
t.Fatalf("decrypt OCI config %d passphrase: %v", config.ID, err)
}
}
return Credentials{TenancyOCID: config.TenancyOCID, UserOCID: config.UserOCID,
Fingerprint: config.Fingerprint, Region: config.Region, PrivateKey: privateKey,
Passphrase: passphrase, Proxy: loadEffortProxy(t, db, cipher, config.ProxyID)}
}
func loadEffortProxy(t *testing.T, db *gorm.DB, cipher *projectcrypto.Cipher, proxyID *uint) *ProxySpec {
t.Helper()
if proxyID == nil {
return nil
}
var proxy model.Proxy
if err := db.First(&proxy, *proxyID).Error; err != nil {
t.Fatalf("load proxy %d: %v", *proxyID, err)
}
password := ""
if proxy.PasswordEnc != "" {
decrypted, err := cipher.DecryptString(proxy.PasswordEnc)
if err != nil {
t.Fatalf("decrypt proxy %d password: %v", *proxyID, err)
}
password = decrypted
}
return &ProxySpec{Type: proxy.Type, Host: proxy.Host, Port: proxy.Port,
Username: proxy.Username, Password: password}
}
// TestParseEffortCases 断言自定义用例解析:端点后缀、空项与畸形项跳过。
func TestParseEffortCases(t *testing.T) {
got := parseEffortCases("a=low, b=high@responses ,bad,=x,c=")
want := []effortProbeCase{{model: "a", effort: "low", responses: true}, {model: "b", effort: "high", responses: true}}
if len(got) != len(want) {
t.Fatalf("parseEffortCases = %+v, want %+v", got, want)
}
for i := range want {
if got[i] != want[i] {
t.Fatalf("case %d = %+v, want %+v", i, got[i], want[i])
}
}
}
// TestParseEffortPayload 断言三种 usage 形态与 output 类型序列的解析。
func TestParseEffortPayload(t *testing.T) {
cases := []struct {
name string
payload string
completion int
reasoning int
types int
}{
{"typed 驼峰", `{"chatResponse":{"choices":[{"message":{"content":[{"type":"TEXT","text":"24"}]}}],"usage":{"completionTokens":5,"completionTokensDetails":{"reasoningTokens":9}}}}`, 5, 9, 0},
{"responses 面", `{"output":[{"type":"reasoning"},{"type":"message","content":[{"type":"output_text","text":"24"}]}],"usage":{"output_tokens":7,"output_tokens_details":{"reasoning_tokens":3}}}`, 7, 3, 2},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
var result effortProbeResult
parseEffortPayload(&result, []byte(tc.payload))
if result.completionTokens != tc.completion || result.reasoningTokens != tc.reasoning || len(result.outputTypes) != tc.types {
t.Fatalf("parse = %+v, want completion=%d reasoning=%d types=%d", result, tc.completion, tc.reasoning, tc.types)
}
if result.answerLen != 2 {
t.Fatalf("answerLen = %d, want 2", result.answerLen)
}
})
}
}
+73
View File
@@ -0,0 +1,73 @@
package oci
import (
"bytes"
"context"
"fmt"
"io"
"net/http"
"github.com/oracle/oci-go-sdk/v65/common"
)
// compatResponsesLimit 限制直通响应体大小;web_search 输出含多段引用,给足余量。
const compatResponsesLimit = int64(8 << 20)
// GenAiCompatResponses 实现 Client:把 OpenAI Responses 请求体直通到 OCI
// `/20231130/actions/v1/responses`(IAM 签名)。该端点实测可执行 xAI 服务端工具
// (web_search / x_search),但不在 Oracle 文档化工具白名单内,行为可能随服务
// 版本、模型或区域变化;调用方须自行校验并改写请求体(store/stream)。
func (c *RealClient) GenAiCompatResponses(ctx context.Context, cred Credentials, region string, body []byte) ([]byte, error) {
ic, err := c.genAiInferenceClient(cred, region)
if err != nil {
return nil, err
}
client := ic.BaseClient
common.UpdateEndpointTemplateForOptions(&client)
common.SetMissingTemplateParams(&client)
request, err := http.NewRequestWithContext(ctx, http.MethodPost, "/actions/v1/responses", bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("build compat responses request: %w", err)
}
request.Header.Set("Content-Type", "application/json")
request.Header.Set("CompartmentId", cred.TenancyOCID)
request.Header.Set("opc-compartment-id", cred.TenancyOCID)
response, err := client.Call(ctx, request)
if err != nil {
return nil, err
}
defer response.Body.Close()
payload, err := io.ReadAll(io.LimitReader(response.Body, compatResponsesLimit))
if err != nil {
return nil, fmt.Errorf("read compat responses body: %w", err)
}
return payload, nil
}
// GenAiCompatResponsesStream 实现 Client:以流式直通 OCI `/actions/v1/responses`,
// 建立成功(2xx)返回 SSE body(调用方负责 Close);建立失败返回 SDK ServiceError,
// 与既有渠道切换/熔断错误分类兼容。请求体须由调用方置 stream:true。
func (c *RealClient) GenAiCompatResponsesStream(ctx context.Context, cred Credentials, region string, body []byte) (io.ReadCloser, error) {
ic, err := c.genAiInferenceClient(cred, region)
if err != nil {
return nil, err
}
client := ic.BaseClient
common.UpdateEndpointTemplateForOptions(&client)
common.SetMissingTemplateParams(&client)
request, err := http.NewRequestWithContext(ctx, http.MethodPost, "/actions/v1/responses", bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("build compat responses stream request: %w", err)
}
request.Header.Set("Content-Type", "application/json")
request.Header.Set("CompartmentId", cred.TenancyOCID)
request.Header.Set("opc-compartment-id", cred.TenancyOCID)
response, err := client.Call(ctx, request)
if err != nil {
if response != nil && response.Body != nil {
response.Body.Close()
}
return nil, err
}
return response.Body, nil
}
+30
View File
@@ -38,3 +38,33 @@ func TestToGenAiModelRetiredFilter(t *testing.T) {
t.Error("正常模型应入池") t.Error("正常模型应入池")
} }
} }
func TestDedupGenAiModelsFineTunePreference(t *testing.T) {
now := time.Date(2026, 7, 10, 0, 0, 0, 0, time.UTC)
chat := []generativeai.ModelCapabilityEnum{generativeai.ModelCapabilityChat}
chatFT := []generativeai.ModelCapabilityEnum{generativeai.ModelCapabilityChat, generativeai.ModelCapabilityFineTune}
mk := func(id, name string, caps []generativeai.ModelCapabilityEnum) generativeai.ModelSummary {
return generativeai.ModelSummary{
Id: common.String(id), DisplayName: common.String(name), Vendor: common.String("meta"),
Capabilities: caps, LifecycleState: generativeai.ModelLifecycleStateActive,
}
}
const name = "meta.llama-3-70b-instruct"
tests := []struct {
name string
items []generativeai.ModelSummary
want string // 期望保留的 OCID
}{
{"基座在前、纯对话在后:保留纯对话条目", []generativeai.ModelSummary{mk("o-ft", name, chatFT), mk("o-od", name, chat)}, "o-od"},
{"纯对话在前、基座在后:保留纯对话条目", []generativeai.ModelSummary{mk("o-od", name, chat), mk("o-ft", name, chatFT)}, "o-od"},
{"仅基座条目:保留不误删", []generativeai.ModelSummary{mk("o-ft", name, chatFT)}, "o-ft"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := dedupGenAiModels(tt.items, now)
if len(got) != 1 || got[0].Ocid != tt.want {
t.Errorf("dedupGenAiModels() = %+v, want 仅保留 %s", got, tt.want)
}
})
}
}
+10 -176
View File
@@ -1,193 +1,27 @@
package oci package oci
import ( import (
"encoding/json"
"testing" "testing"
"github.com/oracle/oci-go-sdk/v65/generativeaiinference" "github.com/oracle/oci-go-sdk/v65/generativeaiinference"
"oci-portal/internal/aiwire"
) )
// TestIrToSDK 断言 IR→GENERIC 的角色拆装、采样参数与工具映射 // TestSdkUsageCachedTokens 断言 embed 用量映射:缓存命中挂 details,仅命中 >0 时出现
func TestIrToSDK(t *testing.T) {
temp, mt := 0.7, 100
ir := aiwire.ChatRequest{
Model: "meta.llama-3.3-70b-instruct",
Messages: []aiwire.ChatMessage{
{Role: "system", Content: aiwire.NewTextContent("你是助手")},
{Role: "user", Content: aiwire.NewTextContent("你好")},
{Role: "assistant", ToolCalls: []aiwire.ToolCall{{ID: "c1", Type: "function", Function: aiwire.FunctionCall{Name: "get_weather", Arguments: `{"city":"东京"}`}}}},
{Role: "tool", ToolCallID: "c1", Content: aiwire.NewTextContent("晴")},
},
Temperature: &temp,
MaxTokens: &mt,
Stop: aiwire.StringList{"END"},
Tools: []aiwire.Tool{{Type: "function", Function: aiwire.FunctionDef{Name: "get_weather", Parameters: json.RawMessage(`{"type":"object"}`)}}},
ToolChoice: json.RawMessage(`"auto"`),
}
req := irToSDK(ir, true)
if len(req.Messages) != 4 {
t.Fatalf("messages = %d, want 4", len(req.Messages))
}
if _, ok := req.Messages[0].(generativeaiinference.SystemMessage); !ok {
t.Fatalf("messages[0] = %T, want SystemMessage", req.Messages[0])
}
am, ok := req.Messages[2].(generativeaiinference.AssistantMessage)
if !ok || len(am.ToolCalls) != 1 {
t.Fatalf("assistant tool calls 未映射: %T %+v", req.Messages[2], am)
}
tm, ok := req.Messages[3].(generativeaiinference.ToolMessage)
if !ok || deref(tm.ToolCallId) != "c1" {
t.Fatalf("tool message 未映射 toolCallId: %+v", tm)
}
if req.Temperature == nil || *req.Temperature != 0.7 {
t.Fatalf("temperature 未直通")
}
if req.MaxTokens == nil || *req.MaxTokens != 100 || len(req.Stop) != 1 {
t.Fatalf("maxTokens/stop 未直通")
}
if _, ok := req.ToolChoice.(generativeaiinference.ToolChoiceAuto); !ok {
t.Fatalf("toolChoice = %T, want auto", req.ToolChoice)
}
if len(req.Tools) != 1 || req.IsStream == nil || !*req.IsStream || req.StreamOptions == nil {
t.Fatalf("tools/stream 选项未映射")
}
}
// TestSdkToIR 断言响应文本、工具调用与用量的反向映射。
func TestSdkToIR(t *testing.T) {
text, fr := "东京晴", "tool_calls"
idx, pt, ct, tt := 0, 10, 5, 15
id, name, args := "c1", "get_weather", `{"city":"东京"}`
resp := generativeaiinference.GenericChatResponse{
Choices: []generativeaiinference.ChatChoice{{
Index: &idx,
Message: generativeaiinference.AssistantMessage{
Content: []generativeaiinference.ChatContent{generativeaiinference.TextContent{Text: &text}},
ToolCalls: []generativeaiinference.ToolCall{generativeaiinference.FunctionCall{Id: &id, Name: &name, Arguments: &args}},
},
FinishReason: &fr,
}},
Usage: &generativeaiinference.Usage{PromptTokens: &pt, CompletionTokens: &ct, TotalTokens: &tt},
}
out := sdkToIR(resp, "m1")
if len(out.Choices) != 1 || out.Choices[0].Message.Content.JoinText() != "东京晴" {
t.Fatalf("文本未映射: %+v", out)
}
if out.Choices[0].FinishReason != "tool_calls" || len(out.Choices[0].Message.ToolCalls) != 1 {
t.Fatalf("finishReason/toolCalls 未映射: %+v", out.Choices[0])
}
if out.Usage == nil || out.Usage.TotalTokens != 15 {
t.Fatalf("usage 未映射: %+v", out.Usage)
}
}
// TestParseGenAiEvent 断言流事件宽容解析:文本增量、结束原因与 usage 事件。
func TestParseGenAiEvent(t *testing.T) {
chunk, ok := parseGenAiEvent([]byte(`{"index":0,"message":{"role":"ASSISTANT","content":[{"type":"TEXT","text":"你"}]}}`), "m1")
if !ok || len(chunk.Choices) != 1 || chunk.Choices[0].Delta.Content != "你" {
t.Fatalf("文本增量解析失败: %+v", chunk)
}
chunk, ok = parseGenAiEvent([]byte(`{"finishReason":"stop","usage":{"promptTokens":3,"completionTokens":2,"totalTokens":5}}`), "m1")
if !ok || chunk.Choices[0].FinishReason == nil || *chunk.Choices[0].FinishReason != "stop" {
t.Fatalf("finishReason 解析失败: %+v", chunk)
}
if chunk.Usage == nil || chunk.Usage.TotalTokens != 5 {
t.Fatalf("usage 解析失败: %+v", chunk.Usage)
}
if _, ok := parseGenAiEvent([]byte(`{}`), "m1"); ok {
t.Fatal("空事件应返回 ok=false")
}
}
func TestSdkUsageCachedTokens(t *testing.T) { func TestSdkUsageCachedTokens(t *testing.T) {
p, c, tot, cached := 10, 5, 15, 8 p, c, tot, cached := 10, 5, 15, 8
u := sdkUsageToIR(&generativeaiinference.Usage{PromptTokens: &p, CompletionTokens: &c, TotalTokens: &tot, u := sdkUsageToIR(&generativeaiinference.Usage{PromptTokens: &p, CompletionTokens: &c, TotalTokens: &tot,
PromptTokensDetails: &generativeaiinference.PromptTokensDetails{CachedTokens: &cached}}) PromptTokensDetails: &generativeaiinference.PromptTokensDetails{CachedTokens: &cached}})
if u.CachedTokens() != 8 { if u.PromptTokens != 10 || u.CompletionTokens != 5 || u.TotalTokens != 15 {
t.Errorf("CachedTokens = %d, want 8", u.CachedTokens()) t.Fatalf("usage 映射错误: %+v", u)
} }
if sdkUsageToIR(&generativeaiinference.Usage{PromptTokens: &p}).CachedTokens() != 0 { if u.PromptTokensDetails == nil || u.PromptTokensDetails.CachedTokens != 8 {
t.Error("无细分时 CachedTokens 应为 0") t.Fatalf("cachedTokens 未映射: %+v", u.PromptTokensDetails)
} }
} zero := 0
if got := sdkUsageToIR(&generativeaiinference.Usage{PromptTokens: &p, PromptTokensDetails: &generativeaiinference.PromptTokensDetails{CachedTokens: &zero}}); got.PromptTokensDetails != nil {
func TestIrToCohereSDK(t *testing.T) { t.Fatalf("零命中不应带 details: %+v", got.PromptTokensDetails)
ir := aiwire.ChatRequest{
Model: "cohere.command-r-plus",
Messages: []aiwire.ChatMessage{
{Role: "system", Content: aiwire.NewTextContent("你是助手")},
{Role: "user", Content: aiwire.NewTextContent("第一问")},
{Role: "assistant", Content: aiwire.NewTextContent("第一答")},
{Role: "user", Content: aiwire.NewTextContent("第二问")},
},
} }
req, err := irToCohereSDK(ir, true) if sdkUsageToIR(nil) != nil {
if err != nil { t.Fatal("nil usage 应返回 nil")
t.Fatalf("irToCohereSDK: %v", err)
}
if *req.Message != "第二问" || *req.PreambleOverride != "你是助手" || len(req.ChatHistory) != 2 {
t.Errorf("拆装错误: msg=%q preamble=%v history=%d", *req.Message, req.PreambleOverride, len(req.ChatHistory))
}
if _, ok := req.ChatHistory[0].(generativeaiinference.CohereUserMessage); !ok {
t.Errorf("history[0] 应为 user: %T", req.ChatHistory[0])
}
if _, ok := req.ChatHistory[1].(generativeaiinference.CohereChatBotMessage); !ok {
t.Errorf("history[1] 应为 chatbot: %T", req.ChatHistory[1])
}
if req.IsStream == nil || !*req.IsStream {
t.Error("IsStream 未设置")
}
if req.StreamOptions == nil || req.StreamOptions.IsIncludeUsage == nil || !*req.StreamOptions.IsIncludeUsage {
t.Error("流式应默认开启 usage 回传(StreamOptions.IsIncludeUsage)")
}
// 工具定义降级为扁平参数表(仅取 JSON Schema 顶层 properties)
ir.Tools = []aiwire.Tool{{Type: "function", Function: aiwire.FunctionDef{
Name: "get_weather", Parameters: json.RawMessage(`{"type":"object","properties":{"city":{"type":"string","description":"城市"}},"required":["city"]}`)}}}
req2, err := irToCohereSDK(ir, false)
if err != nil {
t.Fatalf("工具请求应支持: %v", err)
}
tool := req2.Tools[0]
if *tool.Name != "get_weather" || *tool.Description != "get_weather" {
t.Errorf("tool = %+v", tool)
}
if pd, ok := tool.ParameterDefinitions["city"]; !ok || *pd.Type != "string" || pd.IsRequired == nil || !*pd.IsRequired {
t.Errorf("param city = %+v", pd)
}
// 图片输入仍拒绝
ir.Tools = nil
ir.Messages = append(ir.Messages, aiwire.ChatMessage{Role: "user", Content: aiwire.NewPartsContent([]aiwire.ContentPart{
{Type: "image_url", ImageURL: &aiwire.ImageURL{URL: "data:image/png;base64,xx"}}})})
if _, err := irToCohereSDK(ir, false); err == nil {
t.Error("图片输入应被拒绝")
}
// 无 user 消息被拒
if _, err := irToCohereSDK(aiwire.ChatRequest{Model: "cohere.x", Messages: []aiwire.ChatMessage{{Role: "system", Content: aiwire.NewTextContent("s")}}}, false); err == nil {
t.Error("无 user 消息应被拒绝")
}
}
func TestCohereSDKToIR(t *testing.T) {
text := "回答"
resp := generativeaiinference.CohereChatResponse{
Text: &text,
FinishReason: generativeaiinference.CohereChatResponseFinishReasonComplete,
}
out := cohereSDKToIR(resp, "cohere.command-r-plus")
if out.Choices[0].Message.Content.JoinText() != "回答" || out.Choices[0].FinishReason != "stop" {
t.Errorf("cohereSDKToIR = %+v", out.Choices[0])
}
}
func TestParseGenAiEventCohere(t *testing.T) {
chunk, ok := parseGenAiEvent([]byte(`{"apiFormat":"COHERE","text":"你好"}`), "cohere.command-r")
if !ok || chunk.Choices[0].Delta.Content != "你好" {
t.Errorf("cohere text 事件解析 = %+v, %v", chunk, ok)
}
chunk, ok = parseGenAiEvent([]byte(`{"apiFormat":"COHERE","finishReason":"COMPLETE","usage":{"promptTokens":3,"completionTokens":5,"totalTokens":8}}`), "cohere.command-r")
if !ok || *chunk.Choices[0].FinishReason != "stop" || chunk.Usage.TotalTokens != 8 {
t.Errorf("cohere 终帧解析 = %+v, %v", chunk, ok)
} }
} }
-332
View File
@@ -1,332 +0,0 @@
package service
import (
"encoding/json"
"fmt"
"strings"
"oci-portal/internal/aiwire"
)
// ErrAiUnsupportedBlock 表示请求含网关无法承接的内容块(仅支持文本与图片)。
var ErrAiUnsupportedBlock = fmt.Errorf("暂不支持文本与图片以外的内容块")
// AnthropicToIR 把 Anthropic Messages 请求转为 IR(OpenAI 线格式)。
// tool_result 块拆为独立 tool 角色消息(一拆多),顶层 system 变首条 system 消息。
func AnthropicToIR(req aiwire.MessagesRequest) (aiwire.ChatRequest, error) {
ir := aiwire.ChatRequest{
Model: req.Model,
Temperature: req.Temperature,
TopP: req.TopP,
TopK: req.TopK,
Stop: req.StopSequences,
Stream: req.Stream,
Tools: anthTools(req.Tools),
ToolChoice: anthToolChoice(req.ToolChoice),
}
if req.MaxTokens > 0 {
mt := req.MaxTokens
ir.MaxTokens = &mt
}
if sys := req.SystemText(); sys != "" {
ir.Messages = append(ir.Messages, aiwire.ChatMessage{Role: "system", Content: aiwire.NewTextContent(sys)})
}
for _, m := range req.Messages {
msgs, err := anthMessageToIR(m)
if err != nil {
return ir, err
}
ir.Messages = append(ir.Messages, msgs...)
}
return ir, nil
}
// anthMessageToIR 拆解单条 Anthropic 消息;user 消息里的 tool_result 前置为独立 tool 消息,
// image 块转为 IR image_url 部件。
func anthMessageToIR(m aiwire.AnthMessage) ([]aiwire.ChatMessage, error) {
var out []aiwire.ChatMessage
var parts []aiwire.ContentPart
var toolCalls []aiwire.ToolCall
hasImage := false
for _, b := range m.Content.AllBlocks() {
switch b.Type {
case "text":
parts = append(parts, aiwire.ContentPart{Type: "text", Text: b.Text})
case "image":
p, err := anthImagePart(b.Source)
if err != nil {
return nil, err
}
parts, hasImage = append(parts, p), true
case "tool_use":
toolCalls = append(toolCalls, aiwire.ToolCall{ID: b.ID, Type: "function",
Function: aiwire.FunctionCall{Name: b.Name, Arguments: string(b.Input)}})
case "tool_result":
out = append(out, aiwire.ChatMessage{Role: "tool", ToolCallID: b.ToolUseID,
Content: aiwire.NewTextContent(b.ResultText())})
case "thinking", "redacted_thinking":
// 一期忽略 thinking 块
default:
return nil, ErrAiUnsupportedBlock
}
}
content := partsContent(parts, hasImage)
if hasImage || content.JoinText() != "" || len(toolCalls) > 0 {
out = append(out, aiwire.ChatMessage{Role: m.Role, Content: content, ToolCalls: toolCalls})
}
return out, nil
}
// anthImagePart 把 Anthropic image 块转为 IR image_url 部件(base64 → data URI)。
func anthImagePart(source json.RawMessage) (aiwire.ContentPart, error) {
var src struct {
Type string `json:"type"`
MediaType string `json:"media_type"`
Data string `json:"data"`
URL string `json:"url"`
}
if err := json.Unmarshal(source, &src); err != nil {
return aiwire.ContentPart{}, fmt.Errorf("image source 解析失败: %w", err)
}
switch src.Type {
case "base64":
if src.MediaType == "" || src.Data == "" {
return aiwire.ContentPart{}, fmt.Errorf("image source 缺少 media_type 或 data")
}
url := "data:" + src.MediaType + ";base64," + src.Data
return aiwire.ContentPart{Type: "image_url", ImageURL: &aiwire.ImageURL{URL: url}}, nil
case "url":
if src.URL == "" {
return aiwire.ContentPart{}, fmt.Errorf("image source 缺少 url")
}
return aiwire.ContentPart{Type: "image_url", ImageURL: &aiwire.ImageURL{URL: src.URL}}, nil
}
return aiwire.ContentPart{}, fmt.Errorf("不支持的 image source 类型 %q", src.Type)
}
// partsContent 无图片时退回字符串形态(与上游线格式习惯一致),含图片时保留块数组。
func partsContent(parts []aiwire.ContentPart, hasImage bool) aiwire.Content {
if !hasImage {
var sb strings.Builder
for _, p := range parts {
sb.WriteString(p.Text)
}
return aiwire.NewTextContent(sb.String())
}
return aiwire.NewPartsContent(parts)
}
func anthTools(tools []aiwire.AnthTool) []aiwire.Tool {
if len(tools) == 0 {
return nil
}
out := make([]aiwire.Tool, 0, len(tools))
for _, t := range tools {
out = append(out, aiwire.Tool{Type: "function", Function: aiwire.FunctionDef{
Name: t.Name, Description: t.Description, Parameters: t.InputSchema}})
}
return out
}
// anthToolChoice 映射 {type:auto|any|tool,name} → OpenAI 形态。
func anthToolChoice(raw json.RawMessage) json.RawMessage {
if len(raw) == 0 {
return nil
}
var tc struct {
Type string `json:"type"`
Name string `json:"name"`
}
if json.Unmarshal(raw, &tc) != nil {
return nil
}
switch tc.Type {
case "auto":
return json.RawMessage(`"auto"`)
case "any":
return json.RawMessage(`"required"`)
case "tool":
b, _ := json.Marshal(map[string]any{"type": "function", "function": map[string]string{"name": tc.Name}})
return b
case "none":
return json.RawMessage(`"none"`)
}
return nil
}
// IRRespToAnthropic 把 IR 非流式响应转为 Anthropic Messages 响应。
func IRRespToAnthropic(resp *aiwire.ChatResponse, id string) aiwire.MessagesResponse {
out := aiwire.MessagesResponse{ID: id, Type: "message", Role: "assistant", Model: resp.Model, Content: []aiwire.AnthBlock{}}
if len(resp.Choices) > 0 {
choice := resp.Choices[0]
if text := choice.Message.Content.JoinText(); text != "" {
out.Content = append(out.Content, aiwire.AnthBlock{Type: "text", Text: text})
}
for _, tc := range choice.Message.ToolCalls {
out.Content = append(out.Content, aiwire.AnthBlock{Type: "tool_use", ID: tc.ID,
Name: tc.Function.Name, Input: argsToJSON(tc.Function.Arguments)})
}
out.StopReason = anthStopReason(choice.FinishReason)
}
if resp.Usage != nil {
out.Usage = aiwire.AnthUsage{InputTokens: resp.Usage.PromptTokens, OutputTokens: resp.Usage.CompletionTokens,
CacheReadInputTokens: resp.Usage.CachedTokens()}
}
return out
}
// argsToJSON 保证 tool_use.input 是合法 JSON 对象(模型可能产出非法片段)。
func argsToJSON(args string) json.RawMessage {
trimmed := strings.TrimSpace(args)
if trimmed == "" {
return json.RawMessage(`{}`)
}
if json.Valid([]byte(trimmed)) {
return json.RawMessage(trimmed)
}
b, _ := json.Marshal(map[string]string{"_raw": args})
return b
}
func anthStopReason(finish string) string {
switch finish {
case "length":
return "max_tokens"
case "tool_calls":
return "tool_use"
default:
return "end_turn"
}
}
// ---- Anthropic 流式状态机 ----
// AnthEvent 是一条待写出的 Anthropic SSE 事件。
type AnthEvent struct {
Event string
Data any
}
// AnthStream 把 IR chunk 流聚合为 Anthropic 事件序列:
// message_start → content_block_start/delta/stop(text 与 tool_use 分块)→ message_delta → message_stop。
type AnthStream struct {
id, model string
started bool
blockOpen bool
blockIsTool bool
toolID string
blockIndex int
stopReason string
usage aiwire.AnthUsage
}
// NewAnthStream 构造状态机;id 为响应消息 ID。
func NewAnthStream(id, model string) *AnthStream {
return &AnthStream{id: id, model: model, blockIndex: -1, stopReason: "end_turn"}
}
// Feed 消费一个 IR chunk,返回应立即写出的事件。
func (st *AnthStream) Feed(chunk aiwire.ChatChunk) []AnthEvent {
var events []AnthEvent
if !st.started {
st.started = true
events = append(events, st.startEvent())
}
if chunk.Usage != nil {
st.usage = aiwire.AnthUsage{InputTokens: chunk.Usage.PromptTokens, OutputTokens: chunk.Usage.CompletionTokens,
CacheReadInputTokens: chunk.Usage.CachedTokens()}
}
for _, choice := range chunk.Choices {
events = append(events, st.feedDelta(choice.Delta)...)
if choice.FinishReason != nil && *choice.FinishReason != "" {
st.stopReason = anthStopReason(*choice.FinishReason)
}
}
return events
}
func (st *AnthStream) startEvent() AnthEvent {
return AnthEvent{Event: "message_start", Data: map[string]any{
"type": "message_start",
"message": map[string]any{
"id": st.id, "type": "message", "role": "assistant", "model": st.model,
"content": []any{}, "stop_reason": nil,
"usage": map[string]int{"input_tokens": 0, "output_tokens": 0},
},
}}
}
// feedDelta 处理文本与工具调用增量,必要时切块。
func (st *AnthStream) feedDelta(d aiwire.Delta) []AnthEvent {
var events []AnthEvent
if d.Content != "" {
if !st.blockOpen || st.blockIsTool {
events = append(events, st.openBlock(false, "", "")...)
}
events = append(events, AnthEvent{Event: "content_block_delta", Data: map[string]any{
"type": "content_block_delta", "index": st.blockIndex,
"delta": map[string]string{"type": "text_delta", "text": d.Content},
}})
}
for _, tc := range d.ToolCalls {
if tc.ID != "" && (!st.blockOpen || !st.blockIsTool || st.toolID != tc.ID) {
events = append(events, st.openBlock(true, tc.ID, tc.Function.Name)...)
}
if tc.Function.Arguments != "" && st.blockOpen && st.blockIsTool {
events = append(events, AnthEvent{Event: "content_block_delta", Data: map[string]any{
"type": "content_block_delta", "index": st.blockIndex,
"delta": map[string]string{"type": "input_json_delta", "partial_json": tc.Function.Arguments},
}})
}
}
return events
}
// openBlock 关闭当前块并打开新块(text 或 tool_use)。
func (st *AnthStream) openBlock(isTool bool, toolID, toolName string) []AnthEvent {
var events []AnthEvent
if st.blockOpen {
events = append(events, st.closeBlockEvent())
}
st.blockOpen, st.blockIsTool, st.toolID = true, isTool, toolID
st.blockIndex++
block := map[string]any{"type": "text", "text": ""}
if isTool {
block = map[string]any{"type": "tool_use", "id": toolID, "name": toolName, "input": map[string]any{}}
}
events = append(events, AnthEvent{Event: "content_block_start", Data: map[string]any{
"type": "content_block_start", "index": st.blockIndex, "content_block": block,
}})
return events
}
func (st *AnthStream) closeBlockEvent() AnthEvent {
return AnthEvent{Event: "content_block_stop", Data: map[string]any{
"type": "content_block_stop", "index": st.blockIndex,
}}
}
// Finish 在上游流结束后收尾:关块 → message_delta(stop_reason+usage)→ message_stop。
func (st *AnthStream) Finish() []AnthEvent {
var events []AnthEvent
if !st.started {
events = append(events, st.startEvent())
st.started = true
}
if st.blockOpen {
events = append(events, st.closeBlockEvent())
st.blockOpen = false
}
events = append(events,
AnthEvent{Event: "message_delta", Data: map[string]any{
"type": "message_delta",
"delta": map[string]any{"stop_reason": st.stopReason, "stop_sequence": nil},
"usage": map[string]int{"output_tokens": st.usage.OutputTokens},
}},
AnthEvent{Event: "message_stop", Data: map[string]any{"type": "message_stop"}},
)
return events
}
// Usage 返回聚合到的用量(供调用日志)。
func (st *AnthStream) Usage() aiwire.AnthUsage { return st.usage }
-136
View File
@@ -1,136 +0,0 @@
package service
import (
"encoding/json"
"strings"
"testing"
"oci-portal/internal/aiwire"
)
func mustAnthReq(t *testing.T, body string) aiwire.MessagesRequest {
t.Helper()
var req aiwire.MessagesRequest
if err := json.Unmarshal([]byte(body), &req); err != nil {
t.Fatalf("unmarshal anthropic request: %v", err)
}
return req
}
func TestAnthropicToIR(t *testing.T) {
req := mustAnthReq(t, `{
"model": "meta.llama-3.3-70b-instruct", "max_tokens": 128, "system": "你是助手",
"messages": [
{"role": "user", "content": "东京天气?"},
{"role": "assistant", "content": [
{"type": "text", "text": "查询中"},
{"type": "tool_use", "id": "t1", "name": "get_weather", "input": {"city": "东京"}}
]},
{"role": "user", "content": [
{"type": "tool_result", "tool_use_id": "t1", "content": "晴 25 度"},
{"type": "text", "text": "继续"}
]}
],
"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}],
"tool_choice": {"type": "any"}
}`)
ir, err := AnthropicToIR(req)
if err != nil {
t.Fatalf("AnthropicToIR: %v", err)
}
roles := make([]string, 0, len(ir.Messages))
for _, m := range ir.Messages {
roles = append(roles, m.Role)
}
want := []string{"system", "user", "assistant", "tool", "user"}
if strings.Join(roles, ",") != strings.Join(want, ",") {
t.Fatalf("roles = %v, want %v", roles, want)
}
if ir.Messages[2].ToolCalls[0].Function.Name != "get_weather" {
t.Errorf("tool_use 未映射: %+v", ir.Messages[2])
}
if ir.Messages[3].ToolCallID != "t1" || ir.Messages[3].Content.JoinText() != "晴 25 度" {
t.Errorf("tool_result 未拆为 tool 消息: %+v", ir.Messages[3])
}
if ir.MaxTokens == nil || *ir.MaxTokens != 128 || len(ir.Tools) != 1 {
t.Errorf("max_tokens/tools 未直通")
}
if string(ir.ToolChoice) != `"required"` {
t.Errorf("tool_choice any → %s, want required", ir.ToolChoice)
}
// image 块拒绝
bad := mustAnthReq(t, `{"model":"m","max_tokens":1,"messages":[{"role":"user","content":[{"type":"image","source":{}}]}]}`)
if _, err := AnthropicToIR(bad); err == nil {
t.Error("image 块应被拒绝")
}
}
func TestIRRespToAnthropic(t *testing.T) {
resp := &aiwire.ChatResponse{
Model: "m1",
Choices: []aiwire.Choice{{
Message: aiwire.ChatMessage{
Role: "assistant",
Content: aiwire.NewTextContent("查到了"),
ToolCalls: []aiwire.ToolCall{{ID: "t1", Type: "function",
Function: aiwire.FunctionCall{Name: "get_weather", Arguments: `{"city":"东京"}`}}},
},
FinishReason: "tool_calls",
}},
Usage: &aiwire.Usage{PromptTokens: 10, CompletionTokens: 5},
}
out := IRRespToAnthropic(resp, "msg_1")
if len(out.Content) != 2 || out.Content[0].Type != "text" || out.Content[1].Type != "tool_use" {
t.Fatalf("content 块 = %+v", out.Content)
}
if out.StopReason != "tool_use" || out.Usage.InputTokens != 10 || out.Usage.OutputTokens != 5 {
t.Errorf("stop_reason/usage 未映射: %+v", out)
}
if string(out.Content[1].Input) != `{"city":"东京"}` {
t.Errorf("tool input = %s", out.Content[1].Input)
}
}
// eventTypes 提取事件类型序列便于断言。
func eventTypes(events []AnthEvent) string {
types := make([]string, 0, len(events))
for _, e := range events {
types = append(types, e.Event)
}
return strings.Join(types, ",")
}
func TestAnthStreamTextAndTool(t *testing.T) {
st := NewAnthStream("msg_1", "m1")
fr := "tool_calls"
var events []AnthEvent
events = append(events, st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{Content: "你"}}}})...)
events = append(events, st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{Content: "好"}}}})...)
events = append(events, st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{
ToolCalls: []aiwire.ToolCallDelta{{ID: "t1", Function: aiwire.FunctionCallDelta{Name: "get_weather", Arguments: `{"ci`}}},
}}}})...)
events = append(events, st.Feed(aiwire.ChatChunk{
Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{ToolCalls: []aiwire.ToolCallDelta{{ID: "t1", Function: aiwire.FunctionCallDelta{Arguments: `ty":"东京"}`}}}}, FinishReason: &fr}},
Usage: &aiwire.Usage{PromptTokens: 8, CompletionTokens: 4},
})...)
events = append(events, st.Finish()...)
got := eventTypes(events)
want := "message_start,content_block_start,content_block_delta,content_block_delta," +
"content_block_stop,content_block_start,content_block_delta,content_block_delta," +
"content_block_stop,message_delta,message_stop"
if got != want {
t.Fatalf("事件序列:\n got %s\nwant %s", got, want)
}
if st.Usage().OutputTokens != 4 {
t.Errorf("usage 未聚合: %+v", st.Usage())
}
}
func TestAnthStreamEmpty(t *testing.T) {
st := NewAnthStream("msg_1", "m1")
events := st.Finish()
if got := eventTypes(events); got != "message_start,message_delta,message_stop" {
t.Errorf("空流事件序列 = %s", got)
}
}
+189 -15
View File
@@ -5,6 +5,7 @@ import (
"crypto/rand" "crypto/rand"
"crypto/sha256" "crypto/sha256"
"encoding/hex" "encoding/hex"
"encoding/json"
"errors" "errors"
"fmt" "fmt"
"log" "log"
@@ -79,7 +80,7 @@ func (s *AiGatewayService) fireChannelsChanged(ctx context.Context) {
// CreateKey 生成网关密钥并返回明文(仅此一次);customValue 非空时使用给定值, // CreateKey 生成网关密钥并返回明文(仅此一次);customValue 非空时使用给定值,
// group 非空时该密钥只在同分组渠道内路由。 // group 非空时该密钥只在同分组渠道内路由。
func (s *AiGatewayService) CreateKey(ctx context.Context, name, customValue, group string) (string, *model.AiKey, error) { func (s *AiGatewayService) CreateKey(ctx context.Context, name, customValue, group string, models []string) (string, *model.AiKey, error) {
name = strings.TrimSpace(name) name = strings.TrimSpace(name)
if name == "" { if name == "" {
return "", nil, fmt.Errorf("密钥名称不能为空") return "", nil, fmt.Errorf("密钥名称不能为空")
@@ -95,13 +96,28 @@ func (s *AiGatewayService) CreateKey(ctx context.Context, name, customValue, gro
if len(raw) < 8 { if len(raw) < 8 {
return "", nil, fmt.Errorf("自定义密钥至少 8 个字符") return "", nil, fmt.Errorf("自定义密钥至少 8 个字符")
} }
key := &model.AiKey{Name: name, KeyHash: hashKey(raw), Tail: raw[len(raw)-4:], Group: strings.TrimSpace(group), Enabled: true} key := &model.AiKey{Name: name, KeyHash: hashKey(raw), Tail: raw[len(raw)-4:], Group: strings.TrimSpace(group), Models: normalizeKeyModels(models), Enabled: true}
if err := s.db.WithContext(ctx).Create(key).Error; err != nil { if err := s.db.WithContext(ctx).Create(key).Error; err != nil {
return "", nil, fmt.Errorf("密钥名称或取值与现有密钥重复") return "", nil, fmt.Errorf("密钥名称或取值与现有密钥重复")
} }
return raw, key, nil return raw, key, nil
} }
// normalizeKeyModels 规范化模型白名单:trim、剔空串、保序去重;结果为空返回 nil(= 不限)。
func normalizeKeyModels(models []string) []string {
var out []string
seen := map[string]bool{}
for _, m := range models {
m = strings.TrimSpace(m)
if m == "" || seen[m] {
continue
}
seen[m] = true
out = append(out, m)
}
return out
}
func hashKey(raw string) string { func hashKey(raw string) string {
sum := sha256.Sum256([]byte(raw)) sum := sha256.Sum256([]byte(raw))
return hex.EncodeToString(sum[:]) return hex.EncodeToString(sum[:])
@@ -114,8 +130,8 @@ func (s *AiGatewayService) Keys(ctx context.Context) ([]model.AiKey, error) {
return keys, err return keys, err
} }
// UpdateKey 修改密钥名称 / 启用状态 / 分组(group 指针非空即覆盖,可置空)。 // UpdateKey 修改密钥名称 / 启用状态 / 分组 / 模型白名单(指针非空即覆盖,可置空)。
func (s *AiGatewayService) UpdateKey(ctx context.Context, id uint, name string, enabled *bool, group *string) error { func (s *AiGatewayService) UpdateKey(ctx context.Context, id uint, name string, enabled *bool, group *string, models *[]string) error {
updates := map[string]any{} updates := map[string]any{}
if name = strings.TrimSpace(name); name != "" { if name = strings.TrimSpace(name); name != "" {
updates["name"] = name updates["name"] = name
@@ -126,6 +142,14 @@ func (s *AiGatewayService) UpdateKey(ctx context.Context, id uint, name string,
if group != nil { if group != nil {
updates["key_group"] = strings.TrimSpace(*group) updates["key_group"] = strings.TrimSpace(*group)
} }
if models != nil {
// 手动序列化走 map 更新,不依赖 GORM map 路径对 serializer 的支持
b, err := json.Marshal(normalizeKeyModels(*models))
if err != nil {
return fmt.Errorf("serialize models: %w", err)
}
updates["models"] = string(b)
}
if len(updates) == 0 { if len(updates) == 0 {
return nil return nil
} }
@@ -290,12 +314,17 @@ func (s *AiGatewayService) ProbeChannel(ctx context.Context, id uint) (*model.Ai
return &fresh, nil return &fresh, nil
} }
// probe 执行探测并返回 (状态, 错误摘要);同时完成模型缓存同步。 // probe 执行探测并返回 (状态, 错误摘要);同时完成模型缓存同步(黑名单模型不入库)
func (s *AiGatewayService) probe(ctx context.Context, cred oci.Credentials, ch *model.AiChannel) (string, string) { func (s *AiGatewayService) probe(ctx context.Context, cred oci.Credentials, ch *model.AiChannel) (string, string) {
models, err := s.client.ListGenAiModels(ctx, cred, ch.Region) models, err := s.client.ListGenAiModels(ctx, cred, ch.Region)
if err != nil { if err != nil {
return classifyProbeErr(err), truncateErr(oci.CompactError(err)) return classifyProbeErr(err), truncateErr(oci.CompactError(err))
} }
models = supportedGatewayModels(models)
models, err = s.withoutBlacklisted(ctx, models)
if err != nil {
return "error", truncateErr(err.Error())
}
if len(models) == 0 { if len(models) == 0 {
_ = s.replaceModels(ctx, ch.ID, nil) _ = s.replaceModels(ctx, ch.ID, nil)
return "no_service", "区域无可用模型(GenAI 服务不可用或未开放)" return "no_service", "区域无可用模型(GenAI 服务不可用或未开放)"
@@ -306,40 +335,88 @@ func (s *AiGatewayService) probe(ctx context.Context, cred oci.Credentials, ch *
return s.probeChat(ctx, cred, ch, models) return s.probeChat(ctx, cred, ch, models)
} }
// probeChat 按偏好挑选至多 3 个模型依次试调:部分模型元数据标 CHAT 但实际 // probeChat 按候选顺序试调(上限 8):遇「模型不可按需调用」(微调基座 400 / 实体
// 不可对话(如 voice agent),遇 400/5xx 换下一个;401/403/404 属租户级直接定论。 // 不存在 404)换下一个候选,错误信息带模型名供用户加入黑名单;401/403 与鉴权类 404
// 属租户级直接定论 no_quota;其余错误(元数据标 CHAT 但实际不可对话等)累计 3 次止损。
func (s *AiGatewayService) probeChat(ctx context.Context, cred oci.Credentials, ch *model.AiChannel, models []oci.GenAiModel) (string, string) { func (s *AiGatewayService) probeChat(ctx context.Context, cred oci.Credentials, ch *model.AiChannel, models []oci.GenAiModel) (string, string) {
status, detail := "error", "无可试调对话模型" status, detail := "error", "无可试调对话模型"
errBudget := 3
for _, m := range probeCandidates(models) { for _, m := range probeCandidates(models) {
code, err := s.client.GenAiProbeChat(ctx, cred, ch.Region, m.Ocid, m.Name) code, err := s.client.GenAiProbeChat(ctx, cred, ch.Region, m.Ocid, m.Name)
switch { switch {
case code == 200 || code == 429: case code == 200 || code == 429:
return "ok", "" return "ok", ""
case oci.IsModelUnavailable(err):
status, detail = "error", truncateErr(fmt.Sprintf("%s: 不可按需调用,建议加入模型黑名单", m.Name))
case code == 401 || code == 403 || code == 404: case code == 401 || code == 403 || code == 404:
return "no_quota", truncateErr(oci.CompactError(err)) return "no_quota", truncateErr(oci.CompactError(err))
default: default:
status, detail = "error", truncateErr(fmt.Sprintf("%s: %s", m.Name, oci.CompactError(err))) status, detail = "error", truncateErr(fmt.Sprintf("%s: %s", m.Name, oci.CompactError(err)))
if errBudget--; errBudget == 0 {
return status, detail
}
} }
} }
return status, detail return status, detail
} }
// probeCandidates 只取对话模型并按可靠度排序取前 3:主流文本模型优先, // probeCandidateCap 是单次探测的候选上限:放宽到 8 让一次探测有机会越过整批
// voice 等非常规形态殿后(embedding / rerank 已被能力筛选排除) // 不可按需调用的坏模型找到可用者;其他错误另有 3 次止损预算
const probeCandidateCap = 8
// probeCandidates 只取对话模型,按可靠度排序后跨厂商取候选:主流文本模型优先;
// voice 等负分形态(元数据标 CHAT 但实际不可对话)直接排除,不浪费试调预算;
// 每厂商先取最高分再按分数补位——部分区域某厂商全为微调基座(调用必失败),
// 不能让单一厂商占满候选名额拖垮整个渠道的探测结论。
func probeCandidates(models []oci.GenAiModel) []oci.GenAiModel { func probeCandidates(models []oci.GenAiModel) []oci.GenAiModel {
var sorted []oci.GenAiModel var sorted []oci.GenAiModel
for _, m := range models { for _, m := range models {
if m.Capability == "" || m.Capability == "CHAT" { if (m.Capability == "" || m.Capability == "CHAT") && probeScore(m.Name) >= 0 {
sorted = append(sorted, m) sorted = append(sorted, m)
} }
} }
sort.SliceStable(sorted, func(i, j int) bool { sort.SliceStable(sorted, func(i, j int) bool {
return probeScore(sorted[i].Name) > probeScore(sorted[j].Name) return probeScore(sorted[i].Name) > probeScore(sorted[j].Name)
}) })
if len(sorted) > 3 { return diversifyByVendor(sorted, probeCandidateCap)
sorted = sorted[:3] }
// diversifyByVendor 从已排序列表先每厂商各取一个,不足 limit 再按原序补位。
func diversifyByVendor(sorted []oci.GenAiModel, limit int) []oci.GenAiModel {
picked := make([]oci.GenAiModel, 0, limit)
used := make(map[int]bool)
seenVendor := make(map[string]bool)
for i, m := range sorted {
if len(picked) >= limit {
break
} }
return sorted if v := modelVendor(m); !seenVendor[v] {
seenVendor[v] = true
used[i] = true
picked = append(picked, m)
}
}
for i, m := range sorted {
if len(picked) >= limit {
break
}
if !used[i] {
picked = append(picked, m)
}
}
return picked
}
// modelVendor 取厂商标识;OCI 未回填 vendor 时退化为模型名「.」前缀。
func modelVendor(m oci.GenAiModel) string {
if m.Vendor != "" {
return strings.ToLower(m.Vendor)
}
name := strings.ToLower(m.Name)
if i := strings.Index(name, "."); i > 0 {
return name[:i]
}
return name
} }
func probeScore(name string) int { func probeScore(name string) int {
@@ -377,7 +454,7 @@ func truncateErr(msg string) string {
return msg return msg
} }
// SyncModels 重新拉取渠道区域的模型列表并覆盖缓存。 // SyncModels 重新拉取渠道区域的模型列表并覆盖缓存,黑名单中的模型不入库
func (s *AiGatewayService) SyncModels(ctx context.Context, id uint) ([]model.AiModelCache, error) { func (s *AiGatewayService) SyncModels(ctx context.Context, id uint) ([]model.AiModelCache, error) {
var ch model.AiChannel var ch model.AiChannel
if err := s.db.WithContext(ctx).First(&ch, id).Error; err != nil { if err := s.db.WithContext(ctx).First(&ch, id).Error; err != nil {
@@ -391,6 +468,10 @@ func (s *AiGatewayService) SyncModels(ctx context.Context, id uint) ([]model.AiM
if err != nil { if err != nil {
return nil, fmt.Errorf("同步模型失败:%s", oci.CompactError(err)) return nil, fmt.Errorf("同步模型失败:%s", oci.CompactError(err))
} }
models = supportedGatewayModels(models)
if models, err = s.withoutBlacklisted(ctx, models); err != nil {
return nil, err
}
if err := s.replaceModels(ctx, id, models); err != nil { if err := s.replaceModels(ctx, id, models); err != nil {
return nil, err return nil, err
} }
@@ -422,7 +503,7 @@ func (s *AiGatewayService) channelModels(ctx context.Context, channelID uint) ([
return rows, err return rows, err
} }
// GatewayModels 聚合启用渠道的模型(按名称去重),供 /ai/v1/models; // GatewayModels 聚合启用渠道的可用模型(按名称去重),供 /ai/v1/models;
// group 非空时仅聚合该分组渠道(与密钥分组路由口径一致)。 // group 非空时仅聚合该分组渠道(与密钥分组路由口径一致)。
func (s *AiGatewayService) GatewayModels(ctx context.Context, group string) (aiwire.ModelList, error) { func (s *AiGatewayService) GatewayModels(ctx context.Context, group string) (aiwire.ModelList, error) {
q := s.db.WithContext(ctx). q := s.db.WithContext(ctx).
@@ -502,6 +583,99 @@ func (s *AiGatewayService) ProbeAll(ctx context.Context) (string, error) {
return msg, nil return msg, nil
} }
// ---- 模型黑名单 ----
// Blacklist 列出全部黑名单模型(按名称排序)。
func (s *AiGatewayService) Blacklist(ctx context.Context) ([]model.AiModelBlacklist, error) {
var rows []model.AiModelBlacklist
err := s.db.WithContext(ctx).Order("name ASC").Find(&rows).Error
return rows, err
}
// AddBlacklist 把模型名加入黑名单并删除全部渠道缓存中的同名条目;
// 该模型此后同步 / 探测均被过滤,直到移出黑名单后重新同步。
func (s *AiGatewayService) AddBlacklist(ctx context.Context, name string) (*model.AiModelBlacklist, error) {
name = strings.TrimSpace(name)
if name == "" {
return nil, fmt.Errorf("模型名不能为空")
}
row := model.AiModelBlacklist{Name: name}
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var n int64
if err := tx.Model(&model.AiModelBlacklist{}).Where("name = ?", name).Count(&n).Error; err != nil {
return err
}
if n > 0 {
return fmt.Errorf("模型已在黑名单中")
}
if err := tx.Create(&row).Error; err != nil {
return err
}
return tx.Where("name = ?", name).Delete(&model.AiModelCache{}).Error
})
if err != nil {
return nil, err
}
return &row, nil
}
// RemoveBlacklist 把模型移出黑名单;缓存不回填,下次同步 / 探测自然恢复入池。
func (s *AiGatewayService) RemoveBlacklist(ctx context.Context, id uint) error {
res := s.db.WithContext(ctx).Delete(&model.AiModelBlacklist{}, id)
if res.Error != nil {
return res.Error
}
if res.RowsAffected == 0 {
return fmt.Errorf("黑名单条目不存在")
}
return nil
}
// withoutBlacklisted 过滤掉黑名单中的模型,同步与探测入库前统一经此收口。
// supportedGatewayModels 过滤模型目录:对话模型仅保留实测支持 OpenAI 兼容面的
// 厂商(xai / meta / openai)——typed chat 面已剔除,google / cohere 对话模型无
// 上游通路,不入目录(不出现在列表、路由与探测候选);EMBEDDING 模型不受影响。
func supportedGatewayModels(models []oci.GenAiModel) []oci.GenAiModel {
out := make([]oci.GenAiModel, 0, len(models))
for _, m := range models {
if m.Capability == "CHAT" && !compatChatVendor(m.Name) {
continue
}
out = append(out, m)
}
return out
}
func compatChatVendor(model string) bool {
for _, prefix := range []string{"xai.", "meta.", "openai."} {
if strings.HasPrefix(model, prefix) {
return true
}
}
return false
}
func (s *AiGatewayService) withoutBlacklisted(ctx context.Context, models []oci.GenAiModel) ([]oci.GenAiModel, error) {
var names []string
if err := s.db.WithContext(ctx).Model(&model.AiModelBlacklist{}).Pluck("name", &names).Error; err != nil {
return nil, fmt.Errorf("读取模型黑名单: %w", err)
}
if len(names) == 0 {
return models, nil
}
banned := make(map[string]bool, len(names))
for _, n := range names {
banned[n] = true
}
out := make([]oci.GenAiModel, 0, len(models))
for _, m := range models {
if !banned[m.Name] {
out = append(out, m)
}
}
return out, nil
}
// ---- 调用日志 ---- // ---- 调用日志 ----
// LogCall 落一条调用日志(仅元数据与用量,永不含请求 / 响应正文),返回落库 ID 供内容日志关联(失败为 0)。 // LogCall 落一条调用日志(仅元数据与用量,永不含请求 / 响应正文),返回落库 ID 供内容日志关联(失败为 0)。
+37 -21
View File
@@ -3,6 +3,7 @@ package service
import ( import (
"context" "context"
"errors" "errors"
"io"
"math/rand" "math/rand"
"time" "time"
@@ -26,27 +27,31 @@ type aiCandidate struct {
modelOcid string modelOcid string
} }
// Chat 编排非流式调用:选渠道 → 调用 → 可重试错误换渠道(整请求上限 3 次)。 // RespPassthrough 编排一次非流式直通调用:选渠道(priority→加权随机)→ 调用 →
// group 非空时只在同分组渠道内路由(取自调用密钥) // 可重试错误换渠道(整请求上限 3 次)并维护熔断;group 非空时只在同分组渠道内路由。
func (s *AiGatewayService) Chat(ctx context.Context, ir aiwire.ChatRequest, group string) (*aiwire.ChatResponse, ChatMeta, error) { // 上游为 OpenAI-compatible /actions/v1/responses(实测可用,无 Oracle 文档合同)。
func (s *AiGatewayService) RespPassthrough(ctx context.Context, raw []byte, modelName, group string) ([]byte, ChatMeta, error) {
meta := ChatMeta{} meta := ChatMeta{}
excluded := map[uint]bool{} excluded := map[uint]bool{}
var lastErr error var lastErr error
for attempt := 0; attempt < 3; attempt++ { for attempt := 0; attempt < 3; attempt++ {
cand, err := s.pick(ctx, ir.Model, group, "CHAT", excluded) cand, err := s.pick(ctx, modelName, group, "CHAT", excluded)
if err != nil { if err != nil {
return nil, meta, firstErr(lastErr, err) return nil, meta, firstErr(lastErr, err)
} }
meta.ChannelID, meta.ChannelName = cand.ch.ID, cand.ch.Name meta.ChannelID, meta.ChannelName = cand.ch.ID, cand.ch.Name
resp, err := s.callOnce(ctx, cand, ir) payload, err := s.passthroughOnce(ctx, cand, raw)
if err == nil { if err == nil {
s.markSuccess(ctx, cand.ch.ID) s.markSuccess(ctx, cand.ch.ID)
return resp, meta, nil return payload, meta, nil
} }
if !retryable(err) { retry, penalize := switchable(err)
if !retry {
return nil, meta, err return nil, meta, err
} }
if penalize {
s.markFailure(ctx, cand.ch.ID) s.markFailure(ctx, cand.ch.ID)
}
excluded[cand.ch.ID] = true excluded[cand.ch.ID] = true
meta.Retries++ meta.Retries++
lastErr = err lastErr = err
@@ -54,22 +59,22 @@ func (s *AiGatewayService) Chat(ctx context.Context, ir aiwire.ChatRequest, grou
return nil, meta, lastErr return nil, meta, lastErr
} }
func (s *AiGatewayService) callOnce(ctx context.Context, cand *aiCandidate, ir aiwire.ChatRequest) (*aiwire.ChatResponse, error) { func (s *AiGatewayService) passthroughOnce(ctx context.Context, cand *aiCandidate, raw []byte) ([]byte, error) {
cred, err := s.configs.credentialsByID(ctx, cand.ch.OciConfigID) cred, err := s.configs.credentialsByID(ctx, cand.ch.OciConfigID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
return s.client.GenAiChat(ctx, cred, cand.ch.Region, cand.modelOcid, ir) return s.client.GenAiCompatResponses(ctx, cred, cand.ch.Region, raw)
} }
// OpenStream 编排流式调用:流建立成功即绑定渠道,建立失败可换渠道重试。 // RespPassthroughStream 编排流式直通:流建立成功即绑定渠道,建立失败按 switchable
// group 语义与 Chat 相同 // 换渠道重试;建立后的中断不重试、不计熔断(与 OpenStream 语义一致)
func (s *AiGatewayService) OpenStream(ctx context.Context, ir aiwire.ChatRequest, group string) (oci.GenAiStream, ChatMeta, error) { func (s *AiGatewayService) RespPassthroughStream(ctx context.Context, raw []byte, modelName, group string) (io.ReadCloser, ChatMeta, error) {
meta := ChatMeta{} meta := ChatMeta{}
excluded := map[uint]bool{} excluded := map[uint]bool{}
var lastErr error var lastErr error
for attempt := 0; attempt < 3; attempt++ { for attempt := 0; attempt < 3; attempt++ {
cand, err := s.pick(ctx, ir.Model, group, "CHAT", excluded) cand, err := s.pick(ctx, modelName, group, "CHAT", excluded)
if err != nil { if err != nil {
return nil, meta, firstErr(lastErr, err) return nil, meta, firstErr(lastErr, err)
} }
@@ -78,15 +83,18 @@ func (s *AiGatewayService) OpenStream(ctx context.Context, ir aiwire.ChatRequest
if err != nil { if err != nil {
return nil, meta, err return nil, meta, err
} }
stream, err := s.client.GenAiChatStream(ctx, cred, cand.ch.Region, cand.modelOcid, ir) stream, err := s.client.GenAiCompatResponsesStream(ctx, cred, cand.ch.Region, raw)
if err == nil { if err == nil {
s.markSuccess(ctx, cand.ch.ID) s.markSuccess(ctx, cand.ch.ID)
return stream, meta, nil return stream, meta, nil
} }
if !retryable(err) { retry, penalize := switchable(err)
if !retry {
return nil, meta, err return nil, meta, err
} }
if penalize {
s.markFailure(ctx, cand.ch.ID) s.markFailure(ctx, cand.ch.ID)
}
excluded[cand.ch.ID] = true excluded[cand.ch.ID] = true
meta.Retries++ meta.Retries++
lastErr = err lastErr = err
@@ -102,12 +110,17 @@ func firstErr(lastErr, pickErr error) error {
return pickErr return pickErr
} }
// retryable 判是否换渠道重试:429 / 5xx / 网络错误可重试,其余 4xx 直接透传。 // switchable 判断调用失败是否换渠道重试、是否计入熔断:模型级不可用(微调基座
func retryable(err error) bool { // 400 / 实体不存在 404)换渠道但不计熔断——这是模型×区域供给问题而非渠道健康
if status, ok := oci.ServiceStatus(err); ok { // 问题,持久解决靠加入模型黑名单;429 / 5xx / 网络错误换渠道并计熔断;其余 4xx 直接透传。
return status == 429 || status >= 500 func switchable(err error) (retry, penalize bool) {
if oci.IsModelUnavailable(err) {
return true, false
} }
return true if status, ok := oci.ServiceStatus(err); ok {
return status == 429 || status >= 500, true
}
return true, true
} }
// pick 选出支持该模型的最优渠道:能力匹配 → 启用 → 分组匹配 → 未熔断 → 最小优先级组 → 加权随机。 // pick 选出支持该模型的最优渠道:能力匹配 → 启用 → 分组匹配 → 未熔断 → 最小优先级组 → 加权随机。
@@ -221,10 +234,13 @@ func (s *AiGatewayService) Embeddings(ctx context.Context, req aiwire.Embeddings
s.markSuccess(ctx, cand.ch.ID) s.markSuccess(ctx, cand.ch.ID)
return resp, meta, nil return resp, meta, nil
} }
if !retryable(err) { retry, penalize := switchable(err)
if !retry {
return nil, meta, err return nil, meta, err
} }
if penalize {
s.markFailure(ctx, cand.ch.ID) s.markFailure(ctx, cand.ch.ID)
}
excluded[cand.ch.ID] = true excluded[cand.ch.ID] = true
meta.Retries++ meta.Retries++
lastErr = err lastErr = err
+29
View File
@@ -0,0 +1,29 @@
package service
import (
"strings"
"testing"
"oci-portal/internal/oci"
)
// TestSupportedGatewayModels 断言 CHAT 目录仅保留兼容面厂商,EMBEDDING 不受影响。
func TestSupportedGatewayModels(t *testing.T) {
in := []oci.GenAiModel{
{Name: "xai.grok-4.3", Capability: "CHAT"},
{Name: "meta.llama-4-maverick-17b-128e", Capability: "CHAT"},
{Name: "openai.gpt-oss-120b", Capability: "CHAT"},
{Name: "google.gemini-2.5-flash", Capability: "CHAT"},
{Name: "cohere.command-a-03-2025", Capability: "CHAT"},
{Name: "cohere.embed-v4.0", Capability: "EMBEDDING"},
}
out := supportedGatewayModels(in)
names := make([]string, 0, len(out))
for _, m := range out {
names = append(names, m.Name)
}
want := "xai.grok-4.3,meta.llama-4-maverick-17b-128e,openai.gpt-oss-120b,cohere.embed-v4.0"
if got := strings.Join(names, ","); got != want {
t.Fatalf("过滤结果 = %s, want %s", got, want)
}
}
+418 -74
View File
@@ -1,8 +1,11 @@
package service package service
import ( import (
"bytes"
"context" "context"
"errors" "errors"
"fmt"
"io"
"strings" "strings"
"testing" "testing"
"time" "time"
@@ -13,14 +16,28 @@ import (
) )
// stubServiceError 实现 common.ServiceError,用于模拟带状态码的 OCI 服务端错误。 // stubServiceError 实现 common.ServiceError,用于模拟带状态码的 OCI 服务端错误。
type stubServiceError struct{ status int } type stubServiceError struct {
status int
msg string
}
func (e stubServiceError) Error() string { return "stub service error" } func (e stubServiceError) Error() string { return "stub service error" }
func (e stubServiceError) GetHTTPStatusCode() int { return e.status } func (e stubServiceError) GetHTTPStatusCode() int { return e.status }
func (e stubServiceError) GetMessage() string { return "stub" } func (e stubServiceError) GetMessage() string {
if e.msg != "" {
return e.msg
}
return "stub"
}
func (e stubServiceError) GetCode() string { return "Stub" } func (e stubServiceError) GetCode() string { return "Stub" }
func (e stubServiceError) GetOpcRequestID() string { return "req-1" } func (e stubServiceError) GetOpcRequestID() string { return "req-1" }
// finetuneBaseErr 模拟「微调基座模型不可按需调用」的 OCI 400。
func finetuneBaseErr() stubServiceError {
return stubServiceError{status: 400,
msg: "Not allowed to call finetune base model ocid1.generativeaimodel.oc1.eu-frankfurt-1.tpel5q, use Endpoint: false"}
}
// gatewayStubClient 覆写 GenAI 四方法;chatErrs 逐次弹出以模拟先失败后成功。 // gatewayStubClient 覆写 GenAI 四方法;chatErrs 逐次弹出以模拟先失败后成功。
type gatewayStubClient struct { type gatewayStubClient struct {
*fakeClient *fakeClient
@@ -29,14 +46,47 @@ type gatewayStubClient struct {
modelsErr error modelsErr error
probeCode int probeCode int
probeErr error probeErr error
chatResp *aiwire.ChatResponse // probeSeq 非空时逐次弹出,模拟按候选依次试调;弹尽后回落 probeCode/probeErr
chatErrs []error probeSeq []probeResult
// probedNames 记录试调过的模型名,供断言候选过滤
probedNames []string
chatCalls int chatCalls int
regions []string regions []string
embedVecs [][]float32 embedVecs [][]float32
embedUsage *aiwire.Usage embedUsage *aiwire.Usage
embedErr error embedErr error
passPayload []byte
passErrs []error
passCalls int
passRegions []string
}
func (f *gatewayStubClient) GenAiCompatResponses(ctx context.Context, cred oci.Credentials, region string, body []byte) ([]byte, error) {
f.passCalls++
f.passRegions = append(f.passRegions, region)
if len(f.passErrs) > 0 {
err := f.passErrs[0]
f.passErrs = f.passErrs[1:]
if err != nil {
return nil, err
}
}
return f.passPayload, nil
}
func (f *gatewayStubClient) GenAiCompatResponsesStream(ctx context.Context, cred oci.Credentials, region string, body []byte) (io.ReadCloser, error) {
f.passCalls++
f.passRegions = append(f.passRegions, region)
if len(f.passErrs) > 0 {
err := f.passErrs[0]
f.passErrs = f.passErrs[1:]
if err != nil {
return nil, err
}
}
return io.NopCloser(bytes.NewReader(f.passPayload)), nil
} }
func (f *gatewayStubClient) GenAiEmbed(ctx context.Context, cred oci.Credentials, region, modelOcid string, inputs []string, dimensions *int) ([][]float32, *aiwire.Usage, error) { func (f *gatewayStubClient) GenAiEmbed(ctx context.Context, cred oci.Credentials, region, modelOcid string, inputs []string, dimensions *int) ([][]float32, *aiwire.Usage, error) {
@@ -47,27 +97,26 @@ func (f *gatewayStubClient) ListGenAiModels(ctx context.Context, cred oci.Creden
return f.models, f.modelsErr return f.models, f.modelsErr
} }
func (f *gatewayStubClient) GenAiProbeChat(ctx context.Context, cred oci.Credentials, region, modelOcid, modelName string) (int, error) { // probeResult 是 gatewayStubClient.probeSeq 的单次探测结果。
return f.probeCode, f.probeErr type probeResult struct {
code int
err error
} }
func (f *gatewayStubClient) GenAiChat(ctx context.Context, cred oci.Credentials, region, modelOcid string, ir aiwire.ChatRequest) (*aiwire.ChatResponse, error) { func (f *gatewayStubClient) GenAiProbeChat(ctx context.Context, cred oci.Credentials, region, modelOcid, modelName string) (int, error) {
f.chatCalls++ f.probedNames = append(f.probedNames, modelName)
f.regions = append(f.regions, region) if len(f.probeSeq) > 0 {
if len(f.chatErrs) > 0 { r := f.probeSeq[0]
err := f.chatErrs[0] f.probeSeq = f.probeSeq[1:]
f.chatErrs = f.chatErrs[1:] return r.code, r.err
if err != nil {
return nil, err
} }
} return f.probeCode, f.probeErr
return f.chatResp, nil
} }
func newTestGateway(t *testing.T, client oci.Client) (*AiGatewayService, *OciConfigService) { func newTestGateway(t *testing.T, client oci.Client) (*AiGatewayService, *OciConfigService) {
t.Helper() t.Helper()
svc := newTestService(t, client) svc := newTestService(t, client)
if err := svc.db.AutoMigrate(&model.AiKey{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}); err != nil { if err := svc.db.AutoMigrate(&model.AiKey{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiModelBlacklist{}, &model.AiCallLog{}); err != nil {
t.Fatalf("auto migrate ai tables: %v", err) t.Fatalf("auto migrate ai tables: %v", err)
} }
return NewAiGatewayService(svc.db, svc, client), svc return NewAiGatewayService(svc.db, svc, client), svc
@@ -77,14 +126,14 @@ func TestAiKeyLifecycle(t *testing.T) {
gw, _ := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}}) gw, _ := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}})
ctx := context.Background() ctx := context.Background()
raw, key, err := gw.CreateKey(ctx, "免费-ai-api-key", "test-api-key", "") raw, key, err := gw.CreateKey(ctx, "免费-ai-api-key", "test-api-key", "", nil)
if err != nil || raw != "test-api-key" || key.Tail != "-key" { if err != nil || raw != "test-api-key" || key.Tail != "-key" {
t.Fatalf("CreateKey 自定义值 = %q %+v, %v", raw, key, err) t.Fatalf("CreateKey 自定义值 = %q %+v, %v", raw, key, err)
} }
if _, _, err := gw.CreateKey(ctx, "short", "abc", ""); err == nil { if _, _, err := gw.CreateKey(ctx, "short", "abc", "", nil); err == nil {
t.Error("过短自定义密钥应被拒绝") t.Error("过短自定义密钥应被拒绝")
} }
auto, _, err := gw.CreateKey(ctx, "auto", "", "") auto, _, err := gw.CreateKey(ctx, "auto", "", "", nil)
if err != nil || len(auto) < 40 || auto[:3] != "sk-" { if err != nil || len(auto) < 40 || auto[:3] != "sk-" {
t.Fatalf("CreateKey 随机值 = %q, %v", auto, err) t.Fatalf("CreateKey 随机值 = %q, %v", auto, err)
} }
@@ -96,7 +145,7 @@ func TestAiKeyLifecycle(t *testing.T) {
t.Errorf("错误密钥 err = %v", err) t.Errorf("错误密钥 err = %v", err)
} }
off := false off := false
_ = gw.UpdateKey(ctx, key.ID, "", &off, nil) _ = gw.UpdateKey(ctx, key.ID, "", &off, nil, nil)
if _, err := gw.VerifyKey(ctx, "test-api-key"); !errors.Is(err, ErrAiKeyInvalid) { if _, err := gw.VerifyKey(ctx, "test-api-key"); !errors.Is(err, ErrAiKeyInvalid) {
t.Errorf("禁用后 VerifyKey err = %v", err) t.Errorf("禁用后 VerifyKey err = %v", err)
} }
@@ -159,51 +208,6 @@ func seedChannel(t *testing.T, gw *AiGatewayService, cfgID uint, region string,
return ch return ch
} }
func TestAiChatRetrySwitchesChannel(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
chatResp: &aiwire.ChatResponse{Model: "meta.llama-3.3-70b-instruct", Choices: []aiwire.Choice{{Message: aiwire.ChatMessage{Role: "assistant", Content: aiwire.NewTextContent("hi")}, FinishReason: "stop"}}},
chatErrs: []error{stubServiceError{status: 429}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
// 两个同优先级渠道
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
ctx := context.Background()
resp, meta, err := gw.Chat(ctx, aiwire.ChatRequest{Model: "meta.llama-3.3-70b-instruct", Messages: []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent("你好")}}}, "")
if err != nil || resp == nil {
t.Fatalf("Chat = %v, %v", resp, err)
}
if meta.Retries != 1 || client.chatCalls != 2 {
t.Errorf("应换渠道重试一次: retries=%d calls=%d", meta.Retries, client.chatCalls)
}
if len(client.regions) != 2 && client.regions[0] == client.regions[1] {
t.Errorf("重试未换渠道: %v", client.regions)
}
// 未知模型
if _, _, err := gw.Chat(ctx, aiwire.ChatRequest{Model: "no-such-model"}, ""); !errors.Is(err, ErrAiUnknownModel) {
t.Errorf("未知模型 err = %v", err)
}
}
func TestAiChatNonRetryablePassThrough(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
chatErrs: []error{stubServiceError{status: 400}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
_, meta, err := gw.Chat(context.Background(), aiwire.ChatRequest{Model: "meta.llama-3.3-70b-instruct"}, "")
if err == nil || meta.Retries != 0 || client.chatCalls != 1 {
t.Errorf("400 应直接透传不重试: err=%v retries=%d calls=%d", err, meta.Retries, client.chatCalls)
}
}
func TestPickPriorityAndBreaker(t *testing.T) { func TestPickPriorityAndBreaker(t *testing.T) {
gw, svc := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}}) gw, svc := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}})
cfg := importAliveConfig(t, svc) cfg := importAliveConfig(t, svc)
@@ -255,7 +259,7 @@ func TestMarkFailureBackoff(t *testing.T) {
func TestAiGroupRouting(t *testing.T) { func TestAiGroupRouting(t *testing.T) {
client := &gatewayStubClient{ client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}, fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
chatResp: &aiwire.ChatResponse{Model: "meta.llama-3.3-70b-instruct", Choices: []aiwire.Choice{{Message: aiwire.ChatMessage{Role: "assistant", Content: aiwire.NewTextContent("hi")}, FinishReason: "stop"}}}, passPayload: []byte(`{"id":"resp_1","usage":{"input_tokens":3,"output_tokens":1}}`),
} }
gw, svc := newTestGateway(t, client) gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc) cfg := importAliveConfig(t, svc)
@@ -264,17 +268,17 @@ func TestAiGroupRouting(t *testing.T) {
gw.db.Model(vip).Update("channel_group", "vip") gw.db.Model(vip).Update("channel_group", "vip")
gw.db.Create(&model.AiModelCache{ChannelID: other.ID, ModelOcid: "ocid1..cr", Name: "cohere.command-r", Vendor: "cohere", SyncedAt: time.Now()}) gw.db.Create(&model.AiModelCache{ChannelID: other.ID, ModelOcid: "ocid1..cr", Name: "cohere.command-r", Vendor: "cohere", SyncedAt: time.Now()})
ctx := context.Background() ctx := context.Background()
req := aiwire.ChatRequest{Model: "meta.llama-3.3-70b-instruct", Messages: []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent("你好")}}} body := []byte(`{"model":"meta.llama-3.3-70b-instruct","input":"你好"}`)
// 分组密钥只落同分组渠道 // 分组密钥只落同分组渠道
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
_, meta, err := gw.Chat(ctx, req, "vip") _, meta, err := gw.RespPassthrough(ctx, body, "meta.llama-3.3-70b-instruct", "vip")
if err != nil || meta.ChannelID != vip.ID { if err != nil || meta.ChannelID != vip.ID {
t.Fatalf("vip 组第 %d 次落点 = %d, err %v, want %d", i, meta.ChannelID, err, vip.ID) t.Fatalf("vip 组第 %d 次落点 = %d, err %v, want %d", i, meta.ChannelID, err, vip.ID)
} }
} }
// 分组内无渠道 → 无可用渠道 // 分组内无渠道 → 无可用渠道
if _, _, err := gw.Chat(ctx, req, "nope"); !errors.Is(err, ErrAiNoChannel) { if _, _, err := gw.RespPassthrough(ctx, body, "meta.llama-3.3-70b-instruct", "nope"); !errors.Is(err, ErrAiNoChannel) {
t.Errorf("空分组 err = %v, want ErrAiNoChannel", err) t.Errorf("空分组 err = %v, want ErrAiNoChannel", err)
} }
// 模型列表按分组过滤 // 模型列表按分组过滤
@@ -287,12 +291,12 @@ func TestAiGroupRouting(t *testing.T) {
t.Errorf("不限组模型数 = %d, want 2", len(all.Data)) t.Errorf("不限组模型数 = %d, want 2", len(all.Data))
} }
// 密钥分组落库 // 密钥分组落库
_, key, err := gw.CreateKey(ctx, "vip-key", "vip-secret-1234", "vip") _, key, err := gw.CreateKey(ctx, "vip-key", "vip-secret-1234", "vip", nil)
if err != nil || key.Group != "vip" { if err != nil || key.Group != "vip" {
t.Fatalf("CreateKey group = %+v, %v", key, err) t.Fatalf("CreateKey group = %+v, %v", key, err)
} }
empty := "" empty := ""
_ = gw.UpdateKey(ctx, key.ID, "", nil, &empty) _ = gw.UpdateKey(ctx, key.ID, "", nil, &empty, nil)
var fresh model.AiKey var fresh model.AiKey
gw.db.First(&fresh, key.ID) gw.db.First(&fresh, key.ID)
if fresh.Group != "" { if fresh.Group != "" {
@@ -400,6 +404,218 @@ func TestProbeCandidates(t *testing.T) {
} }
} }
func TestProbeCandidatesVendorDiversity(t *testing.T) {
// 部分区域单一厂商全为微调基座:候选须跨厂商分散,不能被 3 个 llama 占满前排
models := []oci.GenAiModel{
{Ocid: "l1", Name: "meta.llama-3-70b-instruct"},
{Ocid: "l2", Name: "meta.llama-3.1-405b-instruct"},
{Ocid: "l3", Name: "meta.llama-3.3-70b-instruct"},
{Ocid: "c1", Name: "cohere.command-a-03-2025"},
{Ocid: "g1", Name: "xai.grok-4"},
}
got := probeCandidates(models)
if len(got) != 5 || got[0].Name != "meta.llama-3-70b-instruct" {
t.Fatalf("上限内全量返回且最高分居首: %+v", got)
}
vendors := map[string]bool{}
for _, m := range got[:3] {
vendors[modelVendor(m)] = true
}
if len(vendors) != 3 {
t.Errorf("前 3 个候选应覆盖 3 个厂商: %+v", got)
}
// 超过上限时截断到 probeCandidateCap
var many []oci.GenAiModel
for i := 0; i < 12; i++ {
many = append(many, oci.GenAiModel{Ocid: fmt.Sprintf("m%d", i), Name: fmt.Sprintf("meta.llama-%d", i)})
}
if capped := probeCandidates(many); len(capped) != probeCandidateCap {
t.Errorf("候选应截断到 %d: got %d", probeCandidateCap, len(capped))
}
}
// entityNotFoundErr 模拟「实体不存在」404(模型在区域内无按需供给)。
func entityNotFoundErr() stubServiceError {
return stubServiceError{status: 404,
msg: "Entity with key ocid1.generativeaimodel.oc1.eu-frankfurt-1.2flsfq not found"}
}
func TestProbeSkipsUnavailableModels(t *testing.T) {
// 首选 llama 不可按需调用(基座 400 / 实体 404)→ 换候选成功 → 渠道判可用;
// 模型全量入库不再自动标记,持久剔除交由用户手动拉黑
tests := []struct {
name string
bad probeResult
}{
{"微调基座 400", probeResult{400, finetuneBaseErr()}},
{"实体不存在 404", probeResult{404, entityNotFoundErr()}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
models: []oci.GenAiModel{
{Ocid: "m1", Name: "meta.llama-3-70b-instruct", Vendor: "meta"},
{Ocid: "m2", Name: "meta.llama-3.1-70b-instruct", Vendor: "meta"},
{Ocid: "m3", Name: "cohere.command-a-03-2025", Vendor: "cohere"},
},
probeSeq: []probeResult{tt.bad, {200, nil}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
ctx := context.Background()
ch, err := gw.CreateChannel(ctx, ChannelInput{OciConfigID: cfg.ID, Region: "eu-frankfurt-1"})
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
probed, err := gw.ProbeChannel(ctx, ch.ID)
if err != nil || probed.ProbeStatus != "ok" {
t.Fatalf("坏候选后应换候选并判可用: %+v, %v", probed, err)
}
rows, _ := gw.channelModels(ctx, ch.ID)
if len(rows) != 3 {
t.Errorf("模型应全量入库, got %d", len(rows))
}
})
}
}
func TestProbeAuth404StillNoQuota(t *testing.T) {
// 鉴权类 404(NotAuthorizedOrNotFound)仍属租户级,直接定论 no_quota
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
models: []oci.GenAiModel{{Ocid: "m1", Name: "meta.llama-3.3-70b-instruct", Vendor: "meta"}},
probeCode: 404,
probeErr: stubServiceError{status: 404, msg: "Authorization failed or requested resource not found."},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
ctx := context.Background()
ch, _ := gw.CreateChannel(ctx, ChannelInput{OciConfigID: cfg.ID, Region: "eu-frankfurt-1"})
probed, _ := gw.ProbeChannel(ctx, ch.ID)
if probed.ProbeStatus != "no_quota" {
t.Errorf("鉴权 404 status = %q, want no_quota", probed.ProbeStatus)
}
}
func TestAiChatFinetuneSwitchesChannelWithoutPenalty(t *testing.T) {
// 微调基座 400 换渠道重试成功,且不计入熔断失败
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
passPayload: []byte(`{"id":"resp_1","usage":{"input_tokens":3,"output_tokens":1}}`),
passErrs: []error{finetuneBaseErr()},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
resp, meta, err := gw.RespPassthrough(context.Background(), []byte(`{}`), "meta.llama-3.3-70b-instruct", "")
if err != nil || resp == nil {
t.Fatalf("RespPassthrough = %v, %v", resp, err)
}
if meta.Retries != 1 || client.passCalls != 2 {
t.Errorf("应换渠道重试一次: retries=%d calls=%d", meta.Retries, client.passCalls)
}
var chs []model.AiChannel
gw.db.Find(&chs)
for _, ch := range chs {
if ch.FailCount != 0 {
t.Errorf("微调基座 400 不应计入熔断: 渠道 %s failCount=%d", ch.Name, ch.FailCount)
}
}
}
func TestBlacklistLifecycle(t *testing.T) {
// 拉黑:全渠道同名缓存删除、列表与路由立即不可见;重复/空名被拒;
// 移出黑名单后重新同步恢复入池
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
models: []oci.GenAiModel{{Ocid: "ocid1..m-eu-frankfurt-1", Name: "meta.llama-3.3-70b-instruct", Vendor: "meta"}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
ch := seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
ctx := context.Background()
item, err := gw.AddBlacklist(ctx, " meta.llama-3.3-70b-instruct ")
if err != nil || item.Name != "meta.llama-3.3-70b-instruct" {
t.Fatalf("AddBlacklist = %+v, %v", item, err)
}
var left int64
gw.db.Model(&model.AiModelCache{}).Count(&left)
if left != 0 {
t.Errorf("拉黑应删除全部渠道的同名缓存, 剩 %d", left)
}
if _, _, err := gw.RespPassthrough(ctx, []byte(`{}`), "meta.llama-3.3-70b-instruct", ""); !errors.Is(err, ErrAiUnknownModel) {
t.Errorf("拉黑后路由应按未知模型拒绝: %v", err)
}
if _, err := gw.AddBlacklist(ctx, "meta.llama-3.3-70b-instruct"); err == nil {
t.Error("重复拉黑应被拒绝")
}
if _, err := gw.AddBlacklist(ctx, " "); err == nil {
t.Error("空模型名应被拒绝")
}
rows, err := gw.Blacklist(ctx)
if err != nil || len(rows) != 1 {
t.Fatalf("Blacklist = %+v, %v", rows, err)
}
if err := gw.RemoveBlacklist(ctx, rows[0].ID); err != nil {
t.Fatalf("RemoveBlacklist: %v", err)
}
if err := gw.RemoveBlacklist(ctx, rows[0].ID); err == nil {
t.Error("移除不存在的条目应报错")
}
// 移出后重新同步:模型恢复入池
if _, err := gw.SyncModels(ctx, ch.ID); err != nil {
t.Fatalf("SyncModels: %v", err)
}
list, _ := gw.GatewayModels(ctx, "")
if len(list.Data) != 1 {
t.Errorf("移出黑名单并同步后应恢复入池: %+v", list.Data)
}
}
func TestSyncAndProbeFilterBlacklisted(t *testing.T) {
// 同步与探测都过滤黑名单模型:不入库、不进试调候选
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
models: []oci.GenAiModel{
{Ocid: "m1", Name: "meta.llama-3-70b-instruct", Vendor: "meta"},
{Ocid: "m2", Name: "cohere.command-a-03-2025", Vendor: "cohere"},
},
probeCode: 200,
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
ctx := context.Background()
if _, err := gw.AddBlacklist(ctx, "meta.llama-3-70b-instruct"); err != nil {
t.Fatalf("AddBlacklist: %v", err)
}
ch, err := gw.CreateChannel(ctx, ChannelInput{OciConfigID: cfg.ID, Region: "eu-frankfurt-1"})
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
probed, err := gw.ProbeChannel(ctx, ch.ID)
if err != nil || probed.ProbeStatus != "ok" {
t.Fatalf("ProbeChannel = %+v, %v", probed, err)
}
rows, _ := gw.channelModels(ctx, ch.ID)
if len(rows) != 1 || rows[0].Name != "cohere.command-a-03-2025" {
t.Errorf("探测同步应过滤黑名单模型: %+v", rows)
}
if len(client.probedNames) != 1 || client.probedNames[0] != "cohere.command-a-03-2025" {
t.Errorf("试调候选不应包含黑名单模型: %v", client.probedNames)
}
models, err := gw.SyncModels(ctx, ch.ID)
if err != nil || len(models) != 1 || models[0].Name != "cohere.command-a-03-2025" {
t.Errorf("SyncModels 应过滤黑名单模型: %+v, %v", models, err)
}
}
func TestDeprecatingModels(t *testing.T) { func TestDeprecatingModels(t *testing.T) {
gw, _ := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}}) gw, _ := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}})
ctx := context.Background() ctx := context.Background()
@@ -460,7 +676,7 @@ func TestAiEmbeddings(t *testing.T) {
t.Errorf("chat 模型走 embeddings err = %v, want ErrAiUnknownModel", err) t.Errorf("chat 模型走 embeddings err = %v, want ErrAiUnknownModel", err)
} }
// embedding 模型名打 chat:同样未知模型 // embedding 模型名打 chat:同样未知模型
if _, _, err := gw.Chat(ctx, aiwire.ChatRequest{Model: "cohere.embed-v4.0", Messages: []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent("x")}}}, ""); !errors.Is(err, ErrAiUnknownModel) { if _, _, err := gw.RespPassthrough(ctx, []byte(`{}`), "cohere.embed-v4.0", ""); !errors.Is(err, ErrAiUnknownModel) {
t.Errorf("embedding 模型走 chat err = %v, want ErrAiUnknownModel", err) t.Errorf("embedding 模型走 chat err = %v, want ErrAiUnknownModel", err)
} }
} }
@@ -471,7 +687,7 @@ func TestAiContentLogSwitch(t *testing.T) {
t.Fatalf("migrate content log: %v", err) t.Fatalf("migrate content log: %v", err)
} }
ctx := context.Background() ctx := context.Background()
_, key, err := gw.CreateKey(ctx, "k1", "content-key-1234", "") _, key, err := gw.CreateKey(ctx, "k1", "content-key-1234", "", nil)
if err != nil { if err != nil {
t.Fatalf("CreateKey: %v", err) t.Fatalf("CreateKey: %v", err)
} }
@@ -507,3 +723,131 @@ func TestAiContentLogSwitch(t *testing.T) {
t.Fatalf("关闭失败: %+v, %v", fresh.ContentLogUntil, err) t.Fatalf("关闭失败: %+v, %v", fresh.ContentLogUntil, err)
} }
} }
func TestRespPassthroughSwitchesChannel(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
passPayload: []byte(`{"id":"resp_1","usage":{"input_tokens":9,"output_tokens":3}}`),
passErrs: []error{stubServiceError{status: 429}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
payload, meta, err := gw.RespPassthrough(context.Background(), []byte(`{}`), "meta.llama-3.3-70b-instruct", "")
if err != nil || len(payload) == 0 {
t.Fatalf("RespPassthrough = %v, %v", payload, err)
}
if meta.Retries != 1 || client.passCalls != 2 {
t.Errorf("应换渠道重试一次: retries=%d calls=%d", meta.Retries, client.passCalls)
}
if len(client.passRegions) == 2 && client.passRegions[0] == client.passRegions[1] {
t.Errorf("重试未换渠道: %v", client.passRegions)
}
}
func TestRespPassthroughNonRetryable(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
passErrs: []error{stubServiceError{status: 400}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
_, meta, err := gw.RespPassthrough(context.Background(), []byte(`{}`), "meta.llama-3.3-70b-instruct", "")
if err == nil || meta.Retries != 0 || client.passCalls != 1 {
t.Errorf("400 应直接透传不重试: err=%v retries=%d calls=%d", err, meta.Retries, client.passCalls)
}
}
func TestNormalizeKeyModels(t *testing.T) {
tests := []struct {
name string
in []string
want []string
}{
{"nil 输入", nil, nil},
{"全空串", []string{"", " "}, nil},
{"trim 与保序去重", []string{" a ", "b", "a", "", "b "}, []string{"a", "b"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := normalizeKeyModels(tt.in)
if len(got) != len(tt.want) {
t.Fatalf("normalizeKeyModels(%v) = %v, want %v", tt.in, got, tt.want)
}
for i := range tt.want {
if got[i] != tt.want[i] {
t.Fatalf("normalizeKeyModels(%v) = %v, want %v", tt.in, got, tt.want)
}
}
})
}
}
// TestAiKeyModelsPersistence 钉住 serializer:json 字段经 Create 与 Updates(map) 两条路径的往返。
func TestAiKeyModelsPersistence(t *testing.T) {
gw, _ := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}})
ctx := context.Background()
_, key, err := gw.CreateKey(ctx, "m-key", "model-key-1234", "", []string{" x-model ", "x-model", "y-model", ""})
if err != nil {
t.Fatalf("CreateKey: %v", err)
}
fresh, err := gw.VerifyKey(ctx, "model-key-1234")
if err != nil || len(fresh.Models) != 2 || fresh.Models[0] != "x-model" || fresh.Models[1] != "y-model" {
t.Fatalf("创建后读回 models = %v, %v", fresh.Models, err)
}
set := []string{"z-model"}
if err := gw.UpdateKey(ctx, key.ID, "", nil, nil, &set); err != nil {
t.Fatalf("UpdateKey 覆盖: %v", err)
}
fresh, _ = gw.VerifyKey(ctx, "model-key-1234")
if len(fresh.Models) != 1 || fresh.Models[0] != "z-model" {
t.Fatalf("覆盖后 models = %v", fresh.Models)
}
if err := gw.UpdateKey(ctx, key.ID, "renamed", nil, nil, nil); err != nil {
t.Fatalf("UpdateKey 不传 models: %v", err)
}
fresh, _ = gw.VerifyKey(ctx, "model-key-1234")
if len(fresh.Models) != 1 {
t.Fatalf("未传 models 却被改动: %v", fresh.Models)
}
empty := []string{}
if err := gw.UpdateKey(ctx, key.ID, "", nil, nil, &empty); err != nil {
t.Fatalf("UpdateKey 清空: %v", err)
}
fresh, _ = gw.VerifyKey(ctx, "model-key-1234")
if len(fresh.Models) != 0 {
t.Fatalf("清空后 models = %v, want 空", fresh.Models)
}
}
// TestRespPassthroughStreamSwitchesChannel 断言流式直通建立失败按 switchable 换渠道。
func TestRespPassthroughStreamSwitchesChannel(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
passPayload: []byte("data: {\"type\":\"response.completed\"}\n\n"),
passErrs: []error{stubServiceError{status: 503}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
stream, meta, err := gw.RespPassthroughStream(context.Background(), []byte(`{}`), "meta.llama-3.3-70b-instruct", "")
if err != nil || stream == nil {
t.Fatalf("RespPassthroughStream = %v", err)
}
defer stream.Close()
if meta.Retries != 1 || client.passCalls != 2 {
t.Errorf("应换渠道重试一次: retries=%d calls=%d", meta.Retries, client.passCalls)
}
payload, _ := io.ReadAll(stream)
if !bytes.Contains(payload, []byte("response.completed")) {
t.Errorf("流内容未透传: %s", payload)
}
}
+90 -205
View File
@@ -8,38 +8,8 @@ import (
"oci-portal/internal/aiwire" "oci-portal/internal/aiwire"
) )
// ResponsesToIR 把 Responses API 请求转换为网关 IR(OpenAI Chat 形态)。 // respRejectStateful 拒绝有状态特性(网关无状态)。
// 有状态特性与内置工具不支持,直接报错(API 层 400)。 func respRejectStateful(req aiwire.RespRequest) error {
func ResponsesToIR(req aiwire.RespRequest) (aiwire.ChatRequest, error) {
if err := respRejectUnsupported(req); err != nil {
return aiwire.ChatRequest{}, err
}
ir := aiwire.ChatRequest{
Model: req.Model,
MaxTokens: req.MaxOutputTokens,
Temperature: req.Temperature,
TopP: req.TopP,
Stream: req.Stream,
ToolChoice: respToolChoice(req.ToolChoice),
}
if req.Reasoning != nil {
ir.ReasoningEffort = req.Reasoning.Effort
}
if req.Instructions != "" {
ir.Messages = append(ir.Messages, aiwire.ChatMessage{Role: "system", Content: aiwire.NewTextContent(req.Instructions)})
}
msgs, err := respInputToMessages(req.Input)
if err != nil {
return aiwire.ChatRequest{}, err
}
ir.Messages = append(ir.Messages, msgs...)
ir.Tools = respTools(req.Tools)
ir.ResponseFormat = respFormat(req.Text)
return ir, nil
}
// respRejectUnsupported 拒绝无状态网关无法承接的请求特性。
func respRejectUnsupported(req aiwire.RespRequest) error {
if req.PreviousResponseID != "" { if req.PreviousResponseID != "" {
return fmt.Errorf("previous_response_id 不支持:网关不保存历史响应,请在 input 中自带完整上下文(store:false 模式)") return fmt.Errorf("previous_response_id 不支持:网关不保存历史响应,请在 input 中自带完整上下文(store:false 模式)")
} }
@@ -49,195 +19,110 @@ func respRejectUnsupported(req aiwire.RespRequest) error {
if req.Background != nil && *req.Background { if req.Background != nil && *req.Background {
return fmt.Errorf("background 模式不支持") return fmt.Errorf("background 模式不支持")
} }
for _, t := range req.Tools {
if t.Type != "function" {
return fmt.Errorf("不支持的工具类型 %q:仅支持 function 工具", t.Type)
}
}
return nil return nil
} }
// respInputToMessages 把 input(string 或 item 数组)展开为 IR 消息序列 // RespServerTools 报告工具列表是否含 xAI 服务端工具(web_search / x_search)
func respInputToMessages(in aiwire.RespInput) ([]aiwire.ChatMessage, error) { func RespServerTools(tools []aiwire.RespTool) bool {
if !in.IsArray {
if strings.TrimSpace(in.Text) == "" {
return nil, fmt.Errorf("input 不能为空")
}
return []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent(in.Text)}}, nil
}
var msgs []aiwire.ChatMessage
for _, it := range in.Items {
out, err := respItemToMessages(it, msgs)
if err != nil {
return nil, err
}
msgs = out
}
return msgs, nil
}
// respItemToMessages 追加一个 input item;连续 function_call 合并进同一 assistant 消息。
func respItemToMessages(it aiwire.RespItem, msgs []aiwire.ChatMessage) ([]aiwire.ChatMessage, error) {
switch it.Type {
case "", "message":
if it.Content.HasUnsupported() {
return nil, ErrAiUnsupportedBlock
}
role := respRole(it.Role)
return append(msgs, aiwire.ChatMessage{Role: role, Content: respContentToIR(it.Content)}), nil
case "function_call":
call := aiwire.ToolCall{ID: it.CallID, Type: "function",
Function: aiwire.FunctionCall{Name: it.Name, Arguments: it.Arguments}}
if n := len(msgs); n > 0 && msgs[n-1].Role == "assistant" && len(msgs[n-1].ToolCalls) > 0 {
msgs[n-1].ToolCalls = append(msgs[n-1].ToolCalls, call)
return msgs, nil
}
return append(msgs, aiwire.ChatMessage{Role: "assistant", ToolCalls: []aiwire.ToolCall{call}}), nil
case "function_call_output":
return append(msgs, aiwire.ChatMessage{Role: "tool", ToolCallID: it.CallID,
Content: aiwire.NewTextContent(it.OutputText())}), nil
case "reasoning":
return msgs, nil // 推理块不回灌上游
default:
return nil, fmt.Errorf("不支持的 input item 类型 %q", it.Type)
}
}
// respRole 归一 item 角色;developer 视为 system,缺省 user。
func respRole(role string) string {
switch role {
case "system", "developer":
return "system"
case "assistant":
return "assistant"
default:
return "user"
}
}
// respContentToIR 把 Responses 内容转 IR;纯文本拍平,含图片时保留部件顺序。
func respContentToIR(c aiwire.RespContent) aiwire.Content {
var parts []aiwire.ContentPart
hasImage := false
for _, p := range c.Parts {
switch p.Type {
case "input_image":
parts = append(parts, aiwire.ContentPart{Type: "image_url",
ImageURL: &aiwire.ImageURL{URL: p.ImageURL, Detail: p.Detail}})
hasImage = true
default:
if p.Text != "" {
parts = append(parts, aiwire.ContentPart{Type: "text", Text: p.Text})
}
}
}
if !c.IsArray || !hasImage {
return aiwire.NewTextContent(c.JoinText())
}
return aiwire.NewPartsContent(parts)
}
// respTools 把扁平工具定义还原为嵌套 Chat 形态;strict 无 OCI 对应,忽略。
func respTools(tools []aiwire.RespTool) []aiwire.Tool {
if len(tools) == 0 {
return nil
}
out := make([]aiwire.Tool, 0, len(tools))
for _, t := range tools { for _, t := range tools {
out = append(out, aiwire.Tool{Type: "function", if t.Type == "web_search" || t.Type == "x_search" {
Function: aiwire.FunctionDef{Name: t.Name, Description: t.Description, Parameters: t.Parameters}}) return true
} }
return out }
return false
} }
// respToolChoice 转换 tool_choice:字符串直通;{type:function,name} 转嵌套; // RespPassthroughValidate 校验直通请求:模型必填,有状态特性不支持,工具类型
// allowed_tools 等对象形态降级 auto。 // 只放行 function 与已实测的 web_search / x_search;流式仅在含服务端工具时拒绝
func respToolChoice(raw json.RawMessage) json.RawMessage { // (工具流式事件形态未实测,不放开)。
if len(raw) == 0 { func RespPassthroughValidate(req aiwire.RespRequest) error {
return nil if strings.TrimSpace(req.Model) == "" {
return fmt.Errorf("model 不能为空")
} }
t := strings.TrimSpace(string(raw)) if req.Stream && RespServerTools(req.Tools) {
if strings.HasPrefix(t, "\"") { return fmt.Errorf("服务端工具暂不支持流式:请去掉 stream 或改用 function 工具")
return raw
} }
var obj struct { if err := respRejectStateful(req); err != nil {
Type string `json:"type"` return err
Name string `json:"name"`
} }
if err := json.Unmarshal(raw, &obj); err != nil || obj.Type != "function" || obj.Name == "" { for _, t := range req.Tools {
return json.RawMessage(`"auto"`) switch t.Type {
} case "function", "web_search", "x_search":
out, _ := json.Marshal(map[string]any{"type": "function", "function": map[string]string{"name": obj.Name}})
return out
}
// respFormat 把 text.format 转成 response_format;text 形态无需显式指定。
func respFormat(text *aiwire.RespText) *aiwire.ResponseFormat {
if text == nil || text.Format == nil {
return nil
}
switch text.Format.Type {
case "json_object":
return &aiwire.ResponseFormat{Type: "json_object"}
case "json_schema":
spec, _ := json.Marshal(map[string]any{
"name": text.Format.Name, "schema": json.RawMessage(text.Format.Schema), "strict": text.Format.Strict,
})
return &aiwire.ResponseFormat{Type: "json_schema", JSONSchema: spec}
default: default:
return fmt.Errorf("不支持的工具类型 %q:服务端工具仅支持 web_search / x_search", t.Type)
}
}
return nil
}
// RespPassthroughBody 以原始请求体为基构造上游 body:强制 store:false(禁上游
// 存态),stream 原样保留(流式直通);用 json.Number 保真未知字段与数值。
func RespPassthroughBody(raw []byte) ([]byte, error) {
dec := json.NewDecoder(strings.NewReader(string(raw)))
dec.UseNumber()
var body map[string]any
if err := dec.Decode(&body); err != nil {
return nil, fmt.Errorf("解析请求体: %w", err)
}
body["store"] = false
return json.Marshal(body)
}
// RespPassthroughUsage 从直通响应提取用量;缺失时返回 nil(日志记零)。
func RespPassthroughUsage(payload []byte) *aiwire.Usage {
var root struct {
Usage *aiwire.RespUsage `json:"usage"`
}
if json.Unmarshal(payload, &root) != nil {
return nil return nil
} }
return usageFromResp(root.Usage)
} }
// IRRespToResponses 把 IR 非流式响应装配为 Response 对象 // usageFromResp 换算 Responses usage 为 OpenAI 口径(缓存命中透传);nil 原样返回
func IRRespToResponses(resp *aiwire.ChatResponse, id string, created int64) aiwire.Response { func usageFromResp(u *aiwire.RespUsage) *aiwire.Usage {
out := aiwire.Response{ID: id, Object: "response", CreatedAt: created,
Status: "completed", Model: resp.Model, Output: []aiwire.RespOutItem{}, Store: false}
if len(resp.Choices) == 0 {
return out
}
choice := resp.Choices[0]
if text := choice.Message.Content.JoinText(); text != "" {
out.Output = append(out.Output, respMessageItem(id, 0, text))
}
for i, tc := range choice.Message.ToolCalls {
out.Output = append(out.Output, aiwire.RespOutItem{Type: "function_call",
ID: respItemID(id, "fc", len(out.Output)+i), Status: "completed",
CallID: tc.ID, Name: tc.Function.Name, Arguments: respArgs(tc.Function.Arguments)})
}
if choice.FinishReason == "length" {
out.Status = "incomplete"
out.IncompleteDetails = &aiwire.RespIncomplete{Reason: "max_output_tokens"}
}
out.Usage = respUsage(resp.Usage)
return out
}
func respMessageItem(respID string, idx int, text string) aiwire.RespOutItem {
return aiwire.RespOutItem{Type: "message", ID: respItemID(respID, "msg", idx),
Status: "completed", Role: "assistant",
Content: []aiwire.RespOutPart{{Type: "output_text", Text: text, Annotations: []any{}}}}
}
// respItemID 从响应 ID 派生确定性 item ID。
func respItemID(respID, kind string, idx int) string {
return fmt.Sprintf("%s_%s_%d", kind, strings.TrimPrefix(respID, "resp_"), idx)
}
// respArgs 保证 arguments 是合法 JSON 字符串(空实参回退 {})。
func respArgs(args string) string {
if strings.TrimSpace(args) == "" {
return "{}"
}
return args
}
// respUsage 把 IR 用量改写为 Responses 命名口径。
func respUsage(u *aiwire.Usage) *aiwire.RespUsage {
if u == nil { if u == nil {
return nil return nil
} }
return &aiwire.RespUsage{InputTokens: u.PromptTokens, OutputTokens: u.CompletionTokens, usage := &aiwire.Usage{PromptTokens: u.InputTokens,
TotalTokens: u.TotalTokens, CompletionTokens: u.OutputTokens, TotalTokens: u.TotalTokens}
InputTokensDetails: aiwire.RespInDetails{CachedTokens: u.CachedTokens()}} if cached := u.InputTokensDetails.CachedTokens; cached > 0 {
usage.PromptTokensDetails = &aiwire.PromptTokensDetails{CachedTokens: cached}
}
return usage
}
// RespStreamCompletedUsage 从一行 SSE data JSON 中提取 response.completed 事件的
// usage;非 completed 事件或解析失败返回 nil。流式直通逐行喂入,最后一次非 nil 生效。
func RespStreamCompletedUsage(data []byte) *aiwire.Usage {
var ev struct {
Type string `json:"type"`
Response json.RawMessage `json:"response"`
}
if json.Unmarshal(data, &ev) != nil || ev.Type != "response.completed" || len(ev.Response) == 0 {
return nil
}
return RespPassthroughUsage(ev.Response)
}
// RespStreamErrorMsg 从一行 SSE data JSON 中提取 error / response.failed 事件的
// 错误消息;非错误事件返回空串。流式直通据此把上游错误写入调用日志。
func RespStreamErrorMsg(data []byte) string {
var ev respStreamEvent
if json.Unmarshal(data, &ev) != nil {
return ""
}
switch ev.Type {
case "error":
if ev.Message != "" {
return ev.Message
}
return "上游返回错误事件 error"
case "response.failed":
if m := respErrorMsg(ev.Response); m != "" {
return m
}
return "上游返回错误事件 response.failed"
}
return ""
} }
-199
View File
@@ -1,199 +0,0 @@
package service
import (
"strings"
"oci-portal/internal/aiwire"
)
// RespEvent 是一条 Responses SSE 语义事件(event 名 + data 载荷)。
type RespEvent struct {
Event string
Data map[string]any
}
// RespStream 把 IR chunk 流聚合为 Responses 语义事件序列。
// 与 AnthStream 同构,但终态事件须携带完整 Response 快照,故全程缓冲文本与实参。
type RespStream struct {
id string
model string
created int64
seq int
started bool
items []aiwire.RespOutItem
msgOpen bool
text strings.Builder
tool aiwire.RespOutItem
toolIdx int
tOpen bool
args strings.Builder
finish string
usage *aiwire.Usage
}
// NewRespStream 构造状态机;id 形如 resp_*,created 为响应时间戳。
func NewRespStream(id, model string, created int64) *RespStream {
return &RespStream{id: id, model: model, created: created}
}
// Feed 消化一个上游 chunk,返回应立即下发的事件。
func (s *RespStream) Feed(chunk aiwire.ChatChunk) []RespEvent {
var evs []RespEvent
if !s.started {
s.started = true
evs = append(evs, s.respEvent("response.created", "in_progress"),
s.respEvent("response.in_progress", "in_progress"))
}
if chunk.Usage != nil {
s.usage = chunk.Usage
}
if len(chunk.Choices) == 0 {
return evs
}
choice := chunk.Choices[0]
if choice.Delta.Content != "" {
evs = append(evs, s.feedText(choice.Delta.Content)...)
}
for _, tc := range choice.Delta.ToolCalls {
evs = append(evs, s.feedTool(tc)...)
}
if choice.FinishReason != nil {
s.finish = *choice.FinishReason
}
return evs
}
// Finish 关闭未闭合的块并产出终态事件(带完整 Response 快照)。
func (s *RespStream) Finish() []RespEvent {
var evs []RespEvent
if !s.started {
s.started = true
evs = append(evs, s.respEvent("response.created", "in_progress"))
}
evs = append(evs, s.closeMsg()...)
evs = append(evs, s.closeTool()...)
status := "completed"
if s.finish == "length" {
status = "incomplete"
}
return append(evs, s.respEvent("response."+status, status))
}
// Usage 返回聚合到的用量(可能为 nil),供调用日志。
func (s *RespStream) Usage() *aiwire.Usage { return s.usage }
// ev 构造带自增 sequence_number 的事件。
func (s *RespStream) ev(typ string, kv map[string]any) RespEvent {
s.seq++
data := map[string]any{"type": typ, "sequence_number": s.seq}
for k, v := range kv {
data[k] = v
}
return RespEvent{Event: typ, Data: data}
}
// respEvent 构造携带 Response 快照的生命周期事件。
func (s *RespStream) respEvent(typ, status string) RespEvent {
return s.ev(typ, map[string]any{"response": s.snapshot(status)})
}
// snapshot 组装当前累计状态的 Response 对象。
func (s *RespStream) snapshot(status string) aiwire.Response {
resp := aiwire.Response{ID: s.id, Object: "response", CreatedAt: s.created,
Status: status, Model: s.model, Store: false,
Output: append([]aiwire.RespOutItem{}, s.items...)}
resp.Usage = respUsage(s.usage)
if status == "incomplete" {
resp.IncompleteDetails = &aiwire.RespIncomplete{Reason: "max_output_tokens"}
}
return resp
}
// feedText 处理文本增量:必要时开 message item 与 content part。
func (s *RespStream) feedText(delta string) []RespEvent {
evs := s.closeTool()
if !s.msgOpen {
s.msgOpen = true
s.text.Reset()
item := aiwire.RespOutItem{Type: "message", ID: s.itemID("msg"), Status: "in_progress",
Role: "assistant", Content: []aiwire.RespOutPart{}}
evs = append(evs, s.ev("response.output_item.added", map[string]any{
"output_index": len(s.items), "item": item}))
evs = append(evs, s.ev("response.content_part.added", map[string]any{
"item_id": item.ID, "output_index": len(s.items), "content_index": 0,
"part": aiwire.RespOutPart{Type: "output_text", Text: "", Annotations: []any{}}}))
}
s.text.WriteString(delta)
return append(evs, s.ev("response.output_text.delta", map[string]any{
"item_id": s.itemID("msg"), "output_index": len(s.items), "content_index": 0, "delta": delta}))
}
// closeMsg 闭合当前 message item(text done → part done → item done)。
func (s *RespStream) closeMsg() []RespEvent {
if !s.msgOpen {
return nil
}
s.msgOpen = false
id, idx, full := s.itemID("msg"), len(s.items), s.text.String()
part := aiwire.RespOutPart{Type: "output_text", Text: full, Annotations: []any{}}
item := aiwire.RespOutItem{Type: "message", ID: id, Status: "completed",
Role: "assistant", Content: []aiwire.RespOutPart{part}}
evs := []RespEvent{
s.ev("response.output_text.done", map[string]any{"item_id": id, "output_index": idx, "content_index": 0, "text": full}),
s.ev("response.content_part.done", map[string]any{"item_id": id, "output_index": idx, "content_index": 0, "part": part}),
s.ev("response.output_item.done", map[string]any{"output_index": idx, "item": item}),
}
s.items = append(s.items, item)
return evs
}
// feedTool 处理工具调用增量:新 Index 先闭合旧调用再开新 item。
func (s *RespStream) feedTool(tc aiwire.ToolCallDelta) []RespEvent {
evs := s.closeMsg()
if s.tOpen && tc.Index != s.toolIdx {
evs = append(evs, s.closeTool()...)
}
if !s.tOpen {
s.tOpen, s.toolIdx = true, tc.Index
s.args.Reset()
s.tool = aiwire.RespOutItem{Type: "function_call", ID: s.itemID("fc"), Status: "in_progress",
CallID: tc.ID, Name: tc.Function.Name}
evs = append(evs, s.ev("response.output_item.added", map[string]any{
"output_index": len(s.items), "item": s.tool}))
}
if tc.ID != "" && s.tool.CallID == "" {
s.tool.CallID = tc.ID
}
if tc.Function.Name != "" && s.tool.Name == "" {
s.tool.Name = tc.Function.Name
}
if tc.Function.Arguments == "" {
return evs
}
s.args.WriteString(tc.Function.Arguments)
return append(evs, s.ev("response.function_call_arguments.delta", map[string]any{
"item_id": s.tool.ID, "output_index": len(s.items), "delta": tc.Function.Arguments}))
}
// closeTool 闭合当前 function_call item(arguments done → item done)。
func (s *RespStream) closeTool() []RespEvent {
if !s.tOpen {
return nil
}
s.tOpen = false
s.tool.Arguments = respArgs(s.args.String())
s.tool.Status = "completed"
idx := len(s.items)
evs := []RespEvent{
s.ev("response.function_call_arguments.done", map[string]any{
"item_id": s.tool.ID, "output_index": idx, "arguments": s.tool.Arguments}),
s.ev("response.output_item.done", map[string]any{"output_index": idx, "item": s.tool}),
}
s.items = append(s.items, s.tool)
return evs
}
// itemID 按当前 output 序号派生确定性 item ID。
func (s *RespStream) itemID(kind string) string {
return respItemID(s.id, kind, len(s.items))
}
+102 -131
View File
@@ -17,47 +17,6 @@ func respReq(t *testing.T, raw string) aiwire.RespRequest {
return req return req
} }
func TestResponsesToIR(t *testing.T) {
req := respReq(t, `{
"model": "meta.llama-3.3-70b-instruct",
"instructions": "你是助手",
"max_output_tokens": 128,
"input": [
{"role": "user", "content": "查天气"},
{"type": "function_call", "call_id": "call_1", "name": "get_weather", "arguments": "{\"city\":\"上海\"}"},
{"type": "function_call_output", "call_id": "call_1", "output": "{\"temp\":31}"},
{"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "31 度"}]}
],
"tools": [{"type": "function", "name": "get_weather", "description": "查天气", "parameters": {"type": "object"}, "strict": true}],
"tool_choice": {"type": "function", "name": "get_weather"},
"text": {"format": {"type": "json_schema", "name": "out", "schema": {"type": "object"}}}
}`)
ir, err := ResponsesToIR(req)
if err != nil {
t.Fatalf("ResponsesToIR: %v", err)
}
roles := make([]string, 0, len(ir.Messages))
for _, m := range ir.Messages {
roles = append(roles, m.Role)
}
want := []string{"system", "user", "assistant", "tool", "assistant"}
if strings.Join(roles, ",") != strings.Join(want, ",") {
t.Errorf("roles = %v, want %v", roles, want)
}
if ir.Messages[2].ToolCalls[0].ID != "call_1" || ir.Messages[3].ToolCallID != "call_1" {
t.Errorf("tool call 链接错误: %+v", ir.Messages)
}
if deref(ir.MaxTokens) != 128 || len(ir.Tools) != 1 || ir.Tools[0].Function.Name != "get_weather" {
t.Errorf("参数映射错误: max=%v tools=%+v", ir.MaxTokens, ir.Tools)
}
if !strings.Contains(string(ir.ToolChoice), `"function"`) || !strings.Contains(string(ir.ToolChoice), "get_weather") {
t.Errorf("tool_choice = %s", ir.ToolChoice)
}
if ir.ResponseFormat == nil || ir.ResponseFormat.Type != "json_schema" {
t.Errorf("response_format = %+v", ir.ResponseFormat)
}
}
func deref(p *int) int { func deref(p *int) int {
if p == nil { if p == nil {
return 0 return 0
@@ -65,112 +24,124 @@ func deref(p *int) int {
return *p return *p
} }
func TestResponsesToIRRejects(t *testing.T) { func TestRespServerTools(t *testing.T) {
cases := []struct{ name, raw string }{ tests := []struct {
{"previous_response_id", `{"model":"m","input":"hi","previous_response_id":"resp_x"}`}, name string
{"conversation", `{"model":"m","input":"hi","conversation":{"id":"conv_1"}}`}, tools []aiwire.RespTool
{"background", `{"model":"m","input":"hi","background":true}`}, want bool
{"builtin tool", `{"model":"m","input":"hi","tools":[{"type":"web_search"}]}`}, }{
{"file_id image", `{"model":"m","input":[{"role":"user","content":[{"type":"input_image","file_id":"file_1"}]}]}`}, {"web_search", []aiwire.RespTool{{Type: "web_search"}}, true},
{"unknown item", `{"model":"m","input":[{"type":"item_reference","id":"x"}]}`}, {"x_search混用", []aiwire.RespTool{{Type: "function", Name: "f"}, {Type: "x_search"}}, true},
{"仅function", []aiwire.RespTool{{Type: "function", Name: "f"}}, false},
{"空", nil, false},
{"其他类型", []aiwire.RespTool{{Type: "code_interpreter"}}, false},
} }
for _, c := range cases { for _, test := range tests {
if _, err := ResponsesToIR(respReq(t, c.raw)); err == nil { t.Run(test.name, func(t *testing.T) {
t.Errorf("%s 应被拒绝", c.name) if got := RespServerTools(test.tools); got != test.want {
t.Fatalf("got %v, want %v", got, test.want)
} }
})
} }
// input 字符串形态 + store/reasoning 忽略项不报错 }
req := respReq(t, `{"model":"m","input":"你好","store":false,"reasoning":{"effort":"low"},"parallel_tool_calls":true}`)
ir, err := ResponsesToIR(req) func TestRespPassthroughValidate(t *testing.T) {
if err != nil || len(ir.Messages) != 1 || ir.Messages[0].Role != "user" || ir.ReasoningEffort != "low" { prev := "resp_1"
t.Errorf("字符串 input = %+v, %v", ir.Messages, err) bg := true
tests := []struct {
name string
req aiwire.RespRequest
wantErr bool
}{
{"合法", aiwire.RespRequest{Model: "m", Tools: []aiwire.RespTool{{Type: "web_search"}}}, false},
{"混用function", aiwire.RespRequest{Model: "m", Tools: []aiwire.RespTool{{Type: "web_search"}, {Type: "function", Name: "f"}}}, false},
{"缺model", aiwire.RespRequest{Tools: []aiwire.RespTool{{Type: "web_search"}}}, true},
{"工具加流式拒绝", aiwire.RespRequest{Model: "m", Stream: true, Tools: []aiwire.RespTool{{Type: "web_search"}}}, true},
{"无工具流式放行", aiwire.RespRequest{Model: "m", Stream: true}, false},
{"function工具流式放行", aiwire.RespRequest{Model: "m", Stream: true, Tools: []aiwire.RespTool{{Type: "function", Name: "f"}}}, false},
{"有状态拒绝", aiwire.RespRequest{Model: "m", PreviousResponseID: prev}, true},
{"background拒绝", aiwire.RespRequest{Model: "m", Background: &bg}, true},
{"未知工具拒绝", aiwire.RespRequest{Model: "m", Tools: []aiwire.RespTool{{Type: "web_search"}, {Type: "mcp"}}}, true},
} }
// input_image(url 形态)放行并保留图文顺序 for _, test := range tests {
req2 := respReq(t, `{"model":"m","input":[{"role":"user","content":[{"type":"input_text","text":"看图"},{"type":"input_image","image_url":"https://x/1.png","detail":"low"}]}]}`) t.Run(test.name, func(t *testing.T) {
ir2, err := ResponsesToIR(req2) err := RespPassthroughValidate(test.req)
if (err != nil) != test.wantErr {
t.Fatalf("err = %v, wantErr %v", err, test.wantErr)
}
})
}
}
func TestRespPassthroughBody(t *testing.T) {
raw := []byte(`{"model":"m","input":"hi","stream":true,"store":true,"max_output_tokens":128,"custom_field":{"a":1.5}}`)
out, err := RespPassthroughBody(raw)
if err != nil { if err != nil {
t.Fatalf("input_image 应放行: %v", err) t.Fatalf("RespPassthroughBody: %v", err)
} }
parts := ir2.Messages[0].Content.Parts var body map[string]any
if len(parts) != 2 || parts[1].Type != "image_url" || parts[1].ImageURL == nil || parts[1].ImageURL.URL != "https://x/1.png" { if err := json.Unmarshal(out, &body); err != nil {
t.Errorf("parts = %+v", parts) t.Fatalf("unmarshal out: %v", err)
}
if body["store"] != false {
t.Errorf("store 应强制 false, got %v", body["store"])
}
if body["stream"] != true {
t.Errorf("stream 应原样保留, got %v", body["stream"])
}
if string(out) == "" || !strings.Contains(string(out), `"max_output_tokens":128`) {
t.Errorf("数值字段应保真: %s", out)
}
if !strings.Contains(string(out), `"custom_field"`) {
t.Errorf("未知字段应保留: %s", out)
}
if _, err := RespPassthroughBody([]byte("not-json")); err == nil {
t.Error("非法 JSON 应报错")
} }
} }
func TestIRRespToResponses(t *testing.T) { func TestRespPassthroughUsage(t *testing.T) {
resp := &aiwire.ChatResponse{Model: "m", Choices: []aiwire.Choice{{ tests := []struct {
Message: aiwire.ChatMessage{Role: "assistant", Content: aiwire.NewTextContent("你好"), name string
ToolCalls: []aiwire.ToolCall{{ID: "call_9", Type: "function", payload string
Function: aiwire.FunctionCall{Name: "f", Arguments: ""}}}}, wantNil bool
FinishReason: "length", prompt, cached int
}}, Usage: &aiwire.Usage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15}} }{
out := IRRespToResponses(resp, "resp_abc", 1751966400) {"完整", `{"usage":{"input_tokens":9,"output_tokens":3,"total_tokens":12,"input_tokens_details":{"cached_tokens":5}}}`, false, 9, 5},
if out.Status != "incomplete" || out.IncompleteDetails == nil || out.IncompleteDetails.Reason != "max_output_tokens" { {"无细分", `{"usage":{"input_tokens":9,"output_tokens":3,"total_tokens":12}}`, false, 9, 0},
t.Errorf("length 应映射 incomplete: %+v", out) {"无usage", `{"id":"resp_1"}`, true, 0, 0},
{"非法JSON", `xx`, true, 0, 0},
} }
if len(out.Output) != 2 || out.Output[0].Type != "message" || out.Output[1].Type != "function_call" { for _, test := range tests {
t.Fatalf("output = %+v", out.Output) t.Run(test.name, func(t *testing.T) {
usage := RespPassthroughUsage([]byte(test.payload))
if (usage == nil) != test.wantNil {
t.Fatalf("usage = %v, wantNil %v", usage, test.wantNil)
} }
if out.Output[0].Content[0].Text != "你好" || out.Output[1].CallID != "call_9" || out.Output[1].Arguments != "{}" { if usage == nil {
t.Errorf("item 装配错误: %+v", out.Output) return
} }
if out.Usage.InputTokens != 10 || out.Usage.OutputTokens != 5 || out.Store { if usage.PromptTokens != test.prompt || usage.CachedTokens() != test.cached {
t.Errorf("usage/store 错误: %+v", out) t.Fatalf("prompt=%d cached=%d, want %d/%d", usage.PromptTokens, usage.CachedTokens(), test.prompt, test.cached)
}
})
} }
} }
func chunkText(text string) aiwire.ChatChunk { // TestRespStreamCompletedUsage 断言流式 usage 只从 completed 事件提取。
return aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{Content: text}}}} func TestRespStreamCompletedUsage(t *testing.T) {
} completed := []byte(`{"type":"response.completed","response":{"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15,"input_tokens_details":{"cached_tokens":4}}}}`)
if u := RespStreamCompletedUsage(completed); u == nil || u.PromptTokens != 10 || u.CompletionTokens != 5 ||
func TestRespStreamTextAndTool(t *testing.T) { u.PromptTokensDetails == nil || u.PromptTokensDetails.CachedTokens != 4 {
st := NewRespStream("resp_x", "m", 1751966400) t.Fatalf("completed 事件 usage 解析失败: %+v", u)
var types []string
collect := func(evs []RespEvent) {
for _, ev := range evs {
types = append(types, ev.Event)
} }
for _, data := range []string{
`{"type":"response.output_text.delta","delta":"hi"}`,
`{"type":"response.completed"}`,
`not-json`,
} {
if u := RespStreamCompletedUsage([]byte(data)); u != nil {
t.Fatalf("非 completed 或缺 response 应返回 nil: %s", data)
} }
collect(st.Feed(chunkText("你")))
collect(st.Feed(chunkText("好")))
collect(st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{
ToolCalls: []aiwire.ToolCallDelta{{Index: 0, ID: "call_1", Function: aiwire.FunctionCallDelta{Name: "f", Arguments: "{\"a\""}}}}}}}))
collect(st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{
ToolCalls: []aiwire.ToolCallDelta{{Index: 0, Function: aiwire.FunctionCallDelta{Arguments: ":1}"}}}}}},
Usage: &aiwire.Usage{PromptTokens: 3, CompletionTokens: 7, TotalTokens: 10}}))
finish := st.Finish()
collect(finish)
want := []string{
"response.created", "response.in_progress",
"response.output_item.added", "response.content_part.added", "response.output_text.delta",
"response.output_text.delta",
"response.output_text.done", "response.content_part.done", "response.output_item.done",
"response.output_item.added", "response.function_call_arguments.delta",
"response.function_call_arguments.delta",
"response.function_call_arguments.done", "response.output_item.done",
"response.completed",
}
if strings.Join(types, "\n") != strings.Join(want, "\n") {
t.Errorf("事件序列 =\n%s\nwant\n%s", strings.Join(types, "\n"), strings.Join(want, "\n"))
}
last := finish[len(finish)-1].Data
resp := last["response"].(aiwire.Response)
if len(resp.Output) != 2 || resp.Output[0].Content[0].Text != "你好" || resp.Output[1].Arguments != "{\"a\":1}" {
t.Errorf("终态快照 = %+v", resp.Output)
}
if resp.Usage == nil || resp.Usage.TotalTokens != 10 {
t.Errorf("终态 usage = %+v", resp.Usage)
}
// sequence_number 严格递增
if seq, ok := last["sequence_number"].(int); !ok || seq != len(want) {
t.Errorf("最终 sequence_number = %v, want %d", last["sequence_number"], len(want))
}
}
func TestRespStreamEmpty(t *testing.T) {
st := NewRespStream("resp_e", "m", 1)
evs := st.Finish()
if len(evs) != 2 || evs[0].Event != "response.created" || evs[1].Event != "response.completed" {
t.Errorf("空流事件 = %+v", evs)
} }
} }
-305
View File
@@ -1,305 +0,0 @@
package service
import (
"context"
"errors"
"fmt"
"log"
"net/netip"
"slices"
"strings"
"time"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"oci-portal/internal/model"
)
// 告警规则约束与命中记录保留期(窗口计数之外多留几天便于排查)。
const (
alertMaxThreshold = 100
alertMaxWindowMin = 1440
alertHitRetention = 7 * 24 * time.Hour
alertSourceIPIn = "in"
alertSourceIPNotIn = "notin"
)
// ErrInvalidAlertRule 标记规则字段非法,api 层映射 400。
var ErrInvalidAlertRule = fmt.Errorf("告警规则字段非法")
// ListAlertRules 返回全部告警规则(创建顺序)。
func (s *LogEventService) ListAlertRules(ctx context.Context) ([]model.AlertRule, error) {
var rules []model.AlertRule
if err := s.db.WithContext(ctx).Order("id").Find(&rules).Error; err != nil {
return nil, fmt.Errorf("list alert rules: %w", err)
}
return rules, nil
}
// CreateAlertRule 校验并创建规则。
func (s *LogEventService) CreateAlertRule(ctx context.Context, rule model.AlertRule) (model.AlertRule, error) {
if err := validateAlertRule(&rule); err != nil {
return model.AlertRule{}, err
}
rule.ID = 0
if err := s.db.WithContext(ctx).Create(&rule).Error; err != nil {
return model.AlertRule{}, fmt.Errorf("create alert rule: %w", err)
}
return rule, nil
}
// UpdateAlertRule 校验并整体覆盖规则(含启停)。
func (s *LogEventService) UpdateAlertRule(ctx context.Context, id uint, rule model.AlertRule) (model.AlertRule, error) {
if err := validateAlertRule(&rule); err != nil {
return model.AlertRule{}, err
}
var cur model.AlertRule
if err := s.db.WithContext(ctx).First(&cur, id).Error; err != nil {
return model.AlertRule{}, fmt.Errorf("find alert rule %d: %w", id, err)
}
rule.ID, rule.CreatedAt = cur.ID, cur.CreatedAt
if err := s.db.WithContext(ctx).Save(&rule).Error; err != nil {
return model.AlertRule{}, fmt.Errorf("update alert rule: %w", err)
}
return rule, nil
}
// DeleteAlertRule 删除规则及其命中记录。
func (s *LogEventService) DeleteAlertRule(ctx context.Context, id uint) error {
if err := s.db.WithContext(ctx).Delete(&model.AlertRule{}, id).Error; err != nil {
return fmt.Errorf("delete alert rule: %w", err)
}
if err := s.db.WithContext(ctx).Where("rule_id = ?", id).Delete(&model.AlertRuleHit{}).Error; err != nil {
return fmt.Errorf("delete alert rule hits: %w", err)
}
return nil
}
// validateAlertRule 校验字段并归一化;非法时返回含具体原因的 ErrInvalidAlertRule 包装。
func validateAlertRule(rule *model.AlertRule) error {
rule.Name = strings.TrimSpace(rule.Name)
if rule.Name == "" {
return fmt.Errorf("%w: 名称必填", ErrInvalidAlertRule)
}
if rule.SourceIPMode == "" {
rule.SourceIPMode = alertSourceIPIn
}
if rule.SourceIPMode != alertSourceIPIn && rule.SourceIPMode != alertSourceIPNotIn {
return fmt.Errorf("%w: 来源 IP 模式须为 in/notin", ErrInvalidAlertRule)
}
if rule.Threshold < 1 || rule.Threshold > alertMaxThreshold {
return fmt.Errorf("%w: 阈值须在 1-%d 之间", ErrInvalidAlertRule, alertMaxThreshold)
}
if rule.Threshold > 1 && (rule.WindowMinutes < 1 || rule.WindowMinutes > alertMaxWindowMin) {
return fmt.Errorf("%w: 阈值>1 时窗口须在 1-%d 分钟之间", ErrInvalidAlertRule, alertMaxWindowMin)
}
if rule.EventTypes != "" {
rule.EventTypes = normalizeCSV(rule.EventTypes)
}
return validateAlertRuleIPs(rule)
}
// validateAlertRuleIPs 归一化并校验来源 IP 列表(裸 IP 或 CIDR)。
func validateAlertRuleIPs(rule *model.AlertRule) error {
if rule.SourceIPs == "" {
return nil
}
rule.SourceIPs = normalizeCSV(rule.SourceIPs)
for _, item := range strings.Split(rule.SourceIPs, ",") {
if _, err := parseIPMatcher(item); err != nil {
return fmt.Errorf("%w: 来源 IP %q 不是合法的 IP 或 CIDR", ErrInvalidAlertRule, item)
}
}
return nil
}
// normalizeCSV 去除各项空白与空项后重组逗号分隔串。
func normalizeCSV(s string) string {
parts := strings.Split(s, ",")
out := parts[:0]
for _, p := range parts {
if p = strings.TrimSpace(p); p != "" {
out = append(out, p)
}
}
return strings.Join(out, ",")
}
// parseIPMatcher 把裸 IP 或 CIDR 解析为前缀(裸 IP 视为单地址前缀)。
func parseIPMatcher(item string) (netip.Prefix, error) {
if strings.Contains(item, "/") {
return netip.ParsePrefix(item)
}
addr, err := netip.ParseAddr(item)
if err != nil {
return netip.Prefix{}, err
}
return netip.PrefixFrom(addr, addr.BitLen()), nil
}
// ipListMatch 报告 ip 是否命中列表中的任一前缀;ip 解析失败视为未命中。
func ipListMatch(list, ip string) bool {
addr, err := netip.ParseAddr(ip)
if err != nil {
return false
}
for _, item := range strings.Split(list, ",") {
if p, err := parseIPMatcher(item); err == nil && p.Contains(addr) {
return true
}
}
return false
}
// ruleHits 报告事件是否命中规则的全部条件(AND 语义,空条件视为任意)。
func ruleHits(rule model.AlertRule, e *model.LogEvent, p parsedEvent) bool {
if rule.OciConfigID != 0 && rule.OciConfigID != e.OciConfigID {
return false
}
name := relayEventShortName(p.EventType)
if rule.EventTypes != "" && !slices.Contains(strings.Split(rule.EventTypes, ","), name) {
return false
}
if rule.ResourceMatch != "" && !strings.Contains(p.ResourceName, rule.ResourceMatch) {
return false
}
return ruleIPHits(rule, p.SourceIP)
}
// ruleIPHits 按模式判定来源 IP 条件:in 命中列表告警;notin 不在列表才告警,
// 事件缺 IP 字段时 notin 不告警(避免解析缺字段导致白名单误报)。
func ruleIPHits(rule model.AlertRule, ip string) bool {
if rule.SourceIPs == "" {
return true
}
if rule.SourceIPMode == alertSourceIPNotIn {
return ip != "" && !ipListMatch(rule.SourceIPs, ip)
}
return ipListMatch(rule.SourceIPs, ip)
}
// matchAlertRules 对一条已解析事件执行全部启用规则;任何内部错误只记日志,不影响解析主流程。
func (s *LogEventService) matchAlertRules(ctx context.Context, rules []model.AlertRule, e *model.LogEvent, p parsedEvent) {
if s.notifier == nil {
return
}
for _, rule := range rules {
if !rule.Enabled || !ruleHits(rule, e, p) {
continue
}
count, ok := s.recordAlertHit(ctx, rule, e)
if !ok || count < rule.Threshold || !s.alertCooldownPass(rule) {
continue
}
s.notifier.SendTemplateAsync("audit_alert", map[string]string{
"rule": rule.Name, "tenant": s.configAlias(ctx, e.OciConfigID),
"event": relayEventShortName(p.EventType), "resource": p.ResourceName,
"ip": p.SourceIP, "count": fmt.Sprint(count),
})
}
}
// recordAlertHit 落一条命中并返回窗口内累计次数;阈值 1 的规则免计数直接触发。
func (s *LogEventService) recordAlertHit(ctx context.Context, rule model.AlertRule, e *model.LogEvent) (int, bool) {
count, err := s.recordAlertHitTx(ctx, rule, e)
if err != nil {
if !errors.Is(err, gorm.ErrRecordNotFound) {
log.Printf("alert rule hit record: %v", err)
}
return 0, false
}
return count, true
}
func (s *LogEventService) recordAlertHitTx(ctx context.Context, rule model.AlertRule, event *model.LogEvent) (int, error) {
count := 0
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var err error
count, err = insertAlertHit(tx, rule, event)
return err
})
return count, err
}
// insertAlertHit 按 rule→event 锁顺序确认引用存在后插入并统计窗口命中。
func insertAlertHit(tx *gorm.DB, rule model.AlertRule, event *model.LogEvent) (int, error) {
if err := lockAlertRefs(tx, rule.ID, event.ID); err != nil {
return 0, err
}
if rule.Threshold <= 1 {
return 1, nil
}
now := time.Now()
hit := model.AlertRuleHit{RuleID: rule.ID, LogEventID: event.ID, HitAt: now}
if err := tx.Create(&hit).Error; err != nil {
return 0, fmt.Errorf("create alert rule hit: %w", err)
}
var count int64
cutoff := now.Add(-time.Duration(rule.WindowMinutes) * time.Minute)
err := tx.Model(&model.AlertRuleHit{}).
Where("rule_id = ? AND hit_at >= ?", rule.ID, cutoff).Count(&count).Error
if err != nil {
return 0, fmt.Errorf("count alert rule hits: %w", err)
}
return int(count), nil
}
func lockAlertRefs(tx *gorm.DB, ruleID, eventID uint) error {
var rule model.AlertRule
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(&rule, ruleID).Error; err != nil {
return fmt.Errorf("lock alert rule %d: %w", ruleID, err)
}
var event model.LogEvent
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(&event, eventID).Error; err != nil {
return fmt.Errorf("lock log event %d: %w", eventID, err)
}
return nil
}
// alertCooldownPass 报告规则是否已过冷却窗口;通过即记录本次发送时刻。
// 阈值 1 的规则无冷却(每次命中即时告警,与既有云端事件通知一致)。
func (s *LogEventService) alertCooldownPass(rule model.AlertRule) bool {
if rule.Threshold <= 1 {
return true
}
s.alertMu.Lock()
defer s.alertMu.Unlock()
window := time.Duration(rule.WindowMinutes) * time.Minute
if last, ok := s.alertSentAt[rule.ID]; ok && time.Since(last) < window {
return false
}
if s.alertSentAt == nil {
s.alertSentAt = map[uint]time.Time{}
}
s.alertSentAt[rule.ID] = time.Now()
return true
}
// ClearAlertCooldown 清除已删除租户规则的进程内冷却状态。
func (s *LogEventService) ClearAlertCooldown(ruleIDs []uint) {
s.alertMu.Lock()
defer s.alertMu.Unlock()
for _, id := range ruleIDs {
delete(s.alertSentAt, id)
}
}
// loadEnabledAlertRules 载入启用中的规则;失败时返回空集并记日志(解析主流程照常)。
func (s *LogEventService) loadEnabledAlertRules(ctx context.Context) []model.AlertRule {
var rules []model.AlertRule
err := s.db.WithContext(ctx).Where("enabled = ?", true).Order("id").Find(&rules).Error
if err != nil {
log.Printf("load alert rules: %v", err)
return nil
}
return rules
}
// cleanupAlertHits 删除保留期外的命中记录(随 cleanupOnce 周期执行)。
func (s *LogEventService) cleanupAlertHits(ctx context.Context) {
cutoff := time.Now().Add(-alertHitRetention)
if err := s.db.WithContext(ctx).Where("hit_at < ?", cutoff).Delete(&model.AlertRuleHit{}).Error; err != nil {
log.Printf("cleanup alert hits: %v", err)
}
}
-345
View File
@@ -1,345 +0,0 @@
package service
import (
"context"
"fmt"
"strings"
"testing"
"time"
"oci-portal/internal/crypto"
"oci-portal/internal/model"
)
func TestValidateAlertRule(t *testing.T) {
tests := []struct {
name string
rule model.AlertRule
wantErr string
check func(t *testing.T, r model.AlertRule)
}{
{name: "名称必填", rule: model.AlertRule{Threshold: 1}, wantErr: "名称"},
{name: "模式非法", rule: model.AlertRule{Name: "r", Threshold: 1, SourceIPMode: "any"}, wantErr: "in/notin"},
{name: "阈值越界", rule: model.AlertRule{Name: "r", Threshold: 101}, wantErr: "阈值"},
{name: "阈值>1须带窗口", rule: model.AlertRule{Name: "r", Threshold: 3}, wantErr: "窗口"},
{name: "IP 非法", rule: model.AlertRule{Name: "r", Threshold: 1, SourceIPs: "300.1.1.1"}, wantErr: "IP"},
{name: "CIDR 合法", rule: model.AlertRule{Name: "r", Threshold: 1, SourceIPs: "10.0.0.0/8, 1.2.3.4"},
check: func(t *testing.T, r model.AlertRule) {
if r.SourceIPs != "10.0.0.0/8,1.2.3.4" {
t.Errorf("SourceIPs = %q, 应去空白归一化", r.SourceIPs)
}
if r.SourceIPMode != alertSourceIPIn {
t.Errorf("SourceIPMode = %q, 应默认 in", r.SourceIPMode)
}
}},
{name: "事件清单归一化", rule: model.AlertRule{Name: "r", Threshold: 1, EventTypes: " TerminateInstance , CreateApiKey ,"},
check: func(t *testing.T, r model.AlertRule) {
if r.EventTypes != "TerminateInstance,CreateApiKey" {
t.Errorf("EventTypes = %q", r.EventTypes)
}
}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
rule := tt.rule
err := validateAlertRule(&rule)
if tt.wantErr == "" {
if err != nil {
t.Fatalf("validateAlertRule: %v", err)
}
if tt.check != nil {
tt.check(t, rule)
}
return
}
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
t.Fatalf("err = %v, want contains %q", err, tt.wantErr)
}
})
}
}
func TestRuleHits(t *testing.T) {
base := model.AlertRule{Name: "r", Threshold: 1, SourceIPMode: alertSourceIPIn}
ev := &model.LogEvent{OciConfigID: 7}
parsed := parsedEvent{
EventType: "com.oraclecloud.ComputeApi.TerminateInstance",
SourceIP: "203.0.113.8",
ResourceName: "web-server-1",
}
tests := []struct {
name string
mod func(r *model.AlertRule)
p *parsedEvent
want bool
}{
{name: "空条件全命中", mod: func(r *model.AlertRule) {}, want: true},
{name: "租户匹配", mod: func(r *model.AlertRule) { r.OciConfigID = 7 }, want: true},
{name: "租户不匹配", mod: func(r *model.AlertRule) { r.OciConfigID = 8 }, want: false},
{name: "事件短名命中", mod: func(r *model.AlertRule) { r.EventTypes = "LaunchInstance,TerminateInstance" }, want: true},
{name: "事件不在清单", mod: func(r *model.AlertRule) { r.EventTypes = "CreateUser" }, want: false},
{name: "资源子串命中", mod: func(r *model.AlertRule) { r.ResourceMatch = "web-" }, want: true},
{name: "资源不含", mod: func(r *model.AlertRule) { r.ResourceMatch = "db-" }, want: false},
{name: "IP in 命中 CIDR", mod: func(r *model.AlertRule) { r.SourceIPs = "203.0.113.0/24" }, want: true},
{name: "IP in 未命中", mod: func(r *model.AlertRule) { r.SourceIPs = "10.0.0.0/8" }, want: false},
{name: "IP notin 白名单外告警", mod: func(r *model.AlertRule) {
r.SourceIPs, r.SourceIPMode = "10.0.0.0/8", alertSourceIPNotIn
}, want: true},
{name: "IP notin 白名单内不告警", mod: func(r *model.AlertRule) {
r.SourceIPs, r.SourceIPMode = "203.0.113.8", alertSourceIPNotIn
}, want: false},
{name: "notin 事件缺 IP 不告警", mod: func(r *model.AlertRule) {
r.SourceIPs, r.SourceIPMode = "10.0.0.0/8", alertSourceIPNotIn
}, p: &parsedEvent{EventType: parsed.EventType}, want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
rule := base
tt.mod(&rule)
p := parsed
if tt.p != nil {
p = *tt.p
}
if got := ruleHits(rule, ev, p); got != tt.want {
t.Errorf("ruleHits = %v, want %v", got, tt.want)
}
})
}
}
func TestAlertRuleCRUD(t *testing.T) {
svc, _, _ := newLogEventEnv(t)
ctx := context.Background()
created, err := svc.CreateAlertRule(ctx, model.AlertRule{Name: "非白名单终止", Enabled: true, Threshold: 1})
if err != nil {
t.Fatalf("create: %v", err)
}
if created.ID == 0 {
t.Fatal("create 未回填 ID")
}
if _, err := svc.CreateAlertRule(ctx, model.AlertRule{Threshold: 1}); err == nil {
t.Fatal("空名称应校验失败")
}
created.Enabled = false
created.EventTypes = "TerminateInstance"
updated, err := svc.UpdateAlertRule(ctx, created.ID, created)
if err != nil {
t.Fatalf("update: %v", err)
}
if updated.Enabled || updated.EventTypes != "TerminateInstance" {
t.Fatalf("update 未生效: %+v", updated)
}
rules, err := svc.ListAlertRules(ctx)
if err != nil || len(rules) != 1 {
t.Fatalf("list = %v, %v", rules, err)
}
if err := svc.DeleteAlertRule(ctx, created.ID); err != nil {
t.Fatalf("delete: %v", err)
}
if rules, _ := svc.ListAlertRules(ctx); len(rules) != 0 {
t.Fatalf("delete 后仍有 %d 条", len(rules))
}
}
// auditEventPayload 构造一条含资源与来源 IP 的 CloudEvents 审计消息。
func auditEventPayload(event, resource, ip string) string {
return fmt.Sprintf(`{"eventType":"com.oraclecloud.ComputeApi.%s","source":"ComputeApi",`+
`"eventTime":"2026-07-10T08:00:00Z","data":{"resourceName":%q,"identity":{"ipAddress":%q}}}`,
event, resource, ip)
}
// newAlertNotifyEnv 组装带假 Telegram 通道的告警测试环境。
func newAlertNotifyEnv(t *testing.T) (*LogEventService, *telegramCapture, func()) {
t.Helper()
svc, db, _ := newLogEventEnv(t)
srv, rec := newFakeTelegram(t, `{"ok":true}`)
cipher, err := crypto.NewCipher("test-data-key")
if err != nil {
t.Fatalf("new cipher: %v", err)
}
settings := NewSettingService(db, cipher)
token := "123456:AAfake"
if err := settings.UpdateTelegram(context.Background(),
UpdateTelegramInput{Enabled: true, BotToken: &token, ChatID: "42"}); err != nil {
t.Fatalf("update telegram: %v", err)
}
n := NewNotifier(settings)
n.base = srv.URL
svc.SetNotifier(n, settings)
return svc, rec, n.Wait
}
func TestMatchAlertRulesNotify(t *testing.T) {
svc, rec, wait := newAlertNotifyEnv(t)
ctx := context.Background()
_, err := svc.CreateAlertRule(ctx, model.AlertRule{
Name: "白名单外终止", Enabled: true, Threshold: 1,
EventTypes: "TerminateInstance", SourceIPs: "10.0.0.0/8", SourceIPMode: alertSourceIPNotIn,
})
if err != nil {
t.Fatalf("create rule: %v", err)
}
// 命中:白名单外 IP;不命中:白名单内 IP
mustIngest(t, svc, "m1", auditEventPayload("TerminateInstance", "web-1", "203.0.113.8"))
mustIngest(t, svc, "m2", auditEventPayload("TerminateInstance", "web-2", "10.1.2.3"))
svc.parseOnce(ctx)
wait()
alerts := auditAlerts(rec.snapshot())
joined := strings.Join(alerts, "\n---\n")
if !strings.Contains(joined, "白名单外终止") || !strings.Contains(joined, "web-1") {
t.Fatalf("应收到含规则名与资源的告警,got %q", joined)
}
if strings.Contains(joined, "web-2") {
t.Fatalf("白名单内事件不应告警,got %q", joined)
}
}
// auditAlerts 过滤出审计告警推送(排除既有 notifyCritical 的云端事件通知)。
func auditAlerts(texts []string) []string {
var out []string
for _, s := range texts {
if strings.Contains(s, "审计告警") {
out = append(out, s)
}
}
return out
}
func TestAlertThresholdWindow(t *testing.T) {
svc, rec, wait := newAlertNotifyEnv(t)
ctx := context.Background()
_, err := svc.CreateAlertRule(ctx, model.AlertRule{
Name: "登录风暴", Enabled: true, Threshold: 3, WindowMinutes: 5, EventTypes: "InteractiveLogin",
})
if err != nil {
t.Fatalf("create rule: %v", err)
}
for i := 1; i <= 4; i++ {
mustIngest(t, svc, fmt.Sprint("login-", i),
auditEventPayload("InteractiveLogin", "user@x.com", "203.0.113.8"))
}
svc.parseOnce(ctx)
wait()
alerts := auditAlerts(rec.snapshot())
if len(alerts) != 1 {
t.Fatalf("窗口内 4 次命中应只告警 1 次(第 3 次触发后冷却),got %d 条: %v", len(alerts), alerts)
}
if !strings.Contains(alerts[0], "3 次") {
t.Errorf("告警文案应含累计次数,got %q", alerts[0])
}
}
// TestAlertRuleBadDataDoesNotBlockParse 验证规则表异常不影响解析主流程。
func TestAlertRuleBadDataDoesNotBlockParse(t *testing.T) {
svc, db, _ := newLogEventEnv(t)
ctx := context.Background()
// 直插一条绕过校验的坏规则(IP 列表非法)
bad := model.AlertRule{Name: "bad", Enabled: true, Threshold: 1, SourceIPs: "not-an-ip"}
if err := db.Create(&bad).Error; err != nil {
t.Fatalf("insert bad rule: %v", err)
}
mustIngest(t, svc, "m1", auditEventPayload("TerminateInstance", "web-1", "1.2.3.4"))
svc.parseOnce(ctx)
var e model.LogEvent
if err := db.First(&e, "message_id = ?", "m1").Error; err != nil {
t.Fatalf("find event: %v", err)
}
if !e.Processed {
t.Fatal("坏规则不应阻塞事件解析")
}
}
// mustIngest 落一条回传事件,失败即终止测试。
func mustIngest(t *testing.T, svc *LogEventService, msgID, payload string) {
t.Helper()
if err := svc.Ingest(context.Background(), 1, msgID, []byte(payload), false); err != nil {
t.Fatalf("ingest %s: %v", msgID, err)
}
}
// TestCleanupAlertHits 验证过期命中记录随清理删除。
func TestCleanupAlertHits(t *testing.T) {
svc, db, _ := newLogEventEnv(t)
old := model.AlertRuleHit{RuleID: 1, HitAt: time.Now().Add(-8 * 24 * time.Hour)}
fresh := model.AlertRuleHit{RuleID: 1, HitAt: time.Now()}
if err := db.Create(&old).Error; err != nil {
t.Fatalf("insert: %v", err)
}
if err := db.Create(&fresh).Error; err != nil {
t.Fatalf("insert: %v", err)
}
svc.cleanupAlertHits(context.Background())
var count int64
db.Model(&model.AlertRuleHit{}).Count(&count)
if count != 1 {
t.Fatalf("清理后应剩 1 条,got %d", count)
}
}
func TestClearAlertCooldown(t *testing.T) {
svc := NewLogEventService(nil)
rule1 := model.AlertRule{ID: 1, Threshold: 2, WindowMinutes: 10}
rule2 := model.AlertRule{ID: 2, Threshold: 2, WindowMinutes: 10}
if !svc.alertCooldownPass(rule1) || !svc.alertCooldownPass(rule2) {
t.Fatal("首次命中应通过冷却检查")
}
svc.ClearAlertCooldown([]uint{rule1.ID})
if !svc.alertCooldownPass(rule1) {
t.Fatal("已清理规则应重新通过冷却检查")
}
if svc.alertCooldownPass(rule2) {
t.Fatal("未清理规则不应通过冷却检查")
}
}
func TestRecordAlertHitRejectsMissingRefs(t *testing.T) {
tests := []struct {
name string
deleteRule bool
}{
{name: "规则已删除", deleteRule: true},
{name: "事件已删除"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
testRecordAlertHitMissingRef(t, tt.deleteRule)
})
}
}
func testRecordAlertHitMissingRef(t *testing.T, deleteRule bool) {
t.Helper()
svc, db, cfgID := newLogEventEnv(t)
rule := model.AlertRule{Name: "r", OciConfigID: cfgID, Threshold: 2, WindowMinutes: 5}
event := model.LogEvent{OciConfigID: cfgID, MessageID: "m"}
if err := db.Create(&rule).Error; err != nil {
t.Fatalf("create rule: %v", err)
}
if err := db.Create(&event).Error; err != nil {
t.Fatalf("create event: %v", err)
}
var err error
if deleteRule {
err = db.Delete(&rule).Error
} else {
err = db.Delete(&event).Error
}
if err != nil {
t.Fatalf("delete reference: %v", err)
}
if count, ok := svc.recordAlertHit(context.Background(), rule, &event); ok || count != 0 {
t.Fatalf("record missing refs = (%d,%v), want (0,false)", count, ok)
}
var hits int64
db.Model(&model.AlertRuleHit{}).Count(&hits)
if hits != 0 {
t.Fatalf("orphan alert hits = %d, want 0", hits)
}
}
+528
View File
@@ -0,0 +1,528 @@
package service
import (
"encoding/json"
"fmt"
"strings"
"oci-portal/internal/aiwire"
)
// Anthropic Messages ↔ OCI OpenAI 兼容面(/actions/v1/responses)直通转换。
// typed chat 面剔除后 Messages 入口的唯一上游通路。语义损失(README 已披露):
// stop_sequences / top_k / metadata / thinking 无对应字段,忽略;上游 reasoning
// 输出项与增量事件丢弃(Anthropic thinking 块含签名语义,不伪造)。
// ErrAiUnsupportedBlock 表示请求含网关无法承接的内容块(仅支持文本/图片/工具块)。
var ErrAiUnsupportedBlock = fmt.Errorf("暂不支持文本与图片以外的内容块")
// AnthropicToResponsesBody 把 Messages 请求转为直通 body(强制 store:false)。
func AnthropicToResponsesBody(req aiwire.MessagesRequest) ([]byte, error) {
input, err := anthInputItems(req.Messages)
if err != nil {
return nil, err
}
body := map[string]any{"model": req.Model, "max_output_tokens": req.MaxTokens,
"input": input, "store": false}
if sys := req.SystemText(); sys != "" {
body["instructions"] = sys
}
if req.Temperature != nil {
body["temperature"] = *req.Temperature
}
if req.TopP != nil {
body["top_p"] = *req.TopP
}
if tools := anthRespTools(req.Tools); tools != nil {
body["tools"] = tools
}
if tc := anthRespToolChoice(req.ToolChoice); tc != nil {
body["tool_choice"] = tc
}
if req.OutputConfig != nil && req.OutputConfig.Effort != "" {
body["reasoning"] = map[string]string{"effort": strings.ToLower(req.OutputConfig.Effort)}
}
if req.Stream {
body["stream"] = true
}
return json.Marshal(body)
}
// anthInputItems 把消息序列展开为 Responses input 项:tool_use / tool_result 为
// 独立 function_call / function_call_output 项,其余聚合为 message 项;遇独立项时
// 先冲刷已聚合部件,保持块间相对顺序。
func anthInputItems(messages []aiwire.AnthMessage) ([]any, error) {
var items []any
for _, m := range messages {
var parts []map[string]any
flush := func() {
if len(parts) > 0 {
items = append(items, map[string]any{"role": m.Role, "content": parts})
parts = nil
}
}
for _, b := range m.Content.AllBlocks() {
item, part, err := anthBlockItem(m.Role, b)
if err != nil {
return nil, err
}
if part != nil {
parts = append(parts, part)
}
if item != nil {
flush()
items = append(items, item)
}
}
flush()
}
return items, nil
}
// anthBlockItem 把单个内容块转为独立项或消息部件(thinking 忽略,未知块拒绝)。
func anthBlockItem(role string, b aiwire.AnthBlock) (item, part map[string]any, err error) {
switch b.Type {
case "text":
return nil, respTextPart(role, b.Text), nil
case "image":
part, err = anthImageInput(b.Source)
return nil, part, err
case "tool_use":
return map[string]any{"type": "function_call", "call_id": b.ID,
"name": b.Name, "arguments": string(b.Input)}, nil, nil
case "tool_result":
return map[string]any{"type": "function_call_output",
"call_id": b.ToolUseID, "output": b.ResultText()}, nil, nil
case "thinking", "redacted_thinking":
return nil, nil, nil
default:
return nil, nil, ErrAiUnsupportedBlock
}
}
// respTextPart 按角色选择 Responses 文本部件类型(assistant 历史为 output_text);
// Anthropic 与 Chat Completions 两条转换链共用。
func respTextPart(role, text string) map[string]any {
if role == "assistant" {
return map[string]any{"type": "output_text", "text": text}
}
return map[string]any{"type": "input_text", "text": text}
}
// anthImageInput 把 Anthropic image source 转为 Responses input_image 部件。
func anthImageInput(source json.RawMessage) (map[string]any, error) {
var src struct {
Type string `json:"type"`
MediaType string `json:"media_type"`
Data string `json:"data"`
URL string `json:"url"`
}
if err := json.Unmarshal(source, &src); err != nil {
return nil, fmt.Errorf("image source 解析失败: %w", err)
}
switch src.Type {
case "base64":
if src.MediaType == "" || src.Data == "" {
return nil, fmt.Errorf("image source 缺少 media_type 或 data")
}
return map[string]any{"type": "input_image",
"image_url": "data:" + src.MediaType + ";base64," + src.Data}, nil
case "url":
if src.URL == "" {
return nil, fmt.Errorf("image source 缺少 url")
}
return map[string]any{"type": "input_image", "image_url": src.URL}, nil
}
return nil, fmt.Errorf("不支持的 image source 类型 %q", src.Type)
}
func anthRespTools(tools []aiwire.AnthTool) []map[string]any {
if len(tools) == 0 {
return nil
}
out := make([]map[string]any, 0, len(tools))
for _, t := range tools {
out = append(out, map[string]any{"type": "function", "name": t.Name,
"description": t.Description, "parameters": t.InputSchema})
}
return out
}
// anthRespToolChoice 映射 {type:auto|any|tool|none,name} → Responses 形态。
func anthRespToolChoice(raw json.RawMessage) any {
if len(raw) == 0 {
return nil
}
var tc struct {
Type string `json:"type"`
Name string `json:"name"`
}
if json.Unmarshal(raw, &tc) != nil {
return nil
}
switch tc.Type {
case "auto":
return "auto"
case "any":
return "required"
case "none":
return "none"
case "tool":
return map[string]string{"type": "function", "name": tc.Name}
}
return nil
}
// respPayload 是直通响应中本转换关心的子集(未知字段忽略)。
type respPayload struct {
Model string `json:"model"`
Status string `json:"status"`
IncompleteDetails *respIncomplete `json:"incomplete_details"`
Output []respOutputItem `json:"output"`
Usage *aiwire.RespUsage `json:"usage"`
Error *map[string]string `json:"error"`
}
type respIncomplete struct {
Reason string `json:"reason"`
}
type respOutputItem struct {
Type string `json:"type"`
CallID string `json:"call_id"`
Name string `json:"name"`
Arguments string `json:"arguments"`
Content []struct {
Type string `json:"type"`
Text string `json:"text"`
} `json:"content"`
}
// ResponsesToAnthropic 把直通非流式响应转为 Anthropic Messages 响应。
func ResponsesToAnthropic(payload []byte, msgID string) (*aiwire.MessagesResponse, error) {
var resp respPayload
if err := json.Unmarshal(payload, &resp); err != nil {
return nil, fmt.Errorf("解析上游响应: %w", err)
}
out := &aiwire.MessagesResponse{ID: msgID, Type: "message", Role: "assistant",
Model: resp.Model, Content: []aiwire.AnthBlock{}, StopReason: "end_turn"}
for _, item := range resp.Output {
switch item.Type {
case "message":
for _, part := range item.Content {
if part.Type == "output_text" && part.Text != "" {
out.Content = append(out.Content, aiwire.AnthBlock{Type: "text", Text: part.Text})
}
}
case "function_call":
out.Content = append(out.Content, aiwire.AnthBlock{Type: "tool_use", ID: item.CallID,
Name: item.Name, Input: argsToJSON(item.Arguments)})
out.StopReason = "tool_use"
}
}
if respTruncated(&resp) {
out.StopReason = "max_tokens"
}
out.Usage = anthUsageFromResp(resp.Usage)
return out, nil
}
func anthUsageFromResp(u *aiwire.RespUsage) aiwire.AnthUsage {
if u == nil {
return aiwire.AnthUsage{}
}
return aiwire.AnthUsage{InputTokens: u.InputTokens, OutputTokens: u.OutputTokens,
CacheReadInputTokens: u.InputTokensDetails.CachedTokens}
}
// argsToJSON 保证 tool_use.input 是合法 JSON 对象(模型可能产出非法片段)。
func argsToJSON(args string) json.RawMessage {
trimmed := strings.TrimSpace(args)
if trimmed == "" {
return json.RawMessage(`{}`)
}
if json.Valid([]byte(trimmed)) {
return json.RawMessage(trimmed)
}
b, _ := json.Marshal(map[string]string{"_raw": args})
return b
}
// ---- Anthropic 流式桥:Responses SSE 事件 → Anthropic 事件序列 ----
// AnthEvent 是一条待写出的 Anthropic SSE 事件。
type AnthEvent struct {
Event string
Data any
}
// AnthRespBridge 把直通 SSE 事件流桥接为 Anthropic 事件序列:
// message_start → content_block_start/delta/stop(text 与 tool_use 分块)→ message_delta → message_stop。
// reasoning 系列事件丢弃;上游 error / response.failed 转 Anthropic error 事件透传。
type AnthRespBridge struct {
id, model string
started bool
blockOpen bool
blockIsTool bool
blockIndex int
stopReason string
usage aiwire.AnthUsage
// sawTerminal 标记收到过终态事件(completed/incomplete/failed);
// 上游流提前 EOF 时据此发 error 事件而非伪装正常结束
sawTerminal bool
// errMsg 记录上游错误消息,非空即本流已失败(供调用日志)
errMsg string
}
// NewAnthRespBridge 构造桥;id 为响应消息 ID。
func NewAnthRespBridge(id, model string) *AnthRespBridge {
return &AnthRespBridge{id: id, model: model, blockIndex: -1, stopReason: "end_turn"}
}
// respStreamEvent 是直通 SSE data JSON 中桥关心的子集。
type respStreamEvent struct {
Type string `json:"type"`
Delta string `json:"delta"`
Message string `json:"message"`
Item *respOutputItem `json:"item"`
Response *respPayload `json:"response"`
}
// Feed 消费一行 SSE data JSON,返回应立即写出的事件。
func (st *AnthRespBridge) Feed(data []byte) []AnthEvent {
var ev respStreamEvent
if json.Unmarshal(data, &ev) != nil {
return nil
}
if ev.Type == "error" || ev.Type == "response.failed" {
return st.failWith(ev)
}
// message_start 延迟到首个可见输出事件:created / reasoning 阶段不向客户端
// 写任何字节,上游此段断流时 handler 才有降级非流式重做的无感窗口
var events []AnthEvent
switch ev.Type {
case "response.output_item.added":
if ev.Item != nil && ev.Item.Type == "function_call" {
events = append(events, st.ensureStarted()...)
events = append(events, st.openBlock(true, ev.Item.CallID, ev.Item.Name)...)
st.stopReason = "tool_use"
}
case "response.output_text.delta":
events = append(events, st.ensureStarted()...)
events = append(events, st.textDelta(ev.Delta)...)
case "response.function_call_arguments.delta":
events = append(events, st.argsDelta(ev.Delta)...)
case "response.completed", "response.incomplete":
st.sawTerminal = true
st.finishFrom(ev.Response)
}
return events
}
// ensureStarted 在首个可见输出事件前补发 message_start(仅一次)。
func (st *AnthRespBridge) ensureStarted() []AnthEvent {
if st.started {
return nil
}
st.started = true
return []AnthEvent{st.startEvent()}
}
// failWith 记录上游错误并产出 Anthropic error 事件(每流至多一次)。
func (st *AnthRespBridge) failWith(ev respStreamEvent) []AnthEvent {
if st.errMsg != "" {
return nil
}
st.sawTerminal = true
msg := ev.Message
if ev.Type == "response.failed" {
st.finishFrom(ev.Response)
if m := respErrorMsg(ev.Response); m != "" {
msg = m
}
}
if msg == "" {
msg = "上游返回错误事件 " + ev.Type
}
st.errMsg = msg
return []AnthEvent{st.errorEvent()}
}
// respErrorMsg 提取 response.failed 载荷中的错误消息。
func respErrorMsg(resp *respPayload) string {
if resp == nil || resp.Error == nil {
return ""
}
return (*resp.Error)["message"]
}
// errorEvent 按 Anthropic 流式协议产出 error 事件。
func (st *AnthRespBridge) errorEvent() AnthEvent {
return AnthEvent{Event: "error", Data: map[string]any{
"type": "error",
"error": map[string]string{"type": "api_error", "message": st.errMsg},
}}
}
// Err 返回上游错误消息;空串表示流正常(供调用日志)。
func (st *AnthRespBridge) Err() string { return st.errMsg }
// SawTerminal 报告是否收到过终态事件;false 即上游流提前终止。
func (st *AnthRespBridge) SawTerminal() bool { return st.sawTerminal }
// AnthMessageEvents 把完整 Messages 响应展开为标准事件序列,
// 供流式上游断流后的非流式降级结果推送(客户端协议不变)。
func AnthMessageEvents(m *aiwire.MessagesResponse) []AnthEvent {
events := []AnthEvent{{Event: "message_start", Data: map[string]any{
"type": "message_start",
"message": map[string]any{
"id": m.ID, "type": "message", "role": m.Role, "model": m.Model,
"content": []any{}, "stop_reason": nil,
"usage": map[string]int{"input_tokens": 0, "output_tokens": 0},
},
}}}
for i, b := range m.Content {
events = append(events, anthBlockEvents(i, b)...)
}
usage := map[string]int{"input_tokens": m.Usage.InputTokens, "output_tokens": m.Usage.OutputTokens}
if m.Usage.CacheReadInputTokens > 0 {
usage["cache_read_input_tokens"] = m.Usage.CacheReadInputTokens
}
events = append(events,
AnthEvent{Event: "message_delta", Data: map[string]any{
"type": "message_delta",
"delta": map[string]any{"stop_reason": m.StopReason, "stop_sequence": nil},
"usage": usage,
}},
AnthEvent{Event: "message_stop", Data: map[string]any{"type": "message_stop"}})
return events
}
// anthBlockEvents 把一个内容块展开为 start / delta / stop 三事件。
func anthBlockEvents(index int, b aiwire.AnthBlock) []AnthEvent {
var start, delta map[string]any
if b.Type == "tool_use" {
start = map[string]any{"type": "tool_use", "id": b.ID, "name": b.Name, "input": map[string]any{}}
delta = map[string]any{"type": "input_json_delta", "partial_json": string(b.Input)}
} else {
start = map[string]any{"type": "text", "text": ""}
delta = map[string]any{"type": "text_delta", "text": b.Text}
}
return []AnthEvent{
{Event: "content_block_start", Data: map[string]any{
"type": "content_block_start", "index": index, "content_block": start}},
{Event: "content_block_delta", Data: map[string]any{
"type": "content_block_delta", "index": index, "delta": delta}},
{Event: "content_block_stop", Data: map[string]any{
"type": "content_block_stop", "index": index}},
}
}
func (st *AnthRespBridge) textDelta(delta string) []AnthEvent {
var events []AnthEvent
if !st.blockOpen || st.blockIsTool {
events = append(events, st.openBlock(false, "", "")...)
}
events = append(events, AnthEvent{Event: "content_block_delta", Data: map[string]any{
"type": "content_block_delta", "index": st.blockIndex,
"delta": map[string]string{"type": "text_delta", "text": delta},
}})
return events
}
func (st *AnthRespBridge) argsDelta(delta string) []AnthEvent {
if !st.blockOpen || !st.blockIsTool {
return nil
}
return []AnthEvent{{Event: "content_block_delta", Data: map[string]any{
"type": "content_block_delta", "index": st.blockIndex,
"delta": map[string]string{"type": "input_json_delta", "partial_json": delta},
}}}
}
// finishFrom 记录终态:usage 与 stop_reason(max_output_tokens 截断 → max_tokens)。
func (st *AnthRespBridge) finishFrom(resp *respPayload) {
if resp == nil {
return
}
st.usage = anthUsageFromResp(resp.Usage)
if respTruncated(resp) {
st.stopReason = "max_tokens"
}
}
func (st *AnthRespBridge) startEvent() AnthEvent {
return AnthEvent{Event: "message_start", Data: map[string]any{
"type": "message_start",
"message": map[string]any{
"id": st.id, "type": "message", "role": "assistant", "model": st.model,
"content": []any{}, "stop_reason": nil,
"usage": map[string]int{"input_tokens": 0, "output_tokens": 0},
},
}}
}
// openBlock 关闭当前块并打开新块(text 或 tool_use)。
func (st *AnthRespBridge) openBlock(isTool bool, toolID, toolName string) []AnthEvent {
var events []AnthEvent
if st.blockOpen {
events = append(events, st.closeBlockEvent())
}
st.blockOpen, st.blockIsTool = true, isTool
st.blockIndex++
block := map[string]any{"type": "text", "text": ""}
if isTool {
block = map[string]any{"type": "tool_use", "id": toolID, "name": toolName, "input": map[string]any{}}
}
events = append(events, AnthEvent{Event: "content_block_start", Data: map[string]any{
"type": "content_block_start", "index": st.blockIndex, "content_block": block,
}})
return events
}
func (st *AnthRespBridge) closeBlockEvent() AnthEvent {
return AnthEvent{Event: "content_block_stop", Data: map[string]any{
"type": "content_block_stop", "index": st.blockIndex,
}}
}
// Finish 在上游流结束后收尾:关块 → message_delta(stop_reason+usage)→ message_stop。
// 已发过 error 事件的流不再补终态;未见终态事件即 EOF 视为上游提前终止,
// 发 error 事件而非伪装正常结束(否则客户端拿到"成功的空消息")。
func (st *AnthRespBridge) Finish() []AnthEvent {
if st.errMsg != "" {
return nil
}
if !st.sawTerminal {
st.errMsg = "上游流提前终止,未返回终态事件"
return []AnthEvent{st.errorEvent()}
}
var events []AnthEvent
if !st.started {
st.started = true
events = append(events, st.startEvent())
}
if st.blockOpen {
events = append(events, st.closeBlockEvent())
st.blockOpen = false
}
usage := map[string]int{"output_tokens": st.usage.OutputTokens}
if st.usage.InputTokens > 0 {
usage["input_tokens"] = st.usage.InputTokens
}
if st.usage.CacheReadInputTokens > 0 {
usage["cache_read_input_tokens"] = st.usage.CacheReadInputTokens
}
events = append(events,
AnthEvent{Event: "message_delta", Data: map[string]any{
"type": "message_delta",
"delta": map[string]any{"stop_reason": st.stopReason, "stop_sequence": nil},
"usage": usage,
}},
AnthEvent{Event: "message_stop", Data: map[string]any{"type": "message_stop"}},
)
return events
}
// Usage 返回聚合到的用量(供调用日志)。
func (st *AnthRespBridge) Usage() aiwire.AnthUsage { return st.usage }
+268
View File
@@ -0,0 +1,268 @@
package service
import (
"encoding/json"
"strings"
"testing"
"oci-portal/internal/aiwire"
)
func mustAnthReq(t *testing.T, raw string) aiwire.MessagesRequest {
t.Helper()
var req aiwire.MessagesRequest
if err := json.Unmarshal([]byte(raw), &req); err != nil {
t.Fatalf("解析请求: %v", err)
}
return req
}
// TestAnthropicToResponsesBody 断言 system/消息/工具/effort 的直通装配与 store 强制。
func TestAnthropicToResponsesBody(t *testing.T) {
req := mustAnthReq(t, `{
"model": "xai.grok-4.3", "max_tokens": 128, "system": "你是助手", "stream": true,
"temperature": 0.5,
"messages": [
{"role": "user", "content": "东京天气?"},
{"role": "assistant", "content": [
{"type": "text", "text": "查询中"},
{"type": "tool_use", "id": "t1", "name": "get_weather", "input": {"city": "东京"}}
]},
{"role": "user", "content": [
{"type": "tool_result", "tool_use_id": "t1", "content": "晴 25 度"},
{"type": "text", "text": "继续"}
]}
],
"tools": [{"name": "get_weather", "description": "查天气", "input_schema": {"type": "object"}}],
"tool_choice": {"type": "any"},
"output_config": {"effort": "HIGH"}
}`)
payload, err := AnthropicToResponsesBody(req)
if err != nil {
t.Fatalf("AnthropicToResponsesBody: %v", err)
}
var body map[string]any
if err := json.Unmarshal(payload, &body); err != nil {
t.Fatalf("unmarshal: %v", err)
}
if body["model"] != "xai.grok-4.3" || body["max_output_tokens"] != float64(128) ||
body["instructions"] != "你是助手" || body["store"] != false || body["stream"] != true {
t.Fatalf("顶层字段装配错误: %v", body)
}
if body["tool_choice"] != "required" {
t.Fatalf("tool_choice = %v, want required", body["tool_choice"])
}
reasoning, _ := body["reasoning"].(map[string]any)
if reasoning["effort"] != "high" {
t.Fatalf("effort = %v, want high(小写透传)", reasoning)
}
input, _ := body["input"].([]any)
// user 消息、assistant 文本消息、function_call、function_call_output、末条 user 文本
if len(input) != 5 {
t.Fatalf("input 项数 = %d, want 5: %s", len(input), payload)
}
kinds := make([]string, 0, len(input))
for _, it := range input {
m := it.(map[string]any)
if ty, ok := m["type"].(string); ok {
kinds = append(kinds, ty)
} else {
kinds = append(kinds, "message/"+m["role"].(string))
}
}
want := "message/user,message/assistant,function_call,function_call_output,message/user"
if got := strings.Join(kinds, ","); got != want {
t.Fatalf("input 顺序 = %s, want %s", got, want)
}
if !strings.Contains(string(payload), `"output_text"`) {
t.Fatalf("assistant 历史文本应为 output_text 部件: %s", payload)
}
}
// TestAnthropicToResponsesBodyRejects 断言不支持的内容块拒绝与图片装配。
func TestAnthropicToResponsesBodyRejects(t *testing.T) {
bad := mustAnthReq(t, `{"model":"m","max_tokens":8,"messages":[
{"role":"user","content":[{"type":"document","source":{}}]}]}`)
if _, err := AnthropicToResponsesBody(bad); err == nil {
t.Fatal("document 块应拒绝")
}
img := mustAnthReq(t, `{"model":"m","max_tokens":8,"messages":[
{"role":"user","content":[{"type":"image","source":{"type":"base64","media_type":"image/png","data":"QUJD"}}]}]}`)
payload, err := AnthropicToResponsesBody(img)
if err != nil || !strings.Contains(string(payload), "data:image/png;base64,QUJD") {
t.Fatalf("图片应转 data URI: %v %s", err, payload)
}
}
// TestResponsesToAnthropic 断言输出块、stop_reason 与 usage 的回转。
func TestResponsesToAnthropic(t *testing.T) {
payload := []byte(`{"model":"xai.grok-4.3","status":"completed","output":[
{"type":"reasoning","summary":[]},
{"type":"message","content":[{"type":"output_text","text":"你好"}]},
{"type":"function_call","call_id":"c1","name":"get_weather","arguments":"{\"city\":\"东京\"}"}],
"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15,"input_tokens_details":{"cached_tokens":4}}}`)
out, err := ResponsesToAnthropic(payload, "msg_1")
if err != nil {
t.Fatalf("ResponsesToAnthropic: %v", err)
}
if len(out.Content) != 2 || out.Content[0].Type != "text" || out.Content[0].Text != "你好" {
t.Fatalf("content 装配错误: %+v", out.Content)
}
if out.Content[1].Type != "tool_use" || out.Content[1].ID != "c1" || string(out.Content[1].Input) != `{"city":"东京"}` {
t.Fatalf("tool_use 装配错误: %+v", out.Content[1])
}
if out.StopReason != "tool_use" {
t.Fatalf("stop_reason = %s, want tool_use", out.StopReason)
}
if out.Usage.InputTokens != 10 || out.Usage.OutputTokens != 5 || out.Usage.CacheReadInputTokens != 4 {
t.Fatalf("usage = %+v", out.Usage)
}
trunc := []byte(`{"model":"m","status":"incomplete","incomplete_details":{"reason":"max_output_tokens"},
"output":[{"type":"message","content":[{"type":"output_text","text":"半"}]}]}`)
out2, err := ResponsesToAnthropic(trunc, "msg_2")
if err != nil || out2.StopReason != "max_tokens" {
t.Fatalf("截断 stop_reason = %v, %v", out2, err)
}
}
func bridgeEventTypes(events []AnthEvent) string {
kinds := make([]string, 0, len(events))
for _, ev := range events {
kinds = append(kinds, ev.Event)
}
return strings.Join(kinds, ",")
}
// TestAnthRespBridge 断言流桥:文本增量、工具调用切块、reasoning 丢弃与收尾事件。
func TestAnthRespBridge(t *testing.T) {
st := NewAnthRespBridge("msg_1", "m1")
var events []AnthEvent
feed := func(lines ...string) {
for _, l := range lines {
events = append(events, st.Feed([]byte(l))...)
}
}
feed(`{"type":"response.created","response":{"model":"m1"}}`,
`{"type":"response.reasoning_text.delta","delta":"思考中"}`,
`{"type":"response.output_text.delta","delta":"你"}`,
`{"type":"response.output_text.delta","delta":"好"}`,
`{"type":"response.output_item.added","item":{"type":"function_call","call_id":"c1","name":"f"}}`,
`{"type":"response.function_call_arguments.delta","delta":"{\"a\":1}"}`,
`{"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":8,"output_tokens":4}}}`)
events = append(events, st.Finish()...)
got := bridgeEventTypes(events)
want := "message_start,content_block_start,content_block_delta,content_block_delta," +
"content_block_stop,content_block_start,content_block_delta,content_block_stop," +
"message_delta,message_stop"
if got != want {
t.Fatalf("事件序列:\n got %s\nwant %s", got, want)
}
if st.Usage().InputTokens != 8 || st.Usage().OutputTokens != 4 {
t.Fatalf("usage = %+v", st.Usage())
}
}
// TestAnthRespBridgeFailure 断言异常流的错误透传:上游 error / response.failed
// 事件转 Anthropic error 事件,提前 EOF(未见终态)不再伪装正常结束。
func TestAnthRespBridgeFailure(t *testing.T) {
cases := []struct {
name string
lines []string
wantKinds string
wantErr string
}{
{
name: "上游 error 事件透传",
lines: []string{`{"type":"error","message":"model overloaded"}`},
wantKinds: "error",
wantErr: "model overloaded",
},
{
name: "response.failed 提取错误消息",
lines: []string{
`{"type":"response.output_text.delta","delta":"你"}`,
`{"type":"response.failed","response":{"status":"failed","error":{"message":"content filtered"}}}`,
},
wantKinds: "message_start,content_block_start,content_block_delta,error",
wantErr: "content filtered",
},
{
name: "空流提前终止",
lines: nil,
wantKinds: "error",
wantErr: "上游流提前终止,未返回终态事件",
},
{
name: "输出中途 EOF 无终态",
lines: []string{
`{"type":"response.output_text.delta","delta":"你"}`,
},
wantKinds: "message_start,content_block_start,content_block_delta,error",
wantErr: "上游流提前终止,未返回终态事件",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
st := NewAnthRespBridge("msg_1", "m1")
var events []AnthEvent
for _, l := range tc.lines {
events = append(events, st.Feed([]byte(l))...)
}
events = append(events, st.Finish()...)
if got := bridgeEventTypes(events); got != tc.wantKinds {
t.Fatalf("事件序列 = %s, want %s", got, tc.wantKinds)
}
if st.Err() != tc.wantErr {
t.Fatalf("Err() = %q, want %q", st.Err(), tc.wantErr)
}
})
}
}
// TestRespStreamErrorMsg 断言直通流错误事件消息提取。
func TestRespStreamErrorMsg(t *testing.T) {
cases := []struct {
name string
data string
want string
}{
{name: "error 事件", data: `{"type":"error","message":"boom"}`, want: "boom"},
{name: "error 无消息用占位", data: `{"type":"error"}`, want: "上游返回错误事件 error"},
{name: "failed 事件", data: `{"type":"response.failed","response":{"error":{"message":"bad"}}}`, want: "bad"},
{name: "正常事件返回空", data: `{"type":"response.completed","response":{}}`, want: ""},
{name: "非 JSON 返回空", data: `<html>`, want: ""},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if got := RespStreamErrorMsg([]byte(tc.data)); got != tc.want {
t.Fatalf("RespStreamErrorMsg = %q, want %q", got, tc.want)
}
})
}
}
// TestAnthMessageEvents 断言非流式降级结果展开的事件序列与 usage。
func TestAnthMessageEvents(t *testing.T) {
msg := &aiwire.MessagesResponse{
ID: "msg_1", Type: "message", Role: "assistant", Model: "m1",
Content: []aiwire.AnthBlock{
{Type: "text", Text: "好"},
{Type: "tool_use", ID: "c1", Name: "f", Input: json.RawMessage(`{"a":1}`)},
},
StopReason: "tool_use",
Usage: aiwire.AnthUsage{InputTokens: 9, OutputTokens: 3, CacheReadInputTokens: 5},
}
events := AnthMessageEvents(msg)
want := "message_start,content_block_start,content_block_delta,content_block_stop," +
"content_block_start,content_block_delta,content_block_stop,message_delta,message_stop"
if got := bridgeEventTypes(events); got != want {
t.Fatalf("事件序列:\n got %s\nwant %s", got, want)
}
delta := events[len(events)-2].Data.(map[string]any)
usage := delta["usage"].(map[string]int)
if usage["input_tokens"] != 9 || usage["output_tokens"] != 3 || usage["cache_read_input_tokens"] != 5 {
t.Fatalf("usage = %+v", usage)
}
}
+389
View File
@@ -0,0 +1,389 @@
package service
import (
"encoding/json"
"fmt"
"strings"
"oci-portal/internal/aiwire"
)
// OpenAI Chat Completions ↔ OCI OpenAI 兼容面(/actions/v1/responses)直通转换。
// 端点定位 Tier 2 兼容层(承接只会说 CC 的存量客户端),机制与 anthresponses.go
// 同构。语义损失(README 披露):stop / seed / n / penalty 等无对应字段,忽略;
// 上游 reasoning 输出项与增量事件丢弃;非 function 工具类型拒绝。
// ChatToResponsesBody 把 Chat Completions 请求转为直通 body(强制 store:false)。
func ChatToResponsesBody(req aiwire.ChatRequest) ([]byte, error) {
input, instructions, err := chatInputItems(req.Messages)
if err != nil {
return nil, err
}
body := map[string]any{"model": req.Model, "input": input, "store": false}
if instructions != "" {
body["instructions"] = instructions
}
chatSampling(body, req)
if err := chatBodyTools(body, req); err != nil {
return nil, err
}
if tf := chatTextFormat(req.ResponseFormat); tf != nil {
body["text"] = map[string]any{"format": tf}
}
if req.ReasoningEffort != "" {
body["reasoning"] = map[string]string{"effort": strings.ToLower(req.ReasoningEffort)}
}
if req.Stream {
body["stream"] = true
}
return json.Marshal(body)
}
// chatSampling 装配采样与输出预算;max_completion_tokens 优先于已弃用的 max_tokens。
func chatSampling(body map[string]any, req aiwire.ChatRequest) {
if req.MaxCompletionTokens != nil {
body["max_output_tokens"] = *req.MaxCompletionTokens
} else if req.MaxTokens != nil {
body["max_output_tokens"] = *req.MaxTokens
}
if req.Temperature != nil {
body["temperature"] = *req.Temperature
}
if req.TopP != nil {
body["top_p"] = *req.TopP
}
if req.ParallelToolCalls != nil {
body["parallel_tool_calls"] = *req.ParallelToolCalls
}
}
// chatBodyTools 装配工具与 tool_choice;非 function 工具类型拒绝
// (web_search 等新能力仅在 Responses / Messages 端点供给)。
func chatBodyTools(body map[string]any, req aiwire.ChatRequest) error {
if len(req.Tools) == 0 {
return nil
}
tools := make([]map[string]any, 0, len(req.Tools))
for _, t := range req.Tools {
if t.Type != "function" {
return fmt.Errorf("不支持的工具类型 %q:该端点仅支持 function 工具", t.Type)
}
tools = append(tools, map[string]any{"type": "function", "name": t.Function.Name,
"description": t.Function.Description, "parameters": t.Function.Parameters})
}
body["tools"] = tools
if tc := chatRespToolChoice(req.ToolChoice); tc != nil {
body["tool_choice"] = tc
}
return nil
}
// chatInputItems 把消息序列展开为 Responses input 项与 instructions:
// system/developer 文本聚为 instructions;tool 消息为 function_call_output 项;
// 其余消息经 chatMessageItems 展开,保持相对顺序。
func chatInputItems(messages []aiwire.ChatMessage) ([]any, string, error) {
items := make([]any, 0, len(messages))
var sys []string
for _, m := range messages {
switch m.Role {
case "system", "developer":
if txt := m.Content.JoinText(); txt != "" {
sys = append(sys, txt)
}
case "tool":
items = append(items, map[string]any{"type": "function_call_output",
"call_id": m.ToolCallID, "output": m.Content.JoinText()})
default:
msgItems, err := chatMessageItems(m)
if err != nil {
return nil, "", err
}
items = append(items, msgItems...)
}
}
return items, strings.Join(sys, "\n\n"), nil
}
// chatMessageItems 把 user/assistant 消息转为 message 项;assistant 的
// tool_calls 追加为独立 function_call 项(内容在前,与原语序一致)。
func chatMessageItems(m aiwire.ChatMessage) ([]any, error) {
parts, err := chatContentParts(m.Role, m.Content)
if err != nil {
return nil, err
}
items := make([]any, 0, 1+len(m.ToolCalls))
if len(parts) > 0 {
items = append(items, map[string]any{"role": m.Role, "content": parts})
}
for _, tc := range m.ToolCalls {
items = append(items, map[string]any{"type": "function_call", "call_id": tc.ID,
"name": tc.Function.Name, "arguments": tc.Function.Arguments})
}
return items, nil
}
// chatContentParts 把消息内容转为 Responses 部件(文本按角色定型,图片转 input_image)。
func chatContentParts(role string, c aiwire.Content) ([]map[string]any, error) {
if !c.IsArray {
if c.Text == "" {
return nil, nil
}
return []map[string]any{respTextPart(role, c.Text)}, nil
}
parts := make([]map[string]any, 0, len(c.Parts))
for _, p := range c.Parts {
switch p.Type {
case "", "text":
parts = append(parts, respTextPart(role, p.Text))
case "image_url":
if p.ImageURL == nil || p.ImageURL.URL == "" {
return nil, fmt.Errorf("image_url 块缺少 url")
}
// detail 不透传:兼容面对该字段的支持无法证实(实测期间上游视觉请求不稳,
// 无从归因),它仅是质量提示,忽略合法且消除一个风险轴(README 披露)
parts = append(parts, map[string]any{"type": "input_image", "image_url": p.ImageURL.URL})
default:
return nil, ErrAiUnsupportedBlock
}
}
return parts, nil
}
// chatRespToolChoice 映射 tool_choice:string 形态(auto/none/required)原样;
// {"type":"function","function":{"name":N}} 拍平为 Responses 具名形态;其余忽略。
func chatRespToolChoice(raw json.RawMessage) any {
if len(raw) == 0 {
return nil
}
var s string
if json.Unmarshal(raw, &s) == nil {
switch s {
case "auto", "none", "required":
return s
}
return nil
}
var tc struct {
Type string `json:"type"`
Function struct {
Name string `json:"name"`
} `json:"function"`
}
if json.Unmarshal(raw, &tc) != nil || tc.Type != "function" || tc.Function.Name == "" {
return nil
}
return map[string]string{"type": "function", "name": tc.Function.Name}
}
// chatTextFormat 把 response_format 拍平为 text.format(json_schema 提升嵌套字段)。
func chatTextFormat(rf *aiwire.ResponseFormat) map[string]any {
if rf == nil || rf.Type == "" || rf.Type == "text" {
return nil
}
format := map[string]any{"type": rf.Type}
if rf.Type != "json_schema" || len(rf.JSONSchema) == 0 {
return format
}
var js struct {
Name string `json:"name"`
Schema json.RawMessage `json:"schema"`
Strict *bool `json:"strict"`
}
if json.Unmarshal(rf.JSONSchema, &js) != nil {
return format
}
if js.Name != "" {
format["name"] = js.Name
}
if len(js.Schema) > 0 {
format["schema"] = js.Schema
}
if js.Strict != nil {
format["strict"] = *js.Strict
}
return format
}
// ResponsesToChat 把直通非流式响应转为 Chat Completions 响应(reasoning 项丢弃)。
func ResponsesToChat(payload []byte, id string, created int64) (*aiwire.ChatResponse, error) {
var resp respPayload
if err := json.Unmarshal(payload, &resp); err != nil {
return nil, fmt.Errorf("解析上游响应: %w", err)
}
msg := aiwire.ChatMessage{Role: "assistant"}
var text strings.Builder
for _, item := range resp.Output {
switch item.Type {
case "message":
for _, part := range item.Content {
if part.Type == "output_text" {
text.WriteString(part.Text)
}
}
case "function_call":
msg.ToolCalls = append(msg.ToolCalls, aiwire.ToolCall{ID: item.CallID, Type: "function",
Function: aiwire.FunctionCall{Name: item.Name, Arguments: item.Arguments}})
}
}
msg.Content = aiwire.NewTextContent(text.String())
choice := aiwire.Choice{Index: 0, Message: msg,
FinishReason: chatFinishReason(&resp, len(msg.ToolCalls) > 0)}
return &aiwire.ChatResponse{ID: id, Object: "chat.completion", Created: created,
Model: resp.Model, Choices: []aiwire.Choice{choice}, Usage: usageFromResp(resp.Usage)}, nil
}
// chatFinishReason 映射终态:截断优先 → length;含工具调用 → tool_calls;默认 stop。
func chatFinishReason(resp *respPayload, hasTools bool) string {
if respTruncated(resp) {
return "length"
}
if hasTools {
return "tool_calls"
}
return "stop"
}
// respTruncated 报告直通响应是否因 max_output_tokens 截断。
func respTruncated(resp *respPayload) bool {
return resp != nil && resp.Status == "incomplete" && resp.IncompleteDetails != nil &&
resp.IncompleteDetails.Reason == "max_output_tokens"
}
// ---- Chat Completions 流式桥:直通 SSE 事件 → chat.completion.chunk 序列 ----
// ChatRespBridge 把直通 SSE 事件流桥接为 chunk 序列:首个增量补 role,文本与
// 工具增量按 OpenAI 形态输出;reasoning 系列事件丢弃;[DONE] 由传输层写出。
type ChatRespBridge struct {
id, model string
created int64
includeUsage bool
started bool
toolIndex int
inTool bool
finish string
usage *aiwire.Usage
}
// NewChatRespBridge 构造桥;includeUsage 即 stream_options.include_usage。
func NewChatRespBridge(id, model string, created int64, includeUsage bool) *ChatRespBridge {
return &ChatRespBridge{id: id, model: model, created: created,
includeUsage: includeUsage, toolIndex: -1, finish: "stop"}
}
// Feed 消费一行 SSE data JSON,返回应立即写出的 chunk。
func (b *ChatRespBridge) Feed(data []byte) []aiwire.ChatChunk {
var ev respStreamEvent
if json.Unmarshal(data, &ev) != nil {
return nil
}
switch ev.Type {
case "response.output_text.delta":
b.inTool = false
return []aiwire.ChatChunk{b.chunk(b.deltaWithRole(aiwire.Delta{Content: ev.Delta}), nil)}
case "response.output_item.added":
return b.toolOpen(ev.Item)
case "response.function_call_arguments.delta":
return b.argsChunk(ev.Delta)
case "response.completed", "response.incomplete", "response.failed":
b.finishFrom(ev.Response)
}
return nil
}
// chunk 装配一条单 choice 增量事件。
func (b *ChatRespBridge) chunk(delta aiwire.Delta, finish *string) aiwire.ChatChunk {
return aiwire.ChatChunk{ID: b.id, Object: "chat.completion.chunk", Created: b.created,
Model: b.model, Choices: []aiwire.ChunkChoice{{Index: 0, Delta: delta, FinishReason: finish}}}
}
// deltaWithRole 给首个增量补 role(OpenAI 首块携带 role 语义)。
func (b *ChatRespBridge) deltaWithRole(d aiwire.Delta) aiwire.Delta {
if !b.started {
b.started = true
d.Role = "assistant"
}
return d
}
// toolOpen 在新 function_call 输出项开始时发 id/name 增量并推进聚合 index;
// 非工具输出项只复位增量目标。
func (b *ChatRespBridge) toolOpen(item *respOutputItem) []aiwire.ChatChunk {
if item == nil || item.Type != "function_call" {
b.inTool = false
return nil
}
b.toolIndex++
b.inTool = true
b.finish = "tool_calls"
d := aiwire.Delta{ToolCalls: []aiwire.ToolCallDelta{{Index: b.toolIndex, ID: item.CallID,
Type: "function", Function: aiwire.FunctionCallDelta{Name: item.Name}}}}
return []aiwire.ChatChunk{b.chunk(b.deltaWithRole(d), nil)}
}
// argsChunk 输出当前工具的实参增量(仅在工具输出项进行中)。
func (b *ChatRespBridge) argsChunk(delta string) []aiwire.ChatChunk {
if !b.inTool {
return nil
}
d := aiwire.Delta{ToolCalls: []aiwire.ToolCallDelta{{Index: b.toolIndex,
Function: aiwire.FunctionCallDelta{Arguments: delta}}}}
return []aiwire.ChatChunk{b.chunk(d, nil)}
}
// finishFrom 记录终态 usage 与截断语义(max_output_tokens → length)。
func (b *ChatRespBridge) finishFrom(resp *respPayload) {
if resp == nil {
return
}
b.usage = usageFromResp(resp.Usage)
if respTruncated(resp) {
b.finish = "length"
}
}
// Finish 在上游流结束后收尾:终块带 finish_reason;include_usage 时追加 usage 块
// (choices 为空数组,OpenAI 规范形态;上游未给 usage 时输出零值)。
func (b *ChatRespBridge) Finish() []aiwire.ChatChunk {
chunks := []aiwire.ChatChunk{b.chunk(b.deltaWithRole(aiwire.Delta{}), &b.finish)}
if !b.includeUsage {
return chunks
}
usage := b.usage
if usage == nil {
usage = &aiwire.Usage{}
}
return append(chunks, aiwire.ChatChunk{ID: b.id, Object: "chat.completion.chunk",
Created: b.created, Model: b.model, Choices: []aiwire.ChunkChoice{}, Usage: usage})
}
// Usage 返回聚合到的用量(供调用日志),上游未报告时为 nil。
func (b *ChatRespBridge) Usage() *aiwire.Usage { return b.usage }
// ChatResponseChunks 把完整 Chat 响应展开为 chunk 序列(内容与工具调用 →
// 终块 → 可选 usage 块),供流式上游断流后的非流式降级结果推送。
func ChatResponseChunks(resp *aiwire.ChatResponse, includeUsage bool) []aiwire.ChatChunk {
if len(resp.Choices) == 0 {
return nil
}
choice := resp.Choices[0]
mk := func(delta aiwire.Delta, finish *string) aiwire.ChatChunk {
return aiwire.ChatChunk{ID: resp.ID, Object: "chat.completion.chunk", Created: resp.Created,
Model: resp.Model, Choices: []aiwire.ChunkChoice{{Index: 0, Delta: delta, FinishReason: finish}}}
}
delta := aiwire.Delta{Role: "assistant", Content: choice.Message.Content.Text}
for i, tc := range choice.Message.ToolCalls {
delta.ToolCalls = append(delta.ToolCalls, aiwire.ToolCallDelta{Index: i, ID: tc.ID,
Type: "function", Function: aiwire.FunctionCallDelta{Name: tc.Function.Name, Arguments: tc.Function.Arguments}})
}
finish := choice.FinishReason
chunks := []aiwire.ChatChunk{mk(delta, nil), mk(aiwire.Delta{}, &finish)}
if includeUsage {
usage := resp.Usage
if usage == nil {
usage = &aiwire.Usage{}
}
chunks = append(chunks, aiwire.ChatChunk{ID: resp.ID, Object: "chat.completion.chunk",
Created: resp.Created, Model: resp.Model, Choices: []aiwire.ChunkChoice{}, Usage: usage})
}
return chunks
}
+301
View File
@@ -0,0 +1,301 @@
package service
import (
"encoding/json"
"strings"
"testing"
"oci-portal/internal/aiwire"
)
func mustChatReq(t *testing.T, raw string) aiwire.ChatRequest {
t.Helper()
var req aiwire.ChatRequest
if err := json.Unmarshal([]byte(raw), &req); err != nil {
t.Fatalf("解析请求: %v", err)
}
return req
}
func mustChatBody(t *testing.T, raw string) map[string]any {
t.Helper()
payload, err := ChatToResponsesBody(mustChatReq(t, raw))
if err != nil {
t.Fatalf("ChatToResponsesBody: %v", err)
}
var body map[string]any
if err := json.Unmarshal(payload, &body); err != nil {
t.Fatalf("unmarshal: %v", err)
}
return body
}
// TestChatToResponsesBody 断言消息展开、instructions 聚合、工具与顶层字段装配。
func TestChatToResponsesBody(t *testing.T) {
body := mustChatBody(t, `{
"model": "xai.grok-4.3", "max_tokens": 64, "max_completion_tokens": 128,
"temperature": 0.5, "top_p": 0.9, "parallel_tool_calls": false, "stream": true,
"reasoning_effort": "HIGH",
"messages": [
{"role": "system", "content": "你是助手"},
{"role": "developer", "content": [{"type": "text", "text": "简洁作答"}]},
{"role": "user", "content": "东京天气?"},
{"role": "assistant", "content": "查询中", "tool_calls": [
{"id": "t1", "type": "function", "function": {"name": "get_weather", "arguments": "{\"city\":\"东京\"}"}}
]},
{"role": "tool", "tool_call_id": "t1", "content": "晴 25 度"},
{"role": "user", "content": "继续"}
],
"tools": [{"type": "function", "function": {"name": "get_weather", "description": "查天气", "parameters": {"type": "object"}}}],
"tool_choice": "required"
}`)
if body["model"] != "xai.grok-4.3" || body["store"] != false || body["stream"] != true {
t.Fatalf("顶层字段装配错误: %v", body)
}
if body["max_output_tokens"] != float64(128) {
t.Fatalf("max_output_tokens = %v, want 128(max_completion_tokens 优先)", body["max_output_tokens"])
}
if body["instructions"] != "你是助手\n\n简洁作答" {
t.Fatalf("instructions = %v", body["instructions"])
}
if body["temperature"] != 0.5 || body["top_p"] != 0.9 || body["parallel_tool_calls"] != false {
t.Fatalf("采样字段装配错误: %v", body)
}
if body["tool_choice"] != "required" {
t.Fatalf("tool_choice = %v", body["tool_choice"])
}
if reasoning, _ := body["reasoning"].(map[string]any); reasoning["effort"] != "high" {
t.Fatalf("effort = %v, want high(小写透传)", body["reasoning"])
}
kinds := chatInputKinds(t, body)
want := "message/user,message/assistant,function_call,function_call_output,message/user"
if kinds != want {
t.Fatalf("input 顺序 = %s, want %s", kinds, want)
}
tools, _ := body["tools"].([]any)
tool, _ := tools[0].(map[string]any)
if tool["type"] != "function" || tool["name"] != "get_weather" {
t.Fatalf("工具应拍平为 Responses 形态: %v", tool)
}
}
func chatInputKinds(t *testing.T, body map[string]any) string {
t.Helper()
input, _ := body["input"].([]any)
kinds := make([]string, 0, len(input))
for _, it := range input {
m := it.(map[string]any)
if ty, ok := m["type"].(string); ok {
kinds = append(kinds, ty)
} else {
kinds = append(kinds, "message/"+m["role"].(string))
}
}
return strings.Join(kinds, ",")
}
// TestChatToResponsesBodyContent 断言文本部件按角色定型与图片装配/拒绝。
func TestChatToResponsesBodyContent(t *testing.T) {
body := mustChatBody(t, `{"model":"m","messages":[
{"role":"user","content":[
{"type":"text","text":"看图"},
{"type":"image_url","image_url":{"url":"data:image/png;base64,QUJD"}}]},
{"role":"assistant","content":"这是猫"}]}`)
raw, _ := json.Marshal(body["input"])
if !strings.Contains(string(raw), `"input_image"`) ||
!strings.Contains(string(raw), "data:image/png;base64,QUJD") {
t.Fatalf("图片应转 input_image: %s", raw)
}
if !strings.Contains(string(raw), `"output_text"`) {
t.Fatalf("assistant 历史文本应为 output_text 部件: %s", raw)
}
for name, msg := range map[string]string{
"audio 块": `{"role":"user","content":[{"type":"input_audio","input_audio":{}}]}`,
"图片缺 url": `{"role":"user","content":[{"type":"image_url","image_url":{}}]}`,
} {
raw := `{"model":"m","messages":[` + msg + `]}`
if _, err := ChatToResponsesBody(mustChatReq(t, raw)); err == nil {
t.Errorf("%s 应拒绝", name)
}
}
badTool := mustChatReq(t, `{"model":"m","messages":[{"role":"user","content":"hi"}],
"tools":[{"type":"web_search","function":{}}]}`)
if _, err := ChatToResponsesBody(badTool); err == nil {
t.Error("非 function 工具类型应拒绝")
}
}
// TestChatRespToolChoice 断言 tool_choice 各形态映射。
func TestChatRespToolChoice(t *testing.T) {
cases := []struct {
name, raw string
want any
}{
{"auto", `"auto"`, "auto"},
{"none", `"none"`, "none"},
{"required", `"required"`, "required"},
{"具名 function", `{"type":"function","function":{"name":"f1"}}`,
map[string]string{"type": "function", "name": "f1"}},
{"未知 string", `"whatever"`, nil},
{"缺 name", `{"type":"function","function":{}}`, nil},
{"空", ``, nil},
}
for _, tc := range cases {
got := chatRespToolChoice(json.RawMessage(tc.raw))
if gotMap, ok := got.(map[string]string); ok {
wantMap, _ := tc.want.(map[string]string)
if wantMap == nil || gotMap["name"] != wantMap["name"] {
t.Errorf("%s: got %v, want %v", tc.name, got, tc.want)
}
continue
}
if got != tc.want {
t.Errorf("%s: got %v, want %v", tc.name, got, tc.want)
}
}
}
// TestChatTextFormat 断言 response_format 拍平映射。
func TestChatTextFormat(t *testing.T) {
if got := chatTextFormat(nil); got != nil {
t.Errorf("nil 应不下发: %v", got)
}
if got := chatTextFormat(&aiwire.ResponseFormat{Type: "text"}); got != nil {
t.Errorf("text 应不下发: %v", got)
}
if got := chatTextFormat(&aiwire.ResponseFormat{Type: "json_object"}); got["type"] != "json_object" {
t.Errorf("json_object: %v", got)
}
got := chatTextFormat(&aiwire.ResponseFormat{Type: "json_schema",
JSONSchema: json.RawMessage(`{"name":"out","schema":{"type":"object"},"strict":true}`)})
if got["type"] != "json_schema" || got["name"] != "out" || got["strict"] != true {
t.Errorf("json_schema 应提升嵌套字段: %v", got)
}
if _, ok := got["schema"]; !ok {
t.Errorf("schema 缺失: %v", got)
}
}
// TestResponsesToChat 断言文本聚合、tool_calls、finish_reason 与 usage 回转。
func TestResponsesToChat(t *testing.T) {
payload := []byte(`{"model":"xai.grok-4.3","status":"completed","output":[
{"type":"reasoning","summary":[]},
{"type":"message","content":[{"type":"output_text","text":"你"},{"type":"output_text","text":"好"}]},
{"type":"function_call","call_id":"c1","name":"get_weather","arguments":"{\"city\":\"东京\"}"}],
"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15,"input_tokens_details":{"cached_tokens":4}}}`)
out, err := ResponsesToChat(payload, "chatcmpl-1", 1700000000)
if err != nil {
t.Fatalf("ResponsesToChat: %v", err)
}
if out.Object != "chat.completion" || out.ID != "chatcmpl-1" || out.Created != 1700000000 {
t.Fatalf("响应骨架错误: %+v", out)
}
msg := out.Choices[0].Message
if msg.Role != "assistant" || msg.Content.JoinText() != "你好" {
t.Fatalf("content 装配错误: %+v", msg)
}
if len(msg.ToolCalls) != 1 || msg.ToolCalls[0].ID != "c1" ||
msg.ToolCalls[0].Function.Arguments != `{"city":"东京"}` {
t.Fatalf("tool_calls 装配错误: %+v", msg.ToolCalls)
}
if out.Choices[0].FinishReason != "tool_calls" {
t.Fatalf("finish_reason = %s, want tool_calls", out.Choices[0].FinishReason)
}
if out.Usage.PromptTokens != 10 || out.Usage.CompletionTokens != 5 || out.Usage.CachedTokens() != 4 {
t.Fatalf("usage = %+v", out.Usage)
}
trunc := []byte(`{"model":"m","status":"incomplete","incomplete_details":{"reason":"max_output_tokens"},
"output":[{"type":"message","content":[{"type":"output_text","text":"半"}]}]}`)
out2, err := ResponsesToChat(trunc, "chatcmpl-2", 1)
if err != nil || out2.Choices[0].FinishReason != "length" {
t.Fatalf("截断 finish_reason = %+v, %v", out2.Choices, err)
}
if out2.Usage != nil {
t.Fatalf("无 usage 时应为 nil: %+v", out2.Usage)
}
}
func chatChunkShapes(chunks []aiwire.ChatChunk) string {
kinds := make([]string, 0, len(chunks))
for _, ch := range chunks {
switch {
case len(ch.Choices) == 0:
kinds = append(kinds, "usage")
case ch.Choices[0].FinishReason != nil:
kinds = append(kinds, "finish:"+*ch.Choices[0].FinishReason)
case len(ch.Choices[0].Delta.ToolCalls) > 0:
kinds = append(kinds, "tool")
case ch.Choices[0].Delta.Role != "":
kinds = append(kinds, "role+text")
default:
kinds = append(kinds, "text")
}
}
return strings.Join(kinds, ",")
}
// TestChatRespBridge 断言流桥:首块 role、文本与工具增量、reasoning 丢弃、usage 块。
func TestChatRespBridge(t *testing.T) {
b := NewChatRespBridge("chatcmpl-1", "m1", 1700000000, true)
var chunks []aiwire.ChatChunk
feed := func(lines ...string) {
for _, l := range lines {
chunks = append(chunks, b.Feed([]byte(l))...)
}
}
feed(`{"type":"response.created","response":{"model":"m1"}}`,
`{"type":"response.reasoning_text.delta","delta":"思考中"}`,
`{"type":"response.output_text.delta","delta":"你"}`,
`{"type":"response.output_text.delta","delta":"好"}`,
`{"type":"response.output_item.added","item":{"type":"function_call","call_id":"c1","name":"f"}}`,
`{"type":"response.function_call_arguments.delta","delta":"{\"a\":"}`,
`{"type":"response.function_call_arguments.delta","delta":"1}"}`,
`{"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":8,"output_tokens":4,"total_tokens":12}}}`)
chunks = append(chunks, b.Finish()...)
got := chatChunkShapes(chunks)
want := "role+text,text,tool,tool,tool,finish:tool_calls,usage"
if got != want {
t.Fatalf("chunk 序列:\n got %s\nwant %s", got, want)
}
if chunks[0].Choices[0].Delta.Role != "assistant" || chunks[0].Choices[0].Delta.Content != "你" {
t.Fatalf("首块应含 role 与文本: %+v", chunks[0].Choices[0].Delta)
}
tool := chunks[2].Choices[0].Delta.ToolCalls[0]
if tool.Index != 0 || tool.ID != "c1" || tool.Function.Name != "f" {
t.Fatalf("工具首块 = %+v", tool)
}
args := chunks[3].Choices[0].Delta.ToolCalls[0].Function.Arguments +
chunks[4].Choices[0].Delta.ToolCalls[0].Function.Arguments
if args != `{"a":1}` {
t.Fatalf("实参增量聚合 = %s", args)
}
last := chunks[len(chunks)-1]
if last.Usage == nil || last.Usage.PromptTokens != 8 || len(last.Choices) != 0 {
t.Fatalf("usage 块 = %+v", last)
}
if b.Usage().TotalTokens != 12 {
t.Fatalf("Usage() = %+v", b.Usage())
}
}
// TestChatRespBridgeVariants 断言 include_usage 关闭、截断与空流的收尾形态。
func TestChatRespBridgeVariants(t *testing.T) {
noUsage := NewChatRespBridge("c", "m", 1, false)
noUsage.Feed([]byte(`{"type":"response.output_text.delta","delta":"x"}`))
noUsage.Feed([]byte(`{"type":"response.incomplete","response":{"status":"incomplete","incomplete_details":{"reason":"max_output_tokens"}}}`))
if got := chatChunkShapes(noUsage.Finish()); got != "finish:length" {
t.Errorf("截断且不带 usage 块: %s", got)
}
empty := NewChatRespBridge("c", "m", 1, false)
fin := empty.Finish()
if got := chatChunkShapes(fin); got != "finish:stop" {
t.Errorf("空流收尾 = %s", got)
}
if fin[0].Choices[0].Delta.Role != "assistant" {
t.Errorf("空流终块应补 role: %+v", fin[0].Choices[0].Delta)
}
}
+113 -16
View File
@@ -57,9 +57,6 @@ type LogEventService struct {
relayPollTick time.Duration // 订阅确认轮询间隔,零值用默认 relayPollTick time.Duration // 订阅确认轮询间隔,零值用默认
relayPollTimeout time.Duration // 订阅确认轮询上限,零值用默认 relayPollTimeout time.Duration // 订阅确认轮询上限,零值用默认
alertMu sync.Mutex // 保护告警规则冷却表
alertSentAt map[uint]time.Time // 规则 ID → 上次告警时刻(阈值型规则冷却)
} }
// NewLogEventService 组装依赖;调用 StartParser / StartCleanup 后台协程后生效。 // NewLogEventService 组装依赖;调用 StartParser / StartCleanup 后台协程后生效。
@@ -373,20 +370,18 @@ func (s *LogEventService) parseOnce(ctx context.Context) {
if len(events) == 0 { if len(events) == 0 {
return return
} }
rules := s.loadEnabledAlertRules(ctx)
for i := range events { for i := range events {
s.processLogEvent(ctx, &events[i], rules) s.processLogEvent(ctx, &events[i])
} }
} }
// processLogEvent 仅在条件更新命中原行后触发通知,删除并发胜出时静默跳过。 // processLogEvent 仅在条件更新命中原行后触发通知,删除并发胜出时静默跳过。
func (s *LogEventService) processLogEvent(ctx context.Context, event *model.LogEvent, rules []model.AlertRule) { func (s *LogEventService) processLogEvent(ctx context.Context, event *model.LogEvent) {
parsed := parseLogEvent([]byte(event.Payload)) parsed := parseLogEvent([]byte(event.Payload))
if !s.updateParsedEvent(ctx, event, parsed) { if !s.updateParsedEvent(ctx, event, parsed) {
return return
} }
s.notifyCritical(ctx, event, parsed) s.notifyCritical(ctx, event, parsed)
s.matchAlertRules(ctx, rules, event, parsed)
} }
// updateParsedEvent 用条件 UPDATE 禁止 Save 在删除后隐式重建事件。 // updateParsedEvent 用条件 UPDATE 禁止 Save 在删除后隐式重建事件。
@@ -418,21 +413,44 @@ type onsEnvelope struct {
Source string `json:"source"` Source string `json:"source"`
EventTime string `json:"eventTime"` EventTime string `json:"eventTime"`
ResourceName string `json:"resourceName"` ResourceName string `json:"resourceName"`
Message string `json:"message"`
Data json.RawMessage `json:"data"` Data json.RawMessage `json:"data"`
Identity *onsIdentity `json:"identity"` Identity *onsIdentity `json:"identity"`
AdditionalDetails *onsAddDetails `json:"additionalDetails"`
StateChange *onsStateChange `json:"stateChange"`
} }
// onsIdentity 是 Audit 事件 data.identity 中与展示相关的字段。 // onsIdentity 是 Audit 事件 data.identity 中与展示相关的字段。
type onsIdentity struct { type onsIdentity struct {
IPAddress string `json:"ipAddress"` IPAddress string `json:"ipAddress"`
PrincipalName string `json:"principalName"`
} }
// parsedEvent 是从消息原文提取的展示字段集;ResourceName 仅供 P2 推送文案,不落库。 // onsAddDetails 是 IDCS 登录类事件 data.additionalDetails 的补充字段;
// AuditEventMapValue 为嵌套的 JSON 字符串(含 eventId 成败与失败原因)。
type onsAddDetails struct {
ActorName string `json:"actorName"`
ClientIP string `json:"clientIp"`
AuditEventMapValue string `json:"auditEventMapValue"`
}
// onsStateChange 承载 Audit v2 的资源变更快照;description 供策略类事件推送文案。
type onsStateChange struct {
Current struct {
Description string `json:"description"`
} `json:"current"`
}
// parsedEvent 是从消息原文提取的展示字段集;ResourceName/Actor/Outcome/Detail
// 仅供 P2 推送文案,不落库。
type parsedEvent struct { type parsedEvent struct {
EventType string EventType string
Source string Source string
SourceIP string SourceIP string
ResourceName string ResourceName string
Actor string // 操作者(identity.principalName,登录事件回退 actorName)
Outcome string // 成功 / 失败 / 空(判读不出)
Detail string // 补充说明(策略描述、登录失败原因等)
EventTime *time.Time EventTime *time.Time
} }
@@ -455,7 +473,7 @@ func parseLogEvent(payload []byte) parsedEvent {
return envelopeFields(env) return envelopeFields(env)
} }
// mergeEnvelope 外层缺失字段时以 data 内层补齐(Audit 的 identity 在内层)。 // mergeEnvelope 外层缺失字段时以 data 内层补齐(Audit 的 identity 在内层)。
func mergeEnvelope(outer, inner onsEnvelope) onsEnvelope { func mergeEnvelope(outer, inner onsEnvelope) onsEnvelope {
if outer.EventType == "" { if outer.EventType == "" {
outer.EventType = inner.EventType outer.EventType = inner.EventType
@@ -472,13 +490,22 @@ func mergeEnvelope(outer, inner onsEnvelope) onsEnvelope {
if outer.ResourceName == "" { if outer.ResourceName == "" {
outer.ResourceName = inner.ResourceName outer.ResourceName = inner.ResourceName
} }
if outer.Message == "" {
outer.Message = inner.Message
}
if outer.Identity == nil { if outer.Identity == nil {
outer.Identity = inner.Identity outer.Identity = inner.Identity
} }
if outer.AdditionalDetails == nil {
outer.AdditionalDetails = inner.AdditionalDetails
}
if outer.StateChange == nil {
outer.StateChange = inner.StateChange
}
return outer return outer
} }
// envelopeFields 收敛字段别名并解析事件时间。 // envelopeFields 收敛字段别名并解析事件时间与推送用补充字段
func envelopeFields(env onsEnvelope) parsedEvent { func envelopeFields(env onsEnvelope) parsedEvent {
eventType := env.EventType eventType := env.EventType
if eventType == "" { if eventType == "" {
@@ -490,19 +517,81 @@ func envelopeFields(env onsEnvelope) parsedEvent {
eventTime = &ts eventTime = &ts
} }
} }
ip := "" outcome, detail := envOutcome(env)
if env.Identity != nil {
ip = env.Identity.IPAddress
}
return parsedEvent{ return parsedEvent{
EventType: clip(eventType, 128), EventType: clip(eventType, 128),
Source: clip(env.Source, 64), Source: clip(env.Source, 64),
SourceIP: clip(ip, 64), SourceIP: clip(envIP(env), 64),
ResourceName: clip(env.ResourceName, 128), ResourceName: clip(env.ResourceName, 128),
Actor: clipRunes(envActor(env), 64),
Outcome: outcome,
Detail: clipRunes(detail, 200),
EventTime: eventTime, EventTime: eventTime,
} }
} }
// envIP 取发起方 IP:Audit 的 identity.ipAddress,登录事件回退 additionalDetails.clientIp。
func envIP(env onsEnvelope) string {
if env.Identity != nil && env.Identity.IPAddress != "" {
return env.Identity.IPAddress
}
if env.AdditionalDetails != nil {
return env.AdditionalDetails.ClientIP
}
return ""
}
// envActor 取操作者:Audit 的 identity.principalName,登录事件回退 additionalDetails.actorName。
func envActor(env onsEnvelope) string {
if env.Identity != nil && env.Identity.PrincipalName != "" {
return env.Identity.PrincipalName
}
if env.AdditionalDetails != nil {
return env.AdditionalDetails.ActorName
}
return ""
}
// envOutcome 判读事件成败与补充说明:登录事件解析 auditEventMapValue;其余
// Audit 事件看 message 后缀,补充说明取 stateChange.current.description(策略描述)。
func envOutcome(env onsEnvelope) (outcome, detail string) {
if env.StateChange != nil {
detail = env.StateChange.Current.Description
}
if env.AdditionalDetails != nil && env.AdditionalDetails.AuditEventMapValue != "" {
return ssoOutcome(env.AdditionalDetails.AuditEventMapValue, detail)
}
switch {
case strings.HasSuffix(env.Message, " succeeded"):
outcome = "成功"
case strings.HasSuffix(env.Message, " failed"):
outcome = "失败"
}
return outcome, detail
}
// ssoOutcome 解析 IDCS 审计负载(JSON 字符串):eventId 含 success/failure 定成败,
// 失败时以其 message 作为原因说明。
func ssoOutcome(raw, detail string) (string, string) {
var ev struct {
EventID string `json:"eventId"`
Message string `json:"message"`
}
if json.Unmarshal([]byte(raw), &ev) != nil {
return "", detail
}
switch {
case strings.Contains(ev.EventID, "success"):
return "成功", detail
case strings.Contains(ev.EventID, "failure"):
if ev.Message != "" {
detail = ev.Message
}
return "失败", detail
}
return "", detail
}
// clip 按模型列宽截断解析出的字段。 // clip 按模型列宽截断解析出的字段。
func clip(s string, max int) string { func clip(s string, max int) string {
if len(s) > max { if len(s) > max {
@@ -511,6 +600,15 @@ func clip(s string, max int) string {
return s return s
} }
// clipRunes 按字符数截断(通知文案用,中文安全)。
func clipRunes(s string, max int) string {
r := []rune(s)
if len(r) > max {
return string(r[:max])
}
return s
}
// StartCleanup 启动周期清理:启动即清一次,之后每 24h 一次,随 ctx 取消退出。 // StartCleanup 启动周期清理:启动即清一次,之后每 24h 一次,随 ctx 取消退出。
func (s *LogEventService) StartCleanup(ctx context.Context) { func (s *LogEventService) StartCleanup(ctx context.Context) {
s.wg.Add(1) s.wg.Add(1)
@@ -535,7 +633,6 @@ func (s *LogEventService) cleanupOnce(ctx context.Context) {
if err := s.cleanup(ctx, logEventRetention, logEventMaxRows); err != nil { if err := s.cleanup(ctx, logEventRetention, logEventMaxRows); err != nil {
log.Printf("log event cleanup: %v", err) log.Printf("log event cleanup: %v", err)
} }
s.cleanupAlertHits(ctx)
} }
// cleanup 先删过期记录,再对超量部分删最旧;阈值参数化便于测试。 // cleanup 先删过期记录,再对超量部分删最旧;阈值参数化便于测试。
+31 -2
View File
@@ -30,8 +30,7 @@ func newLogEventEnv(t *testing.T) (*LogEventService, *gorm.DB, uint) {
t.Fatalf("db handle: %v", err) t.Fatalf("db handle: %v", err)
} }
sqlDB.SetMaxOpenConns(1) sqlDB.SetMaxOpenConns(1)
if err := db.AutoMigrate(&model.Setting{}, &model.OciConfig{}, &model.LogEvent{}, if err := db.AutoMigrate(&model.Setting{}, &model.OciConfig{}, &model.LogEvent{}); err != nil {
&model.AlertRule{}, &model.AlertRuleHit{}); err != nil {
t.Fatalf("auto migrate: %v", err) t.Fatalf("auto migrate: %v", err)
} }
cfg := model.OciConfig{Alias: "测试租户"} cfg := model.OciConfig{Alias: "测试租户"}
@@ -246,6 +245,9 @@ func TestParseLogEvent(t *testing.T) {
wantSource string wantSource string
wantIP string wantIP string
wantTime bool wantTime bool
wantActor string
wantOutcome string
wantDetail string
}{ }{
{ {
name: "CloudEvents 单事件", name: "CloudEvents 单事件",
@@ -283,6 +285,29 @@ func TestParseLogEvent(t *testing.T) {
payload: `{"eventType":"x","data":{"identity":{}}}`, payload: `{"eventType":"x","data":{"identity":{}}}`,
wantType: "x", wantType: "x",
}, },
{
name: "Audit v2 提取操作者成败与策略描述",
payload: `{"eventType":"com.oraclecloud.identityControlPlane.CreatePolicy","data":{
"identity":{"principalName":"Alfonso Garcia","ipAddress":"129.159.43.9"},
"message":"ociportal-logs-sch CreatePolicy succeeded",
"stateChange":{"current":{"description":"允许 Connector 发布到 ONS"}}}}`,
wantType: "com.oraclecloud.identityControlPlane.CreatePolicy", wantIP: "129.159.43.9",
wantActor: "Alfonso Garcia", wantOutcome: "成功", wantDetail: "允许 Connector 发布到 ONS",
},
{
name: "IDCS 登录失败提取用户 IP 与原因",
payload: `{"eventType":"com.oraclecloud.IdentitySignOn.InteractiveLogin","data":{
"additionalDetails":{"actorName":"oci","clientIp":"155.117.82.111",
"auditEventMapValue":"{\"eventId\":\"sso.authentication.failure\",\"message\":\"Authentication failure : incorrect password.\"}"}}}`,
wantType: "com.oraclecloud.IdentitySignOn.InteractiveLogin", wantIP: "155.117.82.111",
wantActor: "oci", wantOutcome: "失败", wantDetail: "Authentication failure : incorrect password.",
},
{
name: "message failed 后缀判失败",
payload: `{"eventType":"x","data":{"message":"vm TerminateInstance failed"}}`,
wantType: "x",
wantOutcome: "失败",
},
} }
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
@@ -291,6 +316,10 @@ func TestParseLogEvent(t *testing.T) {
t.Errorf("parse = (%q,%q,%q), want (%q,%q,%q)", t.Errorf("parse = (%q,%q,%q), want (%q,%q,%q)",
got.EventType, got.Source, got.SourceIP, tt.wantType, tt.wantSource, tt.wantIP) got.EventType, got.Source, got.SourceIP, tt.wantType, tt.wantSource, tt.wantIP)
} }
if got.Actor != tt.wantActor || got.Outcome != tt.wantOutcome || got.Detail != tt.wantDetail {
t.Errorf("推送字段 = (%q,%q,%q), want (%q,%q,%q)",
got.Actor, got.Outcome, got.Detail, tt.wantActor, tt.wantOutcome, tt.wantDetail)
}
if (got.EventTime != nil) != tt.wantTime { if (got.EventTime != nil) != tt.wantTime {
t.Errorf("eventTime present = %v, want %v", got.EventTime != nil, tt.wantTime) t.Errorf("eventTime present = %v, want %v", got.EventTime != nil, tt.wantTime)
} }
+23 -1
View File
@@ -322,12 +322,34 @@ func relayEventShortName(eventType string) string {
} }
// criticalEventVars 判定关键事件并生成模板变量;非关键事件 ok 为 false。 // criticalEventVars 判定关键事件并生成模板变量;非关键事件 ok 为 false。
// actor/ip/outcome/resource 缺失时兜底 —,detail 包装为独立行(空则不占行)。
func criticalEventVars(alias string, p parsedEvent) (map[string]string, bool) { func criticalEventVars(alias string, p parsedEvent) (map[string]string, bool) {
name := relayEventShortName(p.EventType) name := relayEventShortName(p.EventType)
if name == "" || !relayCriticalSet[name] { if name == "" || !relayCriticalSet[name] {
return nil, false return nil, false
} }
return map[string]string{"tenant": alias, "event": name, "resource": p.ResourceName}, true return map[string]string{
"tenant": alias, "event": name,
"resource": orDash(p.ResourceName), "actor": orDash(p.Actor),
"ip": orDash(p.SourceIP), "outcome": orDash(p.Outcome),
"detail": detailLine(p.Detail),
}, true
}
// orDash 空值兜底为 —,避免模板出现悬空标签。
func orDash(v string) string {
if v == "" {
return "—"
}
return v
}
// detailLine 把补充说明包装为独立行;为空时不产生多余空行。
func detailLine(v string) string {
if v == "" {
return ""
}
return "\n" + v
} }
// notifyCritical 对关键事件推送告警;依赖未注入或对应子类开关关闭时跳过。 // notifyCritical 对关键事件推送告警;依赖未注入或对应子类开关关闭时跳过。
+25 -13
View File
@@ -316,23 +316,33 @@ func TestCriticalEventText(t *testing.T) {
name string name string
event parsedEvent event parsedEvent
wantOK bool wantOK bool
wantText string want map[string]string
}{ }{
{ {
name: "实例终止命中", name: "实例终止含操作者与成败",
event: parsedEvent{EventType: "com.oraclecloud.ComputeApi.TerminateInstance", ResourceName: "vm-1"}, event: parsedEvent{EventType: "com.oraclecloud.ComputeApi.TerminateInstance",
ResourceName: "vm-1", Actor: "demo@example.com", SourceIP: "203.0.113.8", Outcome: "成功"},
wantOK: true, wantOK: true,
wantText: "☁️ 云端事件:免费01 TerminateInstance vm-1", want: map[string]string{"event": "TerminateInstance", "resource": "vm-1",
"actor": "demo@example.com", "ip": "203.0.113.8", "outcome": "成功", "detail": ""},
}, },
{ {
name: "无资源名省略尾段", name: "字段缺失兜底破折号",
event: parsedEvent{EventType: "com.oraclecloud.IdentityControlPlane.CreateApiKey"}, event: parsedEvent{EventType: "com.oraclecloud.IdentityControlPlane.CreateApiKey"},
wantOK: true, wantOK: true,
wantText: "☁️ 云端事件:免费01 CreateApiKey", want: map[string]string{"event": "CreateApiKey", "resource": "—",
"actor": "—", "ip": "—", "outcome": "—", "detail": ""},
},
{
name: "补充说明独立成行",
event: parsedEvent{EventType: "CreatePolicy", Detail: "允许发布到 ONS Topic"},
wantOK: true,
want: map[string]string{"event": "CreatePolicy", "detail": "\n允许发布到 ONS Topic"},
}, },
{name: "List 噪声不推", event: parsedEvent{EventType: "com.oraclecloud.ComputeApi.ListInstances"}, wantOK: false}, {name: "List 噪声不推", event: parsedEvent{EventType: "com.oraclecloud.ComputeApi.ListInstances"}, wantOK: false},
{name: "空类型不推", event: parsedEvent{}, wantOK: false}, {name: "空类型不推", event: parsedEvent{}, wantOK: false},
{name: "短名直接命中", event: parsedEvent{EventType: "LaunchInstance"}, wantOK: true, wantText: "☁️ 云端事件:免费01 LaunchInstance"}, {name: "短名直接命中", event: parsedEvent{EventType: "LaunchInstance"}, wantOK: true,
want: map[string]string{"event": "LaunchInstance"}},
} }
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
@@ -340,13 +350,15 @@ func TestCriticalEventText(t *testing.T) {
if ok != tt.wantOK { if ok != tt.wantOK {
t.Fatalf("ok = %v, want %v", ok, tt.wantOK) t.Fatalf("ok = %v, want %v", ok, tt.wantOK)
} }
if ok { if !ok {
got := "☁️ 云端事件:" + vars["tenant"] + " " + vars["event"] return
if vars["resource"] != "" {
got += " " + vars["resource"]
} }
if got != tt.wantText { if vars["tenant"] != "免费01" {
t.Errorf("vars 渲染 = %q, want %q", got, tt.wantText) t.Errorf("tenant = %q, want 免费01", vars["tenant"])
}
for k, want := range tt.want {
if vars[k] != want {
t.Errorf("%s = %q, want %q", k, vars[k], want)
} }
} }
}) })
+22 -23
View File
@@ -21,7 +21,6 @@ var notifyTplOrder = []string{
"task_fail", "task_recover", "snatch_success", "tenant_dead", "task_stop", "task_fail", "task_recover", "snatch_success", "tenant_dead", "task_stop",
"login_lock", "model_deprecated", "login_lock", "model_deprecated",
"log_event_instance", "log_event_identity", "log_event_policy", "log_event_region", "log_event_login", "log_event_instance", "log_event_identity", "log_event_policy", "log_event_region", "log_event_login",
"audit_alert",
} }
var notifyTplDefs = map[string]notifyTplDef{ var notifyTplDefs = map[string]notifyTplDef{
@@ -73,40 +72,40 @@ var notifyTplDefs = map[string]notifyTplDef{
}, },
"log_event_instance": { "log_event_instance": {
Label: "实例生命周期", Label: "实例生命周期",
Default: "☁️ 云端事件:{{tenant}} {{event}} {{resource}}", Default: "☁️ 云端事件:{{tenant}} {{event}} {{resource}} · {{outcome}}\n操作者 {{actor}} · 来源 {{ip}}{{detail}}",
Vars: []string{"tenant", "event", "resource"}, Vars: []string{"tenant", "event", "resource", "actor", "ip", "outcome", "detail"},
Sample: map[string]string{"tenant": "免费01", "event": "TerminateInstance", "resource": "web-server-1"}, Sample: map[string]string{"tenant": "免费01", "event": "TerminateInstance", "resource": "web-server-1",
"actor": "demo@example.com", "ip": "203.0.113.8", "outcome": "成功", "detail": ""},
}, },
"log_event_identity": { "log_event_identity": {
Label: "用户与凭据", Label: "用户与凭据",
Default: "☁️ 云端事件:{{tenant}} {{event}} {{resource}}", Default: "☁️ 云端事件:{{tenant}} {{event}} {{resource}} · {{outcome}}\n操作者 {{actor}} · 来源 {{ip}}{{detail}}",
Vars: []string{"tenant", "event", "resource"}, Vars: []string{"tenant", "event", "resource", "actor", "ip", "outcome", "detail"},
Sample: map[string]string{"tenant": "免费01", "event": "CreateApiKey", "resource": "user/demo"}, Sample: map[string]string{"tenant": "免费01", "event": "CreateApiKey", "resource": "user/demo",
"actor": "admin@example.com", "ip": "203.0.113.8", "outcome": "成功", "detail": ""},
}, },
"log_event_policy": { "log_event_policy": {
Label: "策略变更", Label: "策略变更",
Default: "☁️ 云端事件:{{tenant}} {{event}} {{resource}}", Default: "☁️ 云端事件:{{tenant}} {{event}} {{resource}} · {{outcome}}\n操作者 {{actor}} · 来源 {{ip}}{{detail}}",
Vars: []string{"tenant", "event", "resource"}, Vars: []string{"tenant", "event", "resource", "actor", "ip", "outcome", "detail"},
Sample: map[string]string{"tenant": "免费01", "event": "UpdatePolicy", "resource": "admin-policy"}, Sample: map[string]string{"tenant": "免费01", "event": "CreatePolicy", "resource": "admin-policy",
"actor": "admin@example.com", "ip": "203.0.113.8", "outcome": "成功",
"detail": "\n允许 Service Connector 发布到 ONS Topic"},
}, },
"log_event_region": { "log_event_region": {
Label: "区域订阅", Label: "区域订阅",
Default: "☁️ 云端事件:{{tenant}} {{event}} {{resource}}", Default: "☁️ 云端事件:{{tenant}} {{event}} {{resource}} · {{outcome}}\n操作者 {{actor}} · 来源 {{ip}}{{detail}}",
Vars: []string{"tenant", "event", "resource"}, Vars: []string{"tenant", "event", "resource", "actor", "ip", "outcome", "detail"},
Sample: map[string]string{"tenant": "免费01", "event": "CreateRegionSubscription", "resource": "ap-osaka-1"}, Sample: map[string]string{"tenant": "免费01", "event": "CreateRegionSubscription", "resource": "ap-osaka-1",
"actor": "admin@example.com", "ip": "203.0.113.8", "outcome": "成功", "detail": ""},
}, },
"log_event_login": { "log_event_login": {
Label: "控制台登录", Label: "控制台登录",
Default: "☁️ 云端事件:{{tenant}} {{event}} {{resource}}", Default: "☁️ 云端事件:{{tenant}} {{event}} · {{outcome}}\n用户 {{actor}} · 登录 IP {{ip}}{{detail}}",
Vars: []string{"tenant", "event", "resource"}, Vars: []string{"tenant", "event", "resource", "actor", "ip", "outcome", "detail"},
Sample: map[string]string{"tenant": "免费01", "event": "InteractiveLogin", "resource": "demo@example.com"}, Sample: map[string]string{"tenant": "免费01", "event": "InteractiveLogin", "resource": "—",
}, "actor": "demo@example.com", "ip": "155.117.82.111", "outcome": "失败",
"audit_alert": { "detail": "\nAuthentication failure : You entered an incorrect user name or password."},
Label: "审计告警",
Default: "🔔 审计告警:{{rule}}\n租户 {{tenant}} · {{event}} {{resource}}\n来源 {{ip}} · 窗口内 {{count}} 次",
Vars: []string{"rule", "tenant", "event", "resource", "ip", "count"},
Sample: map[string]string{"rule": "非白名单终止实例", "tenant": "免费01",
"event": "TerminateInstance", "resource": "web-server-1", "ip": "203.0.113.8", "count": "1"},
}, },
} }
+2 -4
View File
@@ -20,7 +20,6 @@ type OciConfigService struct {
cipher *crypto.Cipher cipher *crypto.Cipher
client oci.Client client oci.Client
cleanupTasks *TaskService cleanupTasks *TaskService
cleanupEvents *LogEventService
// auditRaw 按 configId:eventId 暂存审计原始事件(TTL 10 分钟), // auditRaw 按 configId:eventId 暂存审计原始事件(TTL 10 分钟),
// 列表响应剥离 raw 后详情接口据此秒开;miss 走小窗重查兜底。 // 列表响应剥离 raw 后详情接口据此秒开;miss 走小窗重查兜底。
auditRaw *cache.Cache auditRaw *cache.Cache
@@ -31,10 +30,9 @@ func NewOciConfigService(db *gorm.DB, cipher *crypto.Cipher, client oci.Client)
return &OciConfigService{db: db, cipher: cipher, client: client, auditRaw: cache.New(auditRawMax)} return &OciConfigService{db: db, cipher: cipher, client: client, auditRaw: cache.New(auditRawMax)}
} }
// SetTenantCleanupDeps 注入租户删除提交后的任务与告警内存状态同步依赖。 // SetTenantCleanupDeps 注入租户删除提交后的任务内存状态同步依赖。
func (s *OciConfigService) SetTenantCleanupDeps(tasks *TaskService, events *LogEventService) { func (s *OciConfigService) SetTenantCleanupDeps(tasks *TaskService) {
s.cleanupTasks = tasks s.cleanupTasks = tasks
s.cleanupEvents = events
} }
// ImportInput 是导入一份 API Key 的输入: // ImportInput 是导入一份 API Key 的输入:
+7
View File
@@ -137,6 +137,13 @@ func newTestService(t *testing.T, client oci.Client) *OciConfigService {
if err != nil { if err != nil {
t.Fatalf("open in-memory sqlite: %v", err) t.Fatalf("open in-memory sqlite: %v", err)
} }
// :memory: 库每条连接各自独立,后台 goroutine(模型验证等)触发第二条连接会看到空库;
// 收敛到单连接,与 tenantdelete 测试环境一致
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("database handle: %v", err)
}
sqlDB.SetMaxOpenConns(1)
if err := db.AutoMigrate( if err := db.AutoMigrate(
&model.OciConfig{}, &model.RegionCache{}, &model.CompartmentCache{}, &model.Proxy{}, &model.OciConfig{}, &model.RegionCache{}, &model.CompartmentCache{}, &model.Proxy{},
); err != nil { ); err != nil {
+57
View File
@@ -229,6 +229,63 @@ func (s *ProxyService) Delete(ctx context.Context, id uint) error {
return nil return nil
} }
// SetTenants 设置关联此代理的租户全集:列表内的租户改挂到本代理(含从
// 其他代理改挂),列表外已关联的解除;整体事务生效,返回新的关联计数。
func (s *ProxyService) SetTenants(ctx context.Context, id uint, cfgIDs []uint) (int64, error) {
if err := s.db.WithContext(ctx).First(&model.Proxy{}, id).Error; err != nil {
return 0, fmt.Errorf("find proxy %d: %w", id, err)
}
ids, err := s.validTenantIDs(ctx, cfgIDs)
if err != nil {
return 0, err
}
err = s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
unbind := tx.Model(&model.OciConfig{}).Where("proxy_id = ?", id)
if len(ids) > 0 {
unbind = unbind.Where("id NOT IN ?", ids)
}
if err := unbind.Update("proxy_id", nil).Error; err != nil {
return fmt.Errorf("unbind tenants: %w", err)
}
if len(ids) == 0 {
return nil
}
err := tx.Model(&model.OciConfig{}).Where("id IN ?", ids).Update("proxy_id", id).Error
if err != nil {
return fmt.Errorf("bind tenants: %w", err)
}
return nil
})
if err != nil {
return 0, err
}
return int64(len(ids)), nil
}
// validTenantIDs 去重并校验租户 ID 均存在;含无效 ID 时报 ErrProxyInvalid。
func (s *ProxyService) validTenantIDs(ctx context.Context, cfgIDs []uint) ([]uint, error) {
seen := map[uint]struct{}{}
ids := make([]uint, 0, len(cfgIDs))
for _, v := range cfgIDs {
if _, ok := seen[v]; !ok {
seen[v] = struct{}{}
ids = append(ids, v)
}
}
if len(ids) == 0 {
return ids, nil
}
var n int64
err := s.db.WithContext(ctx).Model(&model.OciConfig{}).Where("id IN ?", ids).Count(&n).Error
if err != nil {
return nil, fmt.Errorf("count tenants: %w", err)
}
if n != int64(len(ids)) {
return nil, fmt.Errorf("存在无效的租户 ID: %w", ErrProxyInvalid)
}
return ids, nil
}
// SpecOf 解密并组装 SDK 用的代理参数;id 为 nil 返回 nil(直连)。 // SpecOf 解密并组装 SDK 用的代理参数;id 为 nil 返回 nil(直连)。
func (s *ProxyService) SpecOf(ctx context.Context, id *uint) (*oci.ProxySpec, error) { func (s *ProxyService) SpecOf(ctx context.Context, id *uint) (*oci.ProxySpec, error) {
if id == nil { if id == nil {
+97
View File
@@ -2,6 +2,7 @@ package service
import ( import (
"context" "context"
"errors"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"strings" "strings"
@@ -227,3 +228,99 @@ func TestProxyGeoProbe(t *testing.T) {
t.Fatalf("probe view = %+v, want 空地区带探测时间", view) t.Fatalf("probe view = %+v, want 空地区带探测时间", view)
} }
} }
func TestProxySetTenants(t *testing.T) {
svc, db := newProxyEnv(t)
ctx := context.Background()
pa, _ := svc.Create(ctx, ProxyInput{Type: "socks5", Host: "10.0.0.1", Port: 1080})
pb, _ := svc.Create(ctx, ProxyInput{Type: "socks5", Host: "10.0.0.2", Port: 1080})
// t1 已挂 A,t2 挂 B,t3 直连
t1 := model.OciConfig{Alias: "t1", ProxyID: &pa.ID}
t2 := model.OciConfig{Alias: "t2", ProxyID: &pb.ID}
t3 := model.OciConfig{Alias: "t3"}
for _, cfg := range []*model.OciConfig{&t1, &t2, &t3} {
if err := db.Create(cfg).Error; err != nil {
t.Fatalf("seed config: %v", err)
}
}
proxyOf := func(id uint) *uint {
var row model.OciConfig
if err := db.First(&row, id).Error; err != nil {
t.Fatalf("load config %d: %v", id, err)
}
return row.ProxyID
}
cases := []struct {
name string
proxyID uint
cfgIDs []uint
wantN int64
wantErr error
}{
{name: "改挂并新增", proxyID: pa.ID, cfgIDs: []uint{t2.ID, t3.ID, t3.ID}, wantN: 2},
{name: "清空解除全部", proxyID: pa.ID, cfgIDs: nil, wantN: 0},
{name: "无效租户拒绝", proxyID: pa.ID, cfgIDs: []uint{9999}, wantErr: ErrProxyInvalid},
{name: "代理不存在", proxyID: 9999, cfgIDs: nil, wantErr: gorm.ErrRecordNotFound},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
n, err := svc.SetTenants(ctx, tc.proxyID, tc.cfgIDs)
if tc.wantErr != nil {
if !errors.Is(err, tc.wantErr) {
t.Fatalf("err = %v, want %v", err, tc.wantErr)
}
return
}
if err != nil || n != tc.wantN {
t.Fatalf("SetTenants = %d, %v; want %d, nil", n, err, tc.wantN)
}
})
}
// 终态:清空后三租户全部直连(t1 在第一步已被解除,t2/t3 在第二步解除)
for _, id := range []uint{t1.ID, t2.ID, t3.ID} {
if got := proxyOf(id); got != nil {
t.Fatalf("config %d proxyID = %v, want nil", id, *got)
}
}
}
func TestProxySetTenantsRebind(t *testing.T) {
svc, db := newProxyEnv(t)
ctx := context.Background()
pa, _ := svc.Create(ctx, ProxyInput{Type: "socks5", Host: "10.0.1.1", Port: 1080})
pb, _ := svc.Create(ctx, ProxyInput{Type: "socks5", Host: "10.0.1.2", Port: 1080})
cfg := model.OciConfig{Alias: "steal", ProxyID: &pa.ID}
if err := db.Create(&cfg).Error; err != nil {
t.Fatalf("seed config: %v", err)
}
if _, err := svc.SetTenants(ctx, pb.ID, []uint{cfg.ID}); err != nil {
t.Fatalf("SetTenants: %v", err)
}
var row model.OciConfig
if err := db.First(&row, cfg.ID).Error; err != nil {
t.Fatalf("load config: %v", err)
}
if row.ProxyID == nil || *row.ProxyID != pb.ID {
t.Fatalf("proxyID = %v, want %d(改挂到 B)", row.ProxyID, pb.ID)
}
// A 侧计数应归零,B 侧为 1
views, err := svc.List(ctx)
if err != nil {
t.Fatalf("List: %v", err)
}
for _, v := range views {
want := int64(0)
if v.ID == pb.ID {
want = 1
}
if v.UsedBy != want {
t.Fatalf("proxy %d UsedBy = %d, want %d", v.ID, v.UsedBy, want)
}
}
}
+59 -4
View File
@@ -3,6 +3,7 @@ package service
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"log" "log"
"strings" "strings"
@@ -46,6 +47,12 @@ type TaskService struct {
entries map[uint]cron.EntryID entries map[uint]cron.EntryID
// runMu 让租户删除等待在途执行结束,并阻止新执行读取删除前的任务快照。 // runMu 让租户删除等待在途执行结束,并阻止新执行读取删除前的任务快照。
runMu sync.RWMutex runMu sync.RWMutex
// runningMu 保护 running:同一任务不允许并发重复执行
// (cron 重叠触发静默跳过,手动触发返回 ErrTaskRunning)
runningMu sync.Mutex
running map[uint]bool
// runWG 追踪手动触发的后台执行,Stop 时等待收尾
runWG sync.WaitGroup
// snatchVarsMu 保护 snatchVars:runSnatch 抢满时写入成功通知变量, // snatchVarsMu 保护 snatchVars:runSnatch 抢满时写入成功通知变量,
// execute 组装执行后快照时取走(一次性) // execute 组装执行后快照时取走(一次性)
@@ -63,6 +70,7 @@ func NewTaskService(db *gorm.DB, configs *OciConfigService, notifier *Notifier,
settings: settings, settings: settings,
cron: cron.New(), cron: cron.New(),
entries: map[uint]cron.EntryID{}, entries: map[uint]cron.EntryID{},
running: map[uint]bool{},
snatchIPWait: defaultSnatchIPWait, snatchIPWait: defaultSnatchIPWait,
snatchVars: map[uint]map[string]string{}, snatchVars: map[uint]map[string]string{},
} }
@@ -86,10 +94,11 @@ func (s *TaskService) Start() error {
return nil return nil
} }
// Stop 停止调度并等待执行中的任务收尾 // Stop 停止调度并等待执行中的任务收尾(含手动触发的后台执行),
// 保证任务产生的异步通知都已进入 Notifier 的等待队列。 // 保证任务产生的异步通知都已进入 Notifier 的等待队列。
func (s *TaskService) Stop() { func (s *TaskService) Stop() {
<-s.cron.Stop().Done() <-s.cron.Stop().Done()
s.runWG.Wait()
} }
// healthCheckPayload 是测活任务参数;ociConfigIds 为空表示全部配置。 // healthCheckPayload 是测活任务参数;ociConfigIds 为空表示全部配置。
@@ -356,7 +365,7 @@ func (s *TaskService) TaskLogs(ctx context.Context, id uint, limit int) ([]model
return logs, nil return logs, nil
} }
// RunTaskNow 立即执行一次任务并返回本次日志。 // RunTaskNow 同步执行一次任务并返回本次日志(内部与测试使用;API 走 TriggerTask)
func (s *TaskService) RunTaskNow(ctx context.Context, id uint) (*model.TaskLog, error) { func (s *TaskService) RunTaskNow(ctx context.Context, id uint) (*model.TaskLog, error) {
if _, err := s.GetTask(ctx, id); err != nil { if _, err := s.GetTask(ctx, id); err != nil {
return nil, err return nil, err
@@ -364,6 +373,43 @@ func (s *TaskService) RunTaskNow(ctx context.Context, id uint) (*model.TaskLog,
return s.execute(id), nil return s.execute(id), nil
} }
// ErrTaskRunning 表示任务已有一次执行在途,拒绝重复触发。
var ErrTaskRunning = errors.New("任务正在执行中,请稍候")
// TriggerTask 异步触发一次任务执行并立即返回;执行结果照常落任务日志与通知,
// 由前端轮询呈现。同一任务在途时返回 ErrTaskRunning。
func (s *TaskService) TriggerTask(ctx context.Context, id uint) error {
if _, err := s.GetTask(ctx, id); err != nil {
return err
}
if !s.beginRun(id) {
return ErrTaskRunning
}
s.runWG.Add(1)
go func() {
defer s.runWG.Done()
s.runHeld(id)
}()
return nil
}
// beginRun 抢占任务的执行权;在途时返回 false。
func (s *TaskService) beginRun(id uint) bool {
s.runningMu.Lock()
defer s.runningMu.Unlock()
if s.running[id] {
return false
}
s.running[id] = true
return true
}
func (s *TaskService) endRun(id uint) {
s.runningMu.Lock()
delete(s.running, id)
s.runningMu.Unlock()
}
// schedule 把任务注册进 cron 调度。 // schedule 把任务注册进 cron 调度。
func (s *TaskService) schedule(task *model.Task) error { func (s *TaskService) schedule(task *model.Task) error {
s.mu.Lock() s.mu.Lock()
@@ -412,9 +458,18 @@ func (s *TaskService) lockTenantCleanup() func() {
return s.runMu.Unlock return s.runMu.Unlock
} }
// execute 执行一次任务:加载 → 分派 → 更新任务状态并写日志, // execute 执行一次任务:同一任务已有执行在途时静默跳过(cron 重叠触发防抖)。
// 前后状态交给通知判定,只在状态变化时推送。
func (s *TaskService) execute(taskID uint) *model.TaskLog { func (s *TaskService) execute(taskID uint) *model.TaskLog {
if !s.beginRun(taskID) {
return nil
}
return s.runHeld(taskID)
}
// runHeld 在已持有执行权的前提下完成一次执行:加载 → 分派 → 更新任务状态并写日志,
// 前后状态交给通知判定,只在状态变化时推送。
func (s *TaskService) runHeld(taskID uint) *model.TaskLog {
defer s.endRun(taskID)
s.runMu.RLock() s.runMu.RLock()
defer s.runMu.RUnlock() defer s.runMu.RUnlock()
return s.executeLocked(taskID) return s.executeLocked(taskID)
+34
View File
@@ -3,6 +3,7 @@ package service
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"strings" "strings"
"sync" "sync"
@@ -801,3 +802,36 @@ func TestNotifyEventsSnatchVarsMerged(t *testing.T) {
} }
} }
} }
func TestTriggerTaskAsyncAndDedup(t *testing.T) {
client := &blockingClient{fakeClient: &fakeClient{}, started: make(chan struct{}, 1), release: make(chan struct{})}
configs, tasks, db := newTenantDeleteEnv(t, client)
target, _ := seedDeleteTenants(t, db)
setTenantPrivateKey(t, configs, target.ID)
task := createHealthTask(t, tasks, target.ID)
ctx := context.Background()
// 触发立即返回,执行在后台开始
if err := tasks.TriggerTask(ctx, task.ID); err != nil {
t.Fatalf("TriggerTask: %v", err)
}
<-client.started
// 在途重复触发被拒;cron 重叠执行静默跳过
if err := tasks.TriggerTask(ctx, task.ID); !errors.Is(err, ErrTaskRunning) {
t.Errorf("在途重复触发应返回 ErrTaskRunning: %v", err)
}
if entry := tasks.execute(task.ID); entry != nil {
t.Error("在途时 cron 重叠执行应静默跳过")
}
close(client.release)
tasks.Stop() // 等待后台执行收尾
assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{task.ID}, 1)
// 执行权已释放,可再次触发;不存在的任务报错
if !tasks.beginRun(task.ID) {
t.Error("执行结束后应可重新获得执行权")
}
tasks.endRun(task.ID)
if err := tasks.TriggerTask(ctx, 9999); err == nil {
t.Error("不存在的任务应报错")
}
}
+20 -98
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"log"
"strings" "strings"
"gorm.io/gorm" "gorm.io/gorm"
@@ -15,7 +16,6 @@ import (
type tenantDeleteResult struct { type tenantDeleteResult struct {
config model.OciConfig config model.OciConfig
deletedTaskIDs []uint deletedTaskIDs []uint
alertRuleIDs []uint
channelsGone bool channelsGone bool
} }
@@ -47,7 +47,7 @@ func (s *OciConfigService) deleteTenantInTx(tx *gorm.DB, id uint, result *tenant
if err := s.deleteTenantTasks(tx, id, result); err != nil { if err := s.deleteTenantTasks(tx, id, result); err != nil {
return err return err
} }
if err := deleteTenantEvents(tx, id, result); err != nil { if err := deleteTenantEvents(tx, id); err != nil {
return err return err
} }
if err := deleteTenantAI(tx, id, result); err != nil { if err := deleteTenantAI(tx, id, result); err != nil {
@@ -104,7 +104,8 @@ func planTenantTask(task model.Task, id uint) (tenantTaskAction, bool, error) {
func planSnatchTask(task model.Task, id uint) (tenantTaskAction, bool, error) { func planSnatchTask(task model.Task, id uint) (tenantTaskAction, bool, error) {
var payload snatchPayload var payload snatchPayload
if err := json.Unmarshal([]byte(task.Payload), &payload); err != nil { if err := json.Unmarshal([]byte(task.Payload), &payload); err != nil {
return tenantTaskAction{}, false, fmt.Errorf("parse task %d payload: %w", task.ID, err) logSkippedTask(task, err)
return tenantTaskAction{}, false, nil
} }
if payload.OciConfigID != id { if payload.OciConfigID != id {
return tenantTaskAction{}, false, nil return tenantTaskAction{}, false, nil
@@ -112,10 +113,17 @@ func planSnatchTask(task model.Task, id uint) (tenantTaskAction, bool, error) {
return tenantTaskAction{task: task, deleteTask: true}, true, nil return tenantTaskAction{task: task, deleteTask: true}, true, nil
} }
// logSkippedTask 记录 payload 无法解析而被跳过的任务:坏数据无法归属租户,
// fail-closed 会永久阻断删除且只能手工修库;改为保留原任务(不删不改)并放行删除。
func logSkippedTask(task model.Task, err error) {
log.Printf("tenant delete: 任务 %d(%s)payload 无法解析,跳过处理: %v", task.ID, task.Type, err)
}
func planMultiTenantTask(task model.Task, id uint) (tenantTaskAction, bool, error) { func planMultiTenantTask(task model.Task, id uint) (tenantTaskAction, bool, error) {
payload, err := decodeConfigIDs(task) payload, err := decodeConfigIDs(task)
if err != nil { if err != nil {
return tenantTaskAction{}, false, err logSkippedTask(task, err)
return tenantTaskAction{}, false, nil
} }
if len(payload.OciConfigIDs) == 0 { if len(payload.OciConfigIDs) == 0 {
return tenantTaskAction{task: task, payload: task.Payload}, true, nil return tenantTaskAction{task: task, payload: task.Payload}, true, nil
@@ -196,17 +204,9 @@ func resetTenantTask(tx *gorm.DB, action tenantTaskAction) error {
return nil return nil
} }
func deleteTenantEvents(tx *gorm.DB, id uint, result *tenantDeleteResult) error { func deleteTenantEvents(tx *gorm.DB, id uint) error {
ruleIDs, eventIDs, affectedRules, err := loadTenantEventRefs(tx, id) if err := lockTenantEventRows(tx, id); err != nil {
if err != nil { return fmt.Errorf("load tenant log events: %w", err)
return err
}
if err := deleteAlertHits(tx, ruleIDs, eventIDs); err != nil {
return err
}
result.alertRuleIDs = mergeIDs(ruleIDs, affectedRules)
if err := deleteWhere(tx, &model.AlertRule{}, "oci_config_id = ?", id); err != nil {
return fmt.Errorf("delete tenant alert rules: %w", err)
} }
if err := deleteWhere(tx, &model.LogEvent{}, "oci_config_id = ?", id); err != nil { if err := deleteWhere(tx, &model.LogEvent{}, "oci_config_id = ?", id); err != nil {
return fmt.Errorf("delete tenant log events: %w", err) return fmt.Errorf("delete tenant log events: %w", err)
@@ -214,88 +214,13 @@ func deleteTenantEvents(tx *gorm.DB, id uint, result *tenantDeleteResult) error
return nil return nil
} }
func loadTenantEventRefs(tx *gorm.DB, id uint) ([]uint, []uint, []uint, error) { // lockTenantEventRows 对租户全部日志事件行加 FOR UPDATE 锁(SQLite 忽略,
ruleIDs, err := lockedTenantRuleIDs(tx, id) // MySQL/PG 阻塞并发写入);事件可达数万,ID 不回传拼接 SQL,删除用条件语句,
if err != nil { // 避免绑定变量上限。
return nil, nil, nil, fmt.Errorf("load tenant alert rules: %w", err) func lockTenantEventRows(tx *gorm.DB, id uint) error {
}
eventIDs, err := lockedTenantEventIDs(tx, id)
if err != nil {
return nil, nil, nil, fmt.Errorf("load tenant log events: %w", err)
}
affected, err := alertHitRuleIDs(tx, eventIDs)
if err != nil {
return nil, nil, nil, err
}
return ruleIDs, eventIDs, affected, nil
}
func lockedTenantRuleIDs(tx *gorm.DB, id uint) ([]uint, error) {
var rows []model.AlertRule
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").
Where("oci_config_id = ?", id).Order("id").Find(&rows).Error
ids := make([]uint, 0, len(rows))
for _, row := range rows {
ids = append(ids, row.ID)
}
return ids, err
}
func lockedTenantEventIDs(tx *gorm.DB, id uint) ([]uint, error) {
var rows []model.LogEvent var rows []model.LogEvent
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id"). return tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").
Where("oci_config_id = ?", id).Order("id").Find(&rows).Error Where("oci_config_id = ?", id).Order("id").Find(&rows).Error
ids := make([]uint, 0, len(rows))
for _, row := range rows {
ids = append(ids, row.ID)
}
return ids, err
}
func alertHitRuleIDs(tx *gorm.DB, eventIDs []uint) ([]uint, error) {
if len(eventIDs) == 0 {
return nil, nil
}
ids := make([]uint, 0)
err := tx.Model(&model.AlertRuleHit{}).Where("log_event_id IN ?", eventIDs).
Distinct().Pluck("rule_id", &ids).Error
if err != nil {
return nil, fmt.Errorf("load affected alert rules: %w", err)
}
return ids, nil
}
func mergeIDs(groups ...[]uint) []uint {
seen := make(map[uint]struct{})
out := make([]uint, 0)
for _, ids := range groups {
for _, id := range ids {
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
out = append(out, id)
}
}
return out
}
func deleteAlertHits(tx *gorm.DB, ruleIDs, eventIDs []uint) error {
query := tx.Model(&model.AlertRuleHit{})
switch {
case len(ruleIDs) > 0 && len(eventIDs) > 0:
query = query.Where("rule_id IN ? OR log_event_id IN ?", ruleIDs, eventIDs)
case len(ruleIDs) > 0:
query = query.Where("rule_id IN ?", ruleIDs)
case len(eventIDs) > 0:
query = query.Where("log_event_id IN ?", eventIDs)
default:
return nil
}
if err := query.Delete(&model.AlertRuleHit{}).Error; err != nil {
return fmt.Errorf("delete tenant alert hits: %w", err)
}
return nil
} }
func deleteTenantAI(tx *gorm.DB, id uint, result *tenantDeleteResult) error { func deleteTenantAI(tx *gorm.DB, id uint, result *tenantDeleteResult) error {
@@ -448,9 +373,6 @@ func (s *OciConfigService) afterTenantDelete(ctx context.Context, result *tenant
if client, ok := s.client.(tenancyCacheInvalidator); ok { if client, ok := s.client.(tenancyCacheInvalidator); ok {
client.InvalidateTenancy(result.config.TenancyOCID) client.InvalidateTenancy(result.config.TenancyOCID)
} }
if s.cleanupEvents != nil {
s.cleanupEvents.ClearAlertCooldown(result.alertRuleIDs)
}
if s.cleanupTasks != nil { if s.cleanupTasks != nil {
s.cleanupTasks.ApplyTenantCleanup(ctx, result.deletedTaskIDs, result.channelsGone) s.cleanupTasks.ApplyTenantCleanup(ctx, result.deletedTaskIDs, result.channelsGone)
} }
+43 -37
View File
@@ -55,7 +55,7 @@ func newTenantDeleteEnv(t *testing.T, client oci.Client) (*OciConfigService, *Ta
} }
configs := NewOciConfigService(db, cipher, client) configs := NewOciConfigService(db, cipher, client)
tasks := NewTaskService(db, configs, nil, nil) tasks := NewTaskService(db, configs, nil, nil)
configs.SetTenantCleanupDeps(tasks, nil) configs.SetTenantCleanupDeps(tasks)
return configs, tasks, db return configs, tasks, db
} }
@@ -64,7 +64,7 @@ func migrateTenantDeleteModels(t *testing.T, db *gorm.DB) {
err := db.AutoMigrate( err := db.AutoMigrate(
&model.OciConfig{}, &model.Task{}, &model.TaskLog{}, &model.Setting{}, &model.OciConfig{}, &model.Task{}, &model.TaskLog{}, &model.Setting{},
&model.CheckSnapshot{}, &model.CostSnapshot{}, &model.RegionCache{}, &model.CompartmentCache{}, &model.CheckSnapshot{}, &model.CostSnapshot{}, &model.RegionCache{}, &model.CompartmentCache{},
&model.LogEvent{}, &model.AlertRule{}, &model.AlertRuleHit{}, &model.LogEvent{},
&model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}, &model.AiContentLog{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}, &model.AiContentLog{},
&model.Proxy{}, &model.AiKey{}, &model.SystemLog{}, &model.Proxy{}, &model.AiKey{}, &model.SystemLog{},
) )
@@ -108,7 +108,10 @@ func tenantTaskCases() []tenantTaskCase {
{name: "测活多租户", task: taskOf(model.TaskTypeHealthCheck, `{"ociConfigIds":[1,2]}`), wantOK: true, wantPayload: `{"ociConfigIds":[2]}`}, {name: "测活多租户", task: taskOf(model.TaskTypeHealthCheck, `{"ociConfigIds":[1,2]}`), wantOK: true, wantPayload: `{"ociConfigIds":[2]}`},
{name: "成本去重命中", task: taskOf(model.TaskTypeCost, `{"ociConfigIds":[1,1,2]}`), wantOK: true, wantPayload: `{"ociConfigIds":[2]}`}, {name: "成本去重命中", task: taskOf(model.TaskTypeCost, `{"ociConfigIds":[1,1,2]}`), wantOK: true, wantPayload: `{"ociConfigIds":[2]}`},
{name: "成本未命中", task: taskOf(model.TaskTypeCost, `{"ociConfigIds":[2]}`)}, {name: "成本未命中", task: taskOf(model.TaskTypeCost, `{"ociConfigIds":[2]}`)},
{name: "非法 JSON", task: taskOf(model.TaskTypeCost, `{`), wantErr: true}, // 坏 payload 记警告跳过(保留原任务),不再 fail-closed 阻断整个租户删除
{name: "非法 JSON 跳过", task: taskOf(model.TaskTypeCost, `{`)},
{name: "抢机坏 payload 跳过", task: taskOf(model.TaskTypeSnatch, `not-json`)},
{name: "抢机空 payload 跳过", task: taskOf(model.TaskTypeSnatch, ``)},
} }
} }
@@ -159,19 +162,9 @@ func seedTenantEvents(t *testing.T, db *gorm.DB, target, other uint) {
t.Helper() t.Helper()
targetEvent := model.LogEvent{OciConfigID: target, MessageID: "target-event"} targetEvent := model.LogEvent{OciConfigID: target, MessageID: "target-event"}
otherEvent := model.LogEvent{OciConfigID: other, MessageID: "other-event"} otherEvent := model.LogEvent{OciConfigID: other, MessageID: "other-event"}
targetRule := model.AlertRule{Name: "target-rule", OciConfigID: target} for _, value := range []any{&targetEvent, &otherEvent} {
globalRule := model.AlertRule{Name: "global-rule", OciConfigID: 0}
for _, value := range []any{&targetEvent, &otherEvent, &targetRule, &globalRule} {
mustCreate(t, db, value) mustCreate(t, db, value)
} }
hits := []model.AlertRuleHit{
{RuleID: targetRule.ID, LogEventID: otherEvent.ID},
{RuleID: globalRule.ID, LogEventID: targetEvent.ID},
{RuleID: globalRule.ID, LogEventID: otherEvent.ID},
}
for i := range hits {
mustCreate(t, db, &hits[i])
}
} }
func seedTenantAI(t *testing.T, db *gorm.DB, target, other uint) { func seedTenantAI(t *testing.T, db *gorm.DB, target, other uint) {
@@ -199,7 +192,7 @@ func assertTenantRowsGone(t *testing.T, db *gorm.DB, id uint) {
rows := []any{ rows := []any{
&model.OciConfig{}, &model.CheckSnapshot{}, &model.CostSnapshot{}, &model.OciConfig{}, &model.CheckSnapshot{}, &model.CostSnapshot{},
&model.RegionCache{}, &model.CompartmentCache{}, &model.LogEvent{}, &model.RegionCache{}, &model.CompartmentCache{}, &model.LogEvent{},
&model.AlertRule{}, &model.AiChannel{}, &model.AiChannel{},
} }
for _, value := range rows { for _, value := range rows {
column := "oci_config_id" column := "oci_config_id"
@@ -209,7 +202,6 @@ func assertTenantRowsGone(t *testing.T, db *gorm.DB, id uint) {
assertCount(t, db, value, column+" = ?", []any{id}, 0) assertCount(t, db, value, column+" = ?", []any{id}, 0)
} }
assertCount(t, db, &model.Setting{}, "key = ?", []any{secretKey(id)}, 0) assertCount(t, db, &model.Setting{}, "key = ?", []any{secretKey(id)}, 0)
assertCount(t, db, &model.AlertRuleHit{}, "", nil, 1)
assertCount(t, db, &model.AiModelCache{}, "", nil, 1) assertCount(t, db, &model.AiModelCache{}, "", nil, 1)
assertCount(t, db, &model.AiCallLog{}, "", nil, 1) assertCount(t, db, &model.AiCallLog{}, "", nil, 1)
assertCount(t, db, &model.AiContentLog{}, "", nil, 1) assertCount(t, db, &model.AiContentLog{}, "", nil, 1)
@@ -245,31 +237,10 @@ func assertRemainingIndirectRows(t *testing.T, db *gorm.DB, otherID uint) {
t.Fatalf("load other AI call: %v", err) t.Fatalf("load other AI call: %v", err)
} }
assertCount(t, db, &model.AiContentLog{}, "call_log_id = ?", []any{call.ID}, 1) assertCount(t, db, &model.AiContentLog{}, "call_log_id = ?", []any{call.ID}, 1)
assertRemainingAlertHit(t, db, otherID)
}
func assertRemainingAlertHit(t *testing.T, db *gorm.DB, otherID uint) {
t.Helper()
var hit model.AlertRuleHit
if err := db.First(&hit).Error; err != nil {
t.Fatalf("load remaining alert hit: %v", err)
}
var rule model.AlertRule
var event model.LogEvent
if err := db.First(&rule, hit.RuleID).Error; err != nil {
t.Fatalf("load remaining rule: %v", err)
}
if err := db.First(&event, hit.LogEventID).Error; err != nil {
t.Fatalf("load remaining event: %v", err)
}
if rule.OciConfigID != 0 || event.OciConfigID != otherID {
t.Errorf("remaining hit = rule cfg %d/event cfg %d, want global/other", rule.OciConfigID, event.OciConfigID)
}
} }
func assertRetainedGlobals(t *testing.T, db *gorm.DB) { func assertRetainedGlobals(t *testing.T, db *gorm.DB) {
t.Helper() t.Helper()
assertCount(t, db, &model.AlertRule{}, "oci_config_id = 0", nil, 1)
assertCount(t, db, &model.Proxy{}, "", nil, 1) assertCount(t, db, &model.Proxy{}, "", nil, 1)
assertCount(t, db, &model.AiKey{}, "", nil, 1) assertCount(t, db, &model.AiKey{}, "", nil, 1)
assertCount(t, db, &model.SystemLog{}, "", nil, 1) assertCount(t, db, &model.SystemLog{}, "", nil, 1)
@@ -607,3 +578,38 @@ func assertCount(t *testing.T, db *gorm.DB, value any, query string, args []any,
t.Errorf("count %T = %d, want %d", value, got, want) t.Errorf("count %T = %d, want %d", value, got, want)
} }
} }
func TestDeleteTenantSkipsCorruptTaskPayload(t *testing.T) {
configs, _, db := newTenantDeleteEnv(t, &fakeClient{})
target, _ := seedDeleteTenants(t, db)
corrupt := model.Task{Name: "corrupt-snatch", Type: model.TaskTypeSnatch, Payload: `{broken`}
mustCreate(t, db, &corrupt)
mustCreate(t, db, &model.TaskLog{TaskID: corrupt.ID, Message: "keep"})
if err := configs.Delete(context.Background(), target.ID); err != nil {
t.Fatalf("坏 payload 不应阻断租户删除: %v", err)
}
var kept model.Task
if err := db.First(&kept, corrupt.ID).Error; err != nil || kept.Payload != `{broken` {
t.Errorf("坏任务应原样保留: %+v, %v", kept, err)
}
assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{corrupt.ID}, 1)
assertCount(t, db, &model.OciConfig{}, "id = ?", []any{target.ID}, 0)
}
func TestDeleteTenantManyEventsNoVarLimit(t *testing.T) {
// 回归:事件数超 SQLite 绑定变量上限(32766)时删除仍成功(旧实现 IN 展开必失败)
configs, _, db := newTenantDeleteEnv(t, &fakeClient{})
target, _ := seedDeleteTenants(t, db)
events := make([]model.LogEvent, 0, 33000)
for i := 0; i < 33000; i++ {
events = append(events, model.LogEvent{OciConfigID: target.ID, MessageID: fmt.Sprintf("bulk-%d", i)})
}
if err := db.CreateInBatches(&events, 500).Error; err != nil {
t.Fatalf("seed events: %v", err)
}
if err := configs.Delete(context.Background(), target.ID); err != nil {
t.Fatalf("数万事件时删除不应受绑定变量上限影响: %v", err)
}
assertCount(t, db, &model.LogEvent{}, "oci_config_id = ?", []any{target.ID}, 0)
}