From 12c6bd63f1959fddbd682b3a78cae7ceafbdcabe Mon Sep 17 00:00:00 2001 From: Rogee Date: Thu, 8 Oct 2026 23:34:17 +0800 Subject: [PATCH] feat: align AI configuration and task trunk concurrency with shared schema --- AGENTS.md | 8 +- contracts/AGENTS.md | 4 +- .../config-read-legacy-hangup-keywords.json | 2 +- .../config-read-task-asr-with-llm.json | 2 +- ...config-read-task-missing-tts-protocol.json | 2 +- contracts/manifest.json | 7 +- contracts/schema | 2 +- contracts/task_trunks_schema_test.go | 62 ++++ contracts/test_verify.py | 18 +- contracts/verify.py | 23 +- deploys/test/saas-mock/main_test.go | 24 ++ deploys/test/saas-mock/publish.go | 2 +- deploys/test/saas-mock/result_wait.go | 4 +- docs/contracts.md | 25 +- .../schema-task-trunk-concurrency-20261008.md | 42 +++ go.mod | 19 +- go.sum | 6 +- internal/ai/asr_sdk_test.go | 15 +- internal/ai/bailian_models_test.go | 52 +-- internal/ai/bailian_realtime.go | 49 +-- internal/ai/bailian_realtime_test.go | 6 +- internal/ai/bailian_tts.go | 186 +--------- internal/ai/bailian_tts_test.go | 155 ++------- internal/ai/binding.go | 329 +++++------------- internal/ai/binding_test.go | 151 +++----- internal/ai/controls_test.go | 15 +- internal/ai/dependency_test.go | 24 ++ internal/ai/limits_test.go | 12 +- internal/ai/llm_params_test.go | 48 +++ internal/ai/media_profile_test.go | 33 +- internal/ai/opening_test.go | 12 +- internal/ai/pipeline.go | 53 ++- internal/ai/pipeline_test.go | 61 +--- internal/ai/protocol_capabilities_test.go | 80 +---- internal/ai/protocol_models_test.go | 79 +---- internal/ai/protocol_runtime_test.go | 264 +++++--------- internal/ai/protocols.go | 40 +-- internal/ai/provider_endpoints_test.go | 55 +++ internal/ai/real_provider_integration_test.go | 5 +- internal/ai/schema_params_test.go | 40 +++ internal/configread/snapshots.go | 49 ++- internal/configread/snapshots_test.go | 14 +- internal/configread/trunk_concurrency_test.go | 75 ++++ internal/dispatcher/ai_test.go | 10 +- internal/dispatcher/execute.go | 15 +- internal/dispatcher/policy.go | 12 +- internal/dispatcher/policy_test.go | 24 +- .../dispatcher/task_trunk_concurrency_test.go | 164 +++++++++ internal/rpc/approved_execution.go | 8 + internal/rpc/approved_execution_test.go | 7 +- internal/rpc/approved_fixture_test.go | 83 +++++ .../rpc/approved_full_ai_integration_test.go | 42 +-- internal/rpc/approved_integration_test.go | 2 +- internal/rpc/approved_runner_test.go | 45 +-- internal/rpc/approved_snapshot_test.go | 10 +- internal/rpc/recording_server_flow_test.go | 4 +- internal/rpc/task_trunk_concurrency_test.go | 56 +++ internal/store/calls.go | 36 +- internal/store/calls_test.go | 2 +- internal/store/snapshot_test.go | 2 +- internal/store/store.go | 2 +- internal/store/store_test.go | 2 +- internal/store/task_trunk_concurrency_test.go | 195 +++++++++++ scripts/check-current-contracts.py | 8 +- scripts/coverage.sh | 5 +- scripts/test_contract_checkout.py | 16 + scripts/test_coverage.py | 41 +++ scripts/update-contracts.sh | 4 +- 68 files changed, 1524 insertions(+), 1425 deletions(-) create mode 100644 contracts/task_trunks_schema_test.go create mode 100644 docs/evidence/schema-task-trunk-concurrency-20261008.md create mode 100644 internal/ai/dependency_test.go create mode 100644 internal/ai/llm_params_test.go create mode 100644 internal/ai/provider_endpoints_test.go create mode 100644 internal/ai/schema_params_test.go create mode 100644 internal/configread/trunk_concurrency_test.go create mode 100644 internal/dispatcher/task_trunk_concurrency_test.go create mode 100644 internal/rpc/approved_fixture_test.go create mode 100644 internal/rpc/task_trunk_concurrency_test.go create mode 100644 internal/store/task_trunk_concurrency_test.go create mode 100644 scripts/test_coverage.py diff --git a/AGENTS.md b/AGENTS.md index 9fec15f..ee305cc 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -18,6 +18,7 @@ - **采用渐进式、分层的方式构建系统**。首先完成能够端到端运行的最小版本,再基于稳定可用的产品逐步增加功能。不要以尚未成熟的复杂性取代已经可用的产品。 - **保持组件的模块化**,并明确划分不同职责与关注点。 - 当成熟且维护良好的库能够降低整体复杂度或提高可靠性时,应优先采用。除非有明确理由,不要重复实现通用功能。 +- `git.ipao.vip/rogee/doubao-speech-go` 是我们自行维护的独立 SDK,不是火山官方包。仓库为 `git@git.ipao.vip:rogee/doubao-speech-go.git`,本地工作区为 `/home/rogee/Workspace/doubao-speech-go`。发现包的问题,应定位 SDK 内真实根因,直接在该仓库修改并补测试,发布后更新本项目的固定依赖版本;不得为了包的缺陷在业务代码里增加兼容层、绕过或静默兜底。本项目只引用固定版本,不内置重复源码或长期使用本地 replace。 - 在自行实现功能或新增依赖之前,应优先评估项目现有依赖的能力。应先查阅相关文档和类型定义,不应未经确认就认定某个库不具备所需能力。 - 架构决策**应着眼于长期演进**。不要采用仅能解决当前问题、且预期需要在后续替换的权宜方案。 - 在设计解决方案之前,**先研究成熟产品如何解决同类问题**。优先采用经过验证的模式和约定,避免从零开始另行设计一套方案。 @@ -79,7 +80,7 @@ ## 当前范围与权威入口 - P01–P08 及 K01–K16 已完成**项目内隔离 Mock** 核验;本轮把分散的人类可读契约、文档与第三方对接合为唯一当前规范,不重审已确认规则。若未来另获启动开发/审查子 Agent 授权,必须按使用者指定的 `gpt-5.6-luna`、`max` 思考和 `fast: true` 逐项核验并显式配置,不静默换模型、降档或关闭 fast。 -- 唯一现行 SaaS↔Dispatcher 业务规范、Schema、拓扑和正常示例由共享仓库统一管理;本项目的来源/hash、验证脚本、错误示例及历史验证来源放在 `contracts/`,不提交到共享仓库或嵌入运行契约。在本项目的共享入口为 [`contracts/schema/`](contracts/schema/README.md);[`docs/thirds/saas-dispatcher.md`](docs/thirds/saas-dispatcher.md) 只保留导航,不维护另一份规范。内部 Agent RPC 仍在 [`proto/agent/agent.proto`](proto/agent/agent.proto)。本地验收与外部缺口见 [`docs/evidence/saas-dispatcher-p08-acceptance.md`](docs/evidence/saas-dispatcher-p08-acceptance.md)。Markdown 不代替机器合同或外部签收,也不另外维护一份平行字段定义。 +- 唯一现行 SaaS↔Dispatcher 业务规范、Schema、拓扑和正常示例由共享仓库统一管理;现行合同只由 Git submodule 提交号固定,不维护重复逐文件或整包 hash;历史来源证据、验证脚本及错误示例放在 `contracts/`,不提交到共享仓库或嵌入运行契约。在本项目的共享入口为 [`contracts/schema/`](contracts/schema/README.md);[`docs/thirds/saas-dispatcher.md`](docs/thirds/saas-dispatcher.md) 只保留导航,不维护另一份规范。内部 Agent RPC 仍在 [`proto/agent/agent.proto`](proto/agent/agent.proto)。本地验收与外部缺口见 [`docs/evidence/saas-dispatcher-p08-acceptance.md`](docs/evidence/saas-dispatcher-p08-acceptance.md)。Markdown 不代替机器合同或外部签收,也不另外维护一份平行字段定义。 - 已由当前合同来源清单固定哈希的历史提案与计划保留**原字节**于 [`docs/archive/sources/`](docs/archive/sources/README.md),使用者原有未提交的两份旧对接文档也按原字节归档;旧上游 v1 在 [`docs/archive/upstream/`](docs/archive/upstream/README.md) 可离线校验,但不嵌入运行合同。旧 F/W 工作包、旧 MQ-only 合同及归档不作为当前运行入口。固定 MQ `v1`、HTTP `/internal/v1/dispatcher/...` 和业务 revision 是现行通信规则,不是自有实现代次;不得为历史路径新建兼容或回退。 - 当前业务范围仍仅**单节点、单 Dispatcher、单 Agent、单 Cell、单租户**。根命令只接受显式 `agent`/`dispatcher`;`mixed` 和裸 `real` 仍拒绝;隔离 `mock`、只读 `sip-only` 和需主机核验/正式获批指令的 `nonprod-real` 为彼此独立的入口。2026-10-05 正式 SaaS Mock 已驱动真实 Agent/Asterisk 发出五通 SIP INVITE,均由线路返回 480,已证实终结并提交结果但均未接通、未见 LLM 应答;数企→15003164745 受当日 3/3 门禁限制尚未在修复版完成试拨,不能声称六组通过。2026-10-04 获批的同一数企/号码两次单次试拨操作均未观测到 SIP 包或最终结果,第二次已有 Dispatcher 派发回执而 Agent 错误根因未知。第二次派发回执经使用者授权后已私密持久化并确认 MQ,最终结果仍缺失;版本 `d2c4a99` 已通过来源/哈希、非呼出时段不拨号抓包预检与版本级只读主机核验并安装;SaaS、Dispatcher、Agent 在测试机已启动,但没有投递新呼叫,Asterisk 仍为零活动通话。SaaS 测试服务已从获批的新私有快照经 HTTPS 提供租户额度 2/revision 2;2026-10-04 旧事件 `call-669f8829-a111-476d-88c7-3de604bf6dc6` 经使用者单独确认,仅在原始 SQLite 备份、只读通道/抓包核实及操作者证据先私密落盘后,人工标记 Dispatcher inbox 为 `finished` 并释放占用(现为 0),Agent 原 `unknown` 幂等日记、派发回执及原始证据均保留;没有 Agent 签发终结事实或最终结果,不能据此声称真实呼叫成功。不得自动重拨、清理未知执行或以编译/本机 Mock 通过宣称真实通话可用。另有严格隔离的 `--mode sip-only`:只允许 Dispatcher 读完整 SIP、持久接纳归属 `sip.config`、经已激活双向 TLS Agent 会话将完整快照应用到原生 Asterisk,并从运行态核对版本;不发现任务、不启动业务呼叫、不开放准入、不处理其他业务控制。测试机已凭使用者批准,将三条历史登记地址以**测试快照显式声明的** UDP/IP 鉴权、无需 REGISTER 配置写入并核对 Asterisk 运行态;供应商尚未确认这些实际线路是否满足上述参数,虽已观察到 SIP 480,但尚未证实线路可接通或完成生产签收。没有真实 SaaS、management、真实通话或生产签收;真实百炼 LLM 与 OSS 已各做一次**不拨号、独立的最小连接/写入诊断**(见 [`docs/evidence/real-ai-oss-one-shot-20261003.md`](docs/evidence/real-ai-oss-one-shot-20261003.md)),这不是获批任务的 ASR/LLM/TTS、录音和上传全链路签收。生产发布包仍为 `production_approval=false`。任何本机 Mock 或 SIP-only 核验均不授权真实呼叫。 - 使用者批准独立的测试专用 `deploys/test/saas-mock/` 仅模拟缺失的 SaaS:从 0600 的静态快照提供五类正式 HTTP 配置,在专用 RabbitMQ vhost 中由模拟 SaaS 预建现行拓扑;不在 Dispatcher 内注入快照;默认不发布外呼消息,只有受限脚本已为同一 `event_id`/线路/原始号码启动抓包后,才可显式单次 MQ 投递;不替代 Agent/Asterisk/AI/OSS。测试用正式入口必须同时具备只读配置、专用 MQ、真实 Agent/Asterisk/AI/录音/OSS、逐通确认和拨号前活跃抓包证据,缺一项不得试拨。此模拟不构成真实 SaaS 签收。 @@ -97,8 +98,9 @@ - 独立 Dispatcher 的 SQLite 是任务、额度、inbox/outbox 的权威数据;Agent 无业务数据库,录音、执行与上传恢复只写受控私有文件。额度包含未知占用,新 boot/租约到期不得自动清除未知执行;明确 Agent 在拨号前拒绝时须持久拒绝并释放,已接纳但证实 ARI originate 从未提交的失败由原 Agent 经正式终结回执释放,RPC 交付未知时仍占用。2026-10-04 修复版 `d2c4a99` 已安装并通过非呼出时段主机诊断,新隔离测试租户额度 2/revision 2 已启用并由 Dispatcher 只读拉取;原快照保留;该条旧执行仅按前述经使用者授权的人工核实记录解除 Dispatcher 占用,绝非新 boot 或超时自动清除未知执行。仅 Agent 用机器可读标记证实本次未发起的拒绝可自动释放;旧执行的同错误码不算证明。不实现双活数据库、自动跨机热备、多 D 共享额度或第二租户公平。本轮不借本机 D1/D2 隔离夹具宣称多 D 运行。不得建立旧表/旧消息/旧 HTTP 执行兼容通道。 - Dispatcher↔Agent 复用 Unary gRPC 和受控 Endpoint;Agent 预绑定 D UUID 与服务端证书指纹,激活/会话代际、peer mTLS/SAN/SNI 和已签发期限须核对,新 boot 不清未知占用。Agent 不自行向 SaaS 取任务/AI/OSS 授权;Dispatcher 只用已经核验的 Agent `GetLoadedSIP` revision 开执行准入。本机 Mock 的加载回报不证明 Asterisk 已实际加载;仅隔离 SIP-only Agent 在配置原子写入、PJSIP reload 和运行态 endpoint/AOR/UDP transport 一致后持久标记 revision,每次加载查询重新核验。只支持明确的 UDP、IP/none 鉴权、无需 REGISTER 的 IPv4/PCMA 线路;未知字段和其他传输/鉴权/注册方式拒绝,不热更静态 transport。每条 SIP trunk 只允许一个原值 `caller_id`,Agent 把该值同时配置为 Asterisk endpoint 的 `from_user` 和 `callerid`,并在 `GetLoadedSIP` 前核对实际运行值;旧任务 `caller_profile_id` 与线路 `caller_profiles` 必须直接拒绝,不保留兼容路径。SIP 配置的唯一编辑/审批面仍是 management。 - 非生产真实 Agent 的 Asterisk `res_hep`/`res_hep_pjsip` 仅镜像至本机 UDP:在 ARI Dial 前读取本次 SIP Call-ID 并绑定执行,按 Call-ID 和原始 INVITE transaction 只采集真实最终响应的状态码、状态行与完整 `raw`;丢失/无法关联以 `sip_capture_error` 显式报告,不从 ARI、号码或挂断原因推测。镜像缺失不阻止已证实终结的额度释放;未经确认的占用仍保留。配置及回滚见 [`deploys/cell/README.md`](deploys/cell/README.md),不得将完整 SIP 报文写入普通日志、提交或聊天;本地测试通过不等于测试机已部署 HEP 或真实接通。 -- 2026-10-08 用户批准本地配置规则调整及新模型调用支持:SaaS SIP `revision` 可省略,完整读取覆盖旧配置;`transport/auth_mode` 判定忽略大小写,但不补其它缺失线路参数。Dispatcher 持久分配内部 SIP 加载代次并核对 Agent/Asterisk,旧通话排空与未知占用保留规则不变。任务列表仅忽略 `schema_version/control_seq`。任务和连接记录统一使用 `provider_id`;任务决定厂商/模型/ASR–LLM–TTS 用途,provider 只提供连接信息,不要求 `role/adapter/enabled`,未知供应商/协议及不可表达的连接参数显式拒绝。2026-10-08 追加用户确认:移除精确型号/音色业务硬编码;同协议模型及可表达参数从任务取得。TTS 必须明确 protocol(HTTP 或 task-based WebSocket),不按 model 名/地址存在与否猜协议、不失败回退。模型实际可用性由供应商调用返回事实决定,不以本地模拟代签。新增本地调用代码及模拟测试不授权部署、拨号或真实 AI 请求。SQLite 当前布局因新增持久 SIP 快照变为 3;旧布局只读拒绝,不自动迁移或改写旧文件。 -- AI 使用任务内不可变授权快照:仅当前已实现且获批准的协议(火山 ASR、OpenAI 兼容 LLM、百炼 task-based ASR/TTS、百炼 HTTP TTS)可表达的参数进入每通话实例;不限制精确型号或音色,不从 SDK 默认值补任务参数。百炼 task ASR 支持任务明确选定 8/16 kHz;TTS 电话输出仍为单声道 16 kHz PCM16,HTTP speed=1、task speed=0.5–2;ASR-only 不启动 LLM/TTS,完整 AI 不借旧语音测试的授权或参数。只有最终用户 ASR 文本的明确字面关键词可触发拒联/挂断;`agent.conversation.hangup_keywords` 仅接受必填非空 `name`/`triggers`/`closingRemark` 的对象数组,多组命中按配置顺序取第一组,经获批 TTS 完整播放该组结束语后主动挂断,不调用 LLM、不开始新一轮对话;合成/播放失败须明确报错并尝试结束,不重播,旧字符串数组直接拒绝;不由 SDK 默认值、环境、CLI、metadata 或宽松 Schema 改写业务参数,不因 SDK 重试产生第二次发起/收费或重播。日志只存脱敏版本/摘要/计数,不存密钥、prompt、完整对话或音频。 +- 2026-10-08 用户批准本地配置规则调整及新模型调用支持:SaaS SIP `revision` 可省略,完整读取覆盖旧配置;`transport/auth_mode` 判定忽略大小写,但不补其它缺失线路参数。Dispatcher 持久分配内部 SIP 加载代次并核对 Agent/Asterisk,旧通话排空与未知占用保留规则不变。任务列表仅忽略 `schema_version/control_seq`。任务用唯一的 `provider_ref` 对应连接的 `provider_code`,`provider_id` 仅为连接记录标识;一份连接可供多用途复用;任务决定厂商/模型/ASR–LLM–TTS 用途,provider 只提供连接信息,不要求 `role/adapter/enabled`,未知供应商/协议及不可表达的连接参数显式拒绝。2026-10-08 追加用户确认:移除精确型号/音色业务硬编码;同协议模型及可表达参数从任务取得。当前 TTS 固定使用 task-based WebSocket,不按 model 名或地址猜协议、不失败回退 HTTP。模型实际可用性由供应商调用返回事实决定,不以本地模拟代签。新增本地调用代码及模拟测试不授权部署、拨号或真实 AI 请求。SQLite 当前布局因新增持久 SIP 快照变为 3;旧布局只读拒绝,不自动迁移或改写旧文件。 +- 2026-10-08 用户确认按共享合同 `b502ad2` 本地适配任务线路并发:`allowed_trunk_ids` 只接受包含 `trunk_id/concurrency` 的对象数组,重复线路拒绝;非负整数 `concurrency` 的 0 表示该任务不使用该线路。选线及 SQLite 原子占用同时遵守租户、任务总量、线路总量、任务在线路上的上限,未知执行不释放;任务修订在选线后变化时拒绝原选择,不用旧快照发出新呼叫。占用仍由现有 inbox 与冻结任务快照确定,不新增业务表、不转换旧快照或迁移旧库。ASR/TTS 使用连接的 `ws_endpoint`,LLM 使用 `api_endpoint`,缺失不互相回退。本轮仅本地代码、测试和说明,不部署、不拨号、不请求真实 AI。 +- AI 使用任务内不可变授权快照:当前使用火山 ASR、OpenAI 兼容 LLM、百炼 task-based ASR/TTS,TTS 固定 WebSocket。任务的 `params` 为必填、可空的平坦对象,值为字符串、数字、布尔值或 null;原样透传,不校验厂商参数值、不补 SDK 默认值,不限制精确型号/音色。任务/session 的 model、voice、消息及事务身份不得被 params 覆盖,冲突明确拒绝。不设或检查 `prompt.max_bytes`,提示词不截断。原生媒体仍为单声道 16 kHz PCM16,不根据 opaque params 改写或重标,SaaS 须提供匹配的服务参数;ASR-only 不启动 LLM/TTS,完整 AI 不借旧语音测试的授权或参数。只有最终用户 ASR 文本的明确字面关键词可触发拒联/挂断;`agent.conversation.hangup_keywords` 仅接受必填非空 `name`/`triggers`/`closingRemark` 的对象数组,多组命中按配置顺序取第一组,经获批 TTS 完整播放该组结束语后主动挂断,不调用 LLM、不开始新一轮对话;合成/播放失败须明确报错并尝试结束,不重播,旧字符串数组直接拒绝;不由 SDK 默认值、环境、CLI、metadata 或宽松 Schema 改写业务参数,不因 SDK 重试产生第二次发起/收费或重播。日志只存脱敏版本/摘要/计数,不存密钥、prompt、完整对话或音频。 - **私有配置位置(本机路径相对本仓库根目录,只读,绝不提交)**:`.local/provider-ai.env` 是 `0600` 的 `KEY=VALUE` 文件;字段名为 `BAILIAN_API_KEY`、`BAILIAN_BASE_URL`、`BAILIAN_WSS_BASE_URL`、`BAILIAN_TTS_VOICE`、`VOLC_ASR_APP_NAME`、`VOLC_ASR_APP_KEY`、`VOLCENGINE_ACCESS_KEY`、`VOLCENGINE_SECRET_KEY`、`VOLCENGINE_REGION`、`VOLCENGINE_DISABLE_SSL`。根目录 `aliyun-oss.env` 也是 `0600`,**不是 shell env 文件**;它以冒号分隔,字段名准确为 `bucket`、`Endpoint`、`Region`,以及 `RAM` 下的 `username`、`accessKeyId`、`accessKeySecret`(大小写须保持原样)。测试机 `rogee` 用户的现行 ARI 文件位于 `~/.config/go-sip-asterisk/{ari.conf,http.conf,ari-secret}`,不是旧 `.local/asterisk-*/ari.conf`;访问测试机前先核对已登记的 SSH 主机指纹,不展示 `ari-secret`。 - **下次安全读取步骤**:先确认工作目录是本仓库,用 `stat` 仅检查本机两份文件是否存在、所有者与权限 `0600`;不满足即停止。按各自格式在受限本机进程中解析所需字段到内存,不执行 `source`、不打印全文/字段值、不写临时明文副本,不把密钥、签名 URL、音频或完整对话带入聊天、日志、提交及长期证据。AI 的历史文件只可作为**获准凭据来源**,模型/voice/速度等仍由当前获批的 task/providers 快照固定,不能用环境变量覆盖。OSS 历史文件也不能直接传给 `DISPATCHER_OSS_CONFIG_FILE`:该运行配置要求私有 JSON、`dispatcher_id` 和 `oss` 字段,并以环境变量**名称引用**密钥;需按现行合同构造并核验授权后才能使用。普通构建和测试不读取这些私有文件;真实服务测试必须显式启用对应 opt-in 并受现行门禁约束。 - **授权边界**:本轮验收目标是三条已登记线路的真实接通及 LLM 正常应答;旧两个号码已分别在三条线路试拨,六通均为 SIP 480,零接通。新增号码的历史逐次授权不等于本次代码变更获准部署或拨号;本次全局审查本地业务硬编码并以 SaaS 配置快照决定任务、线路和额度,**不部署、不拨号**。上述 AI 与 OSS 私有配置仍仅用于另经明确授权的非生产测试,下次任务须重新确认范围和真实服务调用授权,不能沿用本轮或历史一次性授权。不得把历史配置直接当 SaaS 快照、任务授权或真实呼叫准入,不覆盖/清理旧 OSS 对象。 diff --git a/contracts/AGENTS.md b/contracts/AGENTS.md index b34e8d7..e99377c 100644 --- a/contracts/AGENTS.md +++ b/contracts/AGENTS.md @@ -12,7 +12,7 @@ - 调研或修改契约前必须在项目根目录执行 `make contracts-update`;失败即停止,保留现存改动。 - 普通构建只使用固定 submodule 提交,不自动联网追新,不回退旧契约。 -- `manifest.json` 是本项目来源及文件完整性记录,不是 SaaS 消息或配置。对接文件变更经核对后,可显式执行 `python3 contracts/verify.py --write` 更新当前规范与契约包哈希;历史来源哈希不可改写。 -- 共享变更须提交并推送共享仓库,再同步本项目 submodule 指针、验证记录、实现、测试和文档,并通知 SaaS 开发人员。共享仓库单独提交不算完成。 +- 现行合同只由 Git submodule 提交号固定,不维护逐文件或整包 hash。`manifest.json` 仅保存冻结历史归档的来源证据,不是 SaaS 配置;历史证据不可改写。`verify.py` 仍检查离线 Schema 引用和历史来源完整性。 +- 共享变更须提交并推送共享仓库,再同步本项目 submodule 指针、实现、测试和文档,并通知 SaaS 开发人员。共享仓库单独提交不算完成。 - 交付前执行 `make contract-check`、`make check` 和 `make release-check-local`;本地模拟验证不代替真实 SaaS、线路或生产签收。 - 不暂存、提交、覆盖或清理使用者无关修改。具体同步流程见 `../docs/contracts.md`。 diff --git a/contracts/examples/invalid/config-read-legacy-hangup-keywords.json b/contracts/examples/invalid/config-read-legacy-hangup-keywords.json index eebd23c..0d77446 100644 --- a/contracts/examples/invalid/config-read-legacy-hangup-keywords.json +++ b/contracts/examples/invalid/config-read-legacy-hangup-keywords.json @@ -2,7 +2,7 @@ "resource":"task_config","dispatcher_id":"c046b893-8628-4589-ae50-619d049248a6","tenant_id":1001, "task_id":"task-full","task_revision":2,"status":"running","name":"Full AI example", "max_concurrent_calls":2,"ring_timeout_ms":30000,"max_call_duration_ms":120000, - "route_policy_id":"route-mock","allowed_trunk_ids":["trunk-mock"], + "route_policy_id":"route-mock","allowed_trunk_ids":[{"trunk_id":"trunk-mock","concurrency":2}], "schedule":{"time_zone":"Asia/Shanghai","starts_at":"2026-09-21T00:00:00+08:00","ends_at":null,"weekly_windows":{"monday":[{"start":"09:00","end":"20:00"}],"tuesday":[],"wednesday":[],"thursday":[],"friday":[],"saturday":[],"sunday":[]},"excluded_dates":["2026-10-01"]}, "agent":{ "immutable":true,"mode":"full_ai", diff --git a/contracts/examples/invalid/config-read-task-asr-with-llm.json b/contracts/examples/invalid/config-read-task-asr-with-llm.json index 29ff37d..3b97d31 100644 --- a/contracts/examples/invalid/config-read-task-asr-with-llm.json +++ b/contracts/examples/invalid/config-read-task-asr-with-llm.json @@ -1 +1 @@ -{"resource":"task_config","dispatcher_id":"c046b893-8628-4589-ae50-619d049248a6","tenant_id":1001,"task_id":"task-asr","task_revision":1,"status":"running","max_concurrent_calls":2,"ring_timeout_ms":30000,"max_call_duration_ms":120000,"route_policy_id":"route-mock","allowed_trunk_ids":["trunk-mock"],"schedule":{"time_zone":"Asia/Shanghai","starts_at":null,"ends_at":null,"weekly_windows":{"monday":[],"tuesday":[],"wednesday":[],"thursday":[],"friday":[],"saturday":[],"sunday":[]},"excluded_dates":[]},"agent":{"immutable":true,"mode":"asr_only","asr":{"provider_ref":"asr-example","language":"zh-CN","input":{"encoding":"pcm_s16le","sample_rate_hz":16000,"channels":1,"sample_width_bytes":2}},"llm":{"provider_ref":"llm-example","model":"must-not-be-used"}}} +{"resource":"task_config","dispatcher_id":"c046b893-8628-4589-ae50-619d049248a6","tenant_id":1001,"task_id":"task-asr","task_revision":1,"status":"running","max_concurrent_calls":2,"ring_timeout_ms":30000,"max_call_duration_ms":120000,"route_policy_id":"route-mock","allowed_trunk_ids":[{"trunk_id":"trunk-mock","concurrency":2}],"schedule":{"time_zone":"Asia/Shanghai","starts_at":null,"ends_at":null,"weekly_windows":{"monday":[],"tuesday":[],"wednesday":[],"thursday":[],"friday":[],"saturday":[],"sunday":[]},"excluded_dates":[]},"agent":{"immutable":true,"mode":"asr_only","asr":{"provider_ref":"asr-example","language":"zh-CN","input":{"encoding":"pcm_s16le","sample_rate_hz":16000,"channels":1,"sample_width_bytes":2}},"llm":{"provider_ref":"llm-example","model":"must-not-be-used"}}} diff --git a/contracts/examples/invalid/config-read-task-missing-tts-protocol.json b/contracts/examples/invalid/config-read-task-missing-tts-protocol.json index 45ab1a2..cec4523 100644 --- a/contracts/examples/invalid/config-read-task-missing-tts-protocol.json +++ b/contracts/examples/invalid/config-read-task-missing-tts-protocol.json @@ -2,7 +2,7 @@ "resource":"task_config","dispatcher_id":"c046b893-8628-4589-ae50-619d049248a6","tenant_id":1001, "task_id":"task-missing-protocol","task_revision":2,"status":"running","name":"Isolated invalid fixture", "max_concurrent_calls":2,"ring_timeout_ms":30000,"max_call_duration_ms":120000, - "route_policy_id":"route-mock","allowed_trunk_ids":["trunk-mock"], + "route_policy_id":"route-mock","allowed_trunk_ids":[{"trunk_id":"trunk-mock","concurrency":2}], "schedule":{"time_zone":"Asia/Shanghai","starts_at":"2026-09-21T00:00:00+08:00","ends_at":null,"weekly_windows":{"monday":[{"start":"09:00","end":"20:00"}],"tuesday":[],"wednesday":[],"thursday":[],"friday":[],"saturday":[],"sunday":[]},"excluded_dates":[]}, "agent":{ "immutable":true,"mode":"full_ai", diff --git a/contracts/manifest.json b/contracts/manifest.json index 2448278..4c1c0a6 100644 --- a/contracts/manifest.json +++ b/contracts/manifest.json @@ -3,9 +3,6 @@ "status": "isolated-mock-only-external-unverified", "sources": { "archive/sources/v0.5-proposal.md": "612fdaee50aff6aa7fbef16c2d469d99857646c6d2235617d0e67f6098cd7ada", - "archive/sources/plan-saas-dispatcher-v05-v0.1.md": "666f39e56ea9f4b55661efcac82edd6f9729848e2d60e5f24cdf5aa3ac97ee87", - "schema/saas-dispatcher.md": "ca4b4ce87ab2dea2cf1690543cde84ccd6a91d6ab838823bc8b5b0164d77ff2f" - }, - "bundle_sha256": "7c14ee604a61c013ebc62446a2addeafa6e1923f7cbac54dced77b6c4bf34736", - "bundle_algorithm": "sha256 of sorted path relative to schema/ + space + sha256(file) + newline; only schema root-level JSON and schema/examples/**/*.json; verification files and invalid fixtures excluded" + "archive/sources/plan-saas-dispatcher-v05-v0.1.md": "666f39e56ea9f4b55661efcac82edd6f9729848e2d60e5f24cdf5aa3ac97ee87" + } } diff --git a/contracts/schema b/contracts/schema index 9bbf281..b502ad2 160000 --- a/contracts/schema +++ b/contracts/schema @@ -1 +1 @@ -Subproject commit 9bbf281675fbdcd53065fa87b54d65ae14f4b7ed +Subproject commit b502ad2d47629fece4542b117c90509f005eaedb diff --git a/contracts/task_trunks_schema_test.go b/contracts/task_trunks_schema_test.go new file mode 100644 index 0000000..8a0cf28 --- /dev/null +++ b/contracts/task_trunks_schema_test.go @@ -0,0 +1,62 @@ +package contracts + +import ( + "bytes" + "encoding/json" + "testing" + + "github.com/santhosh-tekuri/jsonschema/v6" +) + +func TestTaskTrunkConcurrencySchema(t *testing.T) { + schema, err := CompileCurrent("http-task-detail.schema.json") + if err != nil { + t.Fatal(err) + } + fixture, err := Files.ReadFile("schema/examples/config-read-task-asr.json") + if err != nil { + t.Fatal(err) + } + for _, tc := range []struct { + name string + value string + valid bool + }{ + {"multiple lines", `[{"trunk_id":"12","concurrency":5},{"trunk_id":"15","concurrency":5}]`, true}, + {"zero concurrency", `[{"trunk_id":"12","concurrency":0}]`, true}, + {"old string array", `["12","15"]`, false}, + {"empty array", `[]`, false}, + {"missing trunk ID", `[{"concurrency":5}]`, false}, + {"empty trunk ID", `[{"trunk_id":"","concurrency":5}]`, false}, + {"numeric trunk ID", `[{"trunk_id":12,"concurrency":5}]`, false}, + {"missing concurrency", `[{"trunk_id":"12"}]`, false}, + {"negative concurrency", `[{"trunk_id":"12","concurrency":-1}]`, false}, + {"fractional concurrency", `[{"trunk_id":"12","concurrency":1.5}]`, false}, + {"string concurrency", `[{"trunk_id":"12","concurrency":"5"}]`, false}, + {"null concurrency", `[{"trunk_id":"12","concurrency":null}]`, false}, + {"boolean concurrency", `[{"trunk_id":"12","concurrency":false}]`, false}, + {"unknown property", `[{"trunk_id":"12","concurrency":5,"unknown":true}]`, false}, + {"duplicate entry", `[{"trunk_id":"12","concurrency":5},{"trunk_id":"12","concurrency":5}]`, false}, + {"null entry", `[null]`, false}, + {"null array", `null`, false}, + } { + t.Run(tc.name, func(t *testing.T) { + var doc map[string]json.RawMessage + if err := json.Unmarshal(fixture, &doc); err != nil { + t.Fatal(err) + } + doc["allowed_trunk_ids"] = json.RawMessage(tc.value) + data, err := json.Marshal(doc) + if err != nil { + t.Fatal(err) + } + instance, err := jsonschema.UnmarshalJSON(bytes.NewReader(data)) + if err != nil { + t.Fatal(err) + } + if err := schema.Validate(instance); (err == nil) != tc.valid { + t.Fatalf("valid=%t, want %t: %v", err == nil, tc.valid, err) + } + }) + } +} diff --git a/contracts/test_verify.py b/contracts/test_verify.py index 14dd4b1..0d18b12 100644 --- a/contracts/test_verify.py +++ b/contracts/test_verify.py @@ -23,19 +23,18 @@ class ContractVerificationTest(unittest.TestCase): def test_current_bundle(self): verify_bundle(self.root) - def test_changed_schema_requires_manifest_update(self): + def test_changed_schema_needs_no_duplicate_content_hash(self): path = self.root / 'schema/http-sip.schema.json' path.write_text(path.read_text() + '\n') - with self.assertRaisesRegex(ValueError, 'bundle hash mismatch'): - verify_bundle(self.root) + verify_bundle(self.root) - def test_missing_schema_reference_is_rejected_even_with_updated_hash(self): + def test_missing_schema_reference_is_rejected(self): path = self.root / 'schema/config-read.schema.json' data = json.loads(path.read_text()) data['oneOf'][0]['$ref'] = './missing.schema.json#' path.write_text(json.dumps(data)) with self.assertRaisesRegex(ValueError, 'missing schema reference'): - verify_bundle(self.root, write=True) + verify_bundle(self.root) def test_invalid_pointer_is_rejected(self): path = self.root / 'schema/config-read.schema.json' @@ -43,20 +42,17 @@ class ContractVerificationTest(unittest.TestCase): data['oneOf'][0]['$ref'] = './http-sip.schema.json#/$defs/missing' path.write_text(json.dumps(data)) with self.assertRaisesRegex(ValueError, 'invalid schema pointer'): - verify_bundle(self.root, write=True) + verify_bundle(self.root) def test_archived_source_is_not_rehashed(self): path = next((self.root / 'archive/sources').glob('*.md')) path.write_text(path.read_text() + '\n') with self.assertRaisesRegex(ValueError, 'source hash mismatch'): - verify_bundle(self.root, write=True) + verify_bundle(self.root) - def test_changed_business_document_requires_manifest_update(self): + def test_changed_business_document_needs_no_duplicate_content_hash(self): path = self.root / 'schema/saas-dispatcher.md' path.write_text(path.read_text() + '\n') - with self.assertRaisesRegex(ValueError, 'source hash mismatch'): - verify_bundle(self.root) - verify_bundle(self.root, write=True) verify_bundle(self.root) def test_shared_directory_contains_only_integration_files(self): diff --git a/contracts/verify.py b/contracts/verify.py index 2fd2d43..9312691 100644 --- a/contracts/verify.py +++ b/contracts/verify.py @@ -1,6 +1,5 @@ #!/usr/bin/env python3 -"""Offline reference/provenance checks; --write updates current document/bundle hashes.""" -import argparse +"""Offline schema references and frozen historical provenance checks.""" import hashlib import json from pathlib import Path @@ -10,15 +9,13 @@ def digest(path): return hashlib.sha256(path.read_bytes()).hexdigest() -def verify_bundle(root, write=False): +def verify_bundle(root): shared = root / 'schema' manifest_path = root / 'manifest.json' manifest = json.loads(manifest_path.read_text(encoding='utf-8')) for relative, expected in manifest['sources'].items(): actual = digest(root / relative) - if write and relative == 'schema/saas-dispatcher.md': - manifest['sources'][relative] = actual - elif actual != expected: + if actual != expected: raise ValueError(f'source hash mismatch: {relative}') schemas = {p.name: json.loads(p.read_text(encoding='utf-8')) for p in shared.glob('*.schema.json')} @@ -55,22 +52,12 @@ def verify_bundle(root, write=False): raise ValueError('missing contract examples or schemas') for path in paths: json.loads(path.read_text(encoding='utf-8')) - listing = ''.join(f'{p.relative_to(shared).as_posix()} {digest(p)}\n' for p in paths) - actual = hashlib.sha256(listing.encode('utf-8')).hexdigest() - if write: - manifest['bundle_sha256'] = actual - manifest_path.write_text(json.dumps(manifest, ensure_ascii=False, indent=2) + '\n', encoding='utf-8') - elif actual != manifest['bundle_sha256']: - raise ValueError(f'contract bundle hash mismatch: {actual} != {manifest["bundle_sha256"]}') return len(paths) if __name__ == '__main__': - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument('--write', action='store_true') - args = parser.parse_args() try: - count = verify_bundle(Path(__file__).resolve().parent, args.write) + count = verify_bundle(Path(__file__).resolve().parent) except (ValueError, OSError) as exc: raise SystemExit(str(exc)) from exc - print(f'shared contract: {count} JSON files; source hashes, bundle hash and offline references valid') + print(f'shared contract: {count} JSON files; offline references and frozen historical provenance valid') diff --git a/deploys/test/saas-mock/main_test.go b/deploys/test/saas-mock/main_test.go index 1099200..e7779b9 100644 --- a/deploys/test/saas-mock/main_test.go +++ b/deploys/test/saas-mock/main_test.go @@ -32,6 +32,30 @@ func testDataDir(t *testing.T) string { if err != nil { t.Fatal(err) } + if target == "providers.json" { + var body map[string]any + if err := json.Unmarshal(original, &body); err != nil { + t.Fatal(err) + } + body["providers"] = append(body["providers"].([]any), map[string]any{"provider_id": "test-bailian", "provider_code": "ali_bailian", "name": "Example only", "api_endpoint": "https://example.invalid", "ws_endpoint": "wss://example.invalid", "api_key": "example-only-not-a-real-secret"}) + original, err = json.Marshal(body) + if err != nil { + t.Fatal(err) + } + } + if target == "tasks/task-full.json" { + var body map[string]any + if err := json.Unmarshal(original, &body); err != nil { + t.Fatal(err) + } + a := body["agent"].(map[string]any) + a["llm"].(map[string]any)["provider_ref"] = "ali_bailian" + a["tts"].(map[string]any)["provider_ref"] = "ali_bailian" + original, err = json.Marshal(body) + if err != nil { + t.Fatal(err) + } + } if err := os.WriteFile(filepath.Join(root, target), original, 0600); err != nil { t.Fatal(err) } diff --git a/deploys/test/saas-mock/publish.go b/deploys/test/saas-mock/publish.go index ef7e72c..90595ba 100644 --- a/deploys/test/saas-mock/publish.go +++ b/deploys/test/saas-mock/publish.go @@ -29,7 +29,7 @@ func buildExecute(data dataset, eventID, taskID, callee string, now time.Time) ( } var task configread.Task // The one-shot host capture arm is bound to one exact trunk before the MQ event is sent. - if err := json.Unmarshal(body, &task); err != nil || task.DispatcherID != data.dispatcherID || task.TenantID != data.tenantID || task.TaskID != taskID || task.Status != "running" || len(task.AllowedTrunkIDs) != 1 { + if err := json.Unmarshal(body, &task); err != nil || task.DispatcherID != data.dispatcherID || task.TenantID != data.tenantID || task.TaskID != taskID || task.Status != "running" || len(task.AllowedTrunks) != 1 || task.AllowedTrunks[0].Concurrency == 0 { return "", nil, errors.New("one-shot command requires a running task pinned to exactly one approved trunk") } command, err := json.Marshal(struct { diff --git a/deploys/test/saas-mock/result_wait.go b/deploys/test/saas-mock/result_wait.go index 1416afd..2217917 100644 --- a/deploys/test/saas-mock/result_wait.go +++ b/deploys/test/saas-mock/result_wait.go @@ -209,7 +209,7 @@ func publishAndAwait(ctx context.Context, brokerURL string, data dataset, eventI return oneShotResult{}, err } var task configread.Task - if err := json.Unmarshal(data.tasks[taskID], &task); err != nil || len(task.AllowedTrunkIDs) != 1 { + if err := json.Unmarshal(data.tasks[taskID], &task); err != nil || len(task.AllowedTrunks) != 1 { return oneShotResult{}, errors.New("one-shot result requires a single approved trunk") } consumer, err := openResultConsumer(brokerURL) @@ -222,5 +222,5 @@ func publishAndAwait(ctx context.Context, brokerURL string, data dataset, eventI if err := publishExecuteAt(publishCtx, brokerURL, data, eventID, taskID, callee, at); err != nil { return oneShotResult{}, err } - return consumer.wait(ctx, resultPath, data.dispatcherID, data.tenantID, eventID, taskID, callee, task.AllowedTrunkIDs[0]) + return consumer.wait(ctx, resultPath, data.dispatcherID, data.tenantID, eventID, taskID, callee, task.AllowedTrunks[0].TrunkID) } diff --git a/docs/contracts.md b/docs/contracts.md index 39310d7..10a1101 100644 --- a/docs/contracts.md +++ b/docs/contracts.md @@ -11,7 +11,7 @@ SaaS 开发人员与本项目共同维护业务规范、Schema、MQ 拓扑和正 | 位置 | 用途 | | --- | --- | | `contracts/schema/` | 双方对接所需的规范、Schema、拓扑和正常示例;README 分开列出 Schema 和 Examples 用途,业务规范及 MQ 拓扑在正文单独链接说明。 | -| `contracts/manifest.json` | 本项目记录的当前规范、冻结历史来源和对接文件完整性 hash;不是 SaaS 对接数据。 | +| `contracts/manifest.json` | 本项目记录的冻结历史来源证据;不登记现行合同 hash,不是 SaaS 对接数据。 | | `contracts/verify.py`、`contracts/test_verify.py` | 离线检查 Schema 引用、来源和文件完整性,并验证共享目录及 README 的范围。 | | `contracts/examples/invalid/` | 本项目拒绝错误消息和配置的测试材料,不作为正常对接示例。 | | `contracts/archive/sources/` | 两份历史来源的原字节,只用于验证,不作为当前规范。 | @@ -27,7 +27,7 @@ git submodule update --init --recursive 构建仅将固定提交的 Schema、拓扑和正常示例嵌入制品,不包含本项目 manifest、错误示例或历史来源, 不在运行时读取 Git checkout 或在线加载引用。普通 `make check` 不自动联网追新。 -缺失/未初始化契约、错误来源、未提交修改、指针不一致、hash 错误或未登记的引用均明确失败,没有旧文件兜底。 +缺失/未初始化契约、错误来源、未提交修改、指针不一致或未登记的引用均明确失败,没有旧文件兜底。 发布清单分别记录共享仓库地址、准确提交、拓扑 hash 和本项目验证 manifest 的 hash。 ## 每次调研或修改契约之前 @@ -41,25 +41,24 @@ make contracts-update `contracts-update` 获取共享仓库最新 `main` 并只允许快进;更新前验证当前指针、来源及干净状态, 更新后使用本项目 `contracts/` 中的验证材料检查共享契约。检查通过后仅暂存 `contracts/schema` 新指针。 -脚本不提交其他文件,不自动改写本项目 manifest;若新提交与本项目记录的 hash 不同,检查失败, -须核对差异后按下述流程显式更新验证记录,不得回退到旧文件。 +现行合同仅以 Git submodule 提交号固定,不再登记逐文件或整包 hash。 +脚本不提交其他文件,不改写历史来源证据,不得回退到旧文件。 网络错误、分叉或现有契约修改时停止;不自动 stash/reset,不改写已有历史。 -若指针有意改变但尚未暂存,须先核对该提交再 `git add contracts/schema`,不能用脚本掩盖未知版本。 +已拉取但尚未暂存的快进允许重新核验;脚本仍拒绝分叉历史、错误来源和子模块内未提交文件,核验成功后才暂存指针。 ## 修改与同步交付 1. 确认已完成上述更新;在共享仓库修改对应 Schema、业务规范、拓扑或正常示例。 定义只维护一次,聚合入口只引用独立 Schema;README 将 Schema 与 Examples 分表,业务规范和 MQ 拓扑使用独立链接说明。 -2. 在本项目补充错误示例或测试,并核对共享文件变更后显式更新验证记录: +2. 在本项目补充错误示例或测试,并核对共享文件变更: ```sh - python3 contracts/verify.py --write python3 contracts/verify.py python3 -m unittest discover -s contracts -p 'test_*.py' -v go test ./contracts ``` - `--write` 只更新现行规范/bundle hash,不允许重算冻结历史来源来掩盖改写。 + 验证保留离线 Schema 引用及冻结历史来源检查;现行合同没有需手动刷新的重复 hash。 3. 提交并推送共享仓库,只选择本次修改的对接文件: ```sh @@ -75,4 +74,12 @@ make contracts-update **仅推送契约仓库不算完成。所有契约变更必须同步修改本项目;双方评审和联调需记录同一准确提交。** 已有 P01–P08、A01–A12、K01–K16 项目内证据见 `docs/evidence/saas-dispatcher-p08-acceptance.md`。 -本次目录整理不改变业务规则、不部署、不拨号;Mock、hash 或离线验证不替代 SaaS、供应商或生产签收。 +Mock、hash 或离线验证不替代 SaaS、供应商或生产签收。 + +## 2026-10-08 本地适配 + +共享版本 `b502ad2` 的任务线路并发结构已用于配置读取、冻结快照、Dispatcher 选线与原子占用、Agent 入站校验及 SaaS Mock。旧字符串数组直接拒绝;已有占用与未知执行不因配置调整或重启清除。并发从现有 inbox 计算,不新增表或自动改写旧快照。 + +AI 连接的 HTTP 与 WebSocket 地址按调用用途分别使用,缺失时明确拒绝、不互相回退。字段定义仍只在共享仓库维护。 + +本轮本地验证与外部边界见 [`evidence/schema-task-trunk-concurrency-20261008.md`](evidence/schema-task-trunk-concurrency-20261008.md)。没有部署、真实拨号或真实 AI 请求。 diff --git a/docs/evidence/schema-task-trunk-concurrency-20261008.md b/docs/evidence/schema-task-trunk-concurrency-20261008.md new file mode 100644 index 0000000..d35c130 --- /dev/null +++ b/docs/evidence/schema-task-trunk-concurrency-20261008.md @@ -0,0 +1,42 @@ +# 2026-10-08 Schema 业务适配:本地验证 + +## 来源与范围 + +- 共享合同:`b502ad2d47629fece4542b117c90509f005eaedb`;字段唯一来源为 [`http-task-detail.schema.json`](../../contracts/schema/http-task-detail.schema.json) 和 [`http-ai-providers.schema.json`](../../contracts/schema/http-ai-providers.schema.json)。 +- 使用者确认:适配任务在线路上的并发上限,允许 0;核验 AI 的 HTTP/WebSocket 地址分工;只进行本地开发和测试。 +- 主项目基于 `3d3d190` 的现有工作区,保留此前未提交改动;本记录不是一个已提交版本的部署证据。 + +## 实现结果 + +- 配置读取、冻结任务快照和 Agent 校验使用新线路对象;旧字符串数组、缺失上限、负数和重复线路拒绝。 +- Dispatcher 按任务配置顺序选择仍有容量的线路,0 或本任务在线路上的占用已达上限时不使用该线路;不会在发出失败后换线重拨。 +- SQLite 原子接纳同时核对租户、任务总量、线路总量和任务在线路上的上限。最后一次选线的任务修订必须与持久快照相同,防止选线期间配置调整后仍发出旧选择。 +- 任务在线路上的占用从现有 inbox 中计算;`dispatching`、`dispatched`、`unknown` 均计入。未知执行不因配置调整、重启或经过时间释放,明确终结后才释放。 +- 不新增数据库表,不迁移旧库,不自动转换旧任务快照。已发出通话的冻结快照不被 edit 改写。 +- 已有 AI 地址分工保持不变:ASR/TTS 使用 `ws_endpoint`,LLM 使用 `api_endpoint`;补充不同地址、同连接多用途及缺失地址不回退的测试。 +- 发出前日志只记录事件标识、修订、线路和并发计数,不记录被叫、密钥、完整配置或对话。 +- SaaS Mock 的受限单通入口和结果核验使用新结构;三个本地错误示例只更新线路形状,保留各自原本要验证的错误。 + +## 验证 + +| 检查 | 结果 | +| --- | --- | +| 配置读取及序列化 | 对象数组与 0 保持原值;旧格式、缺失上限、负数、重复 ID 拒绝 | +| 选线 | 独立任务占用、0、容量不足、配置顺序和重复 ID 均有测试 | +| 原子占用 | 20 个竞争请求在任务线路上限 5 时仅 5 个被接纳 | +| 未知与重启 | 原未知通话保留占用;另一任务有独立的任务线路容量;原通话明确结束后才重新开放 | +| edit | 新上限立即约束后续接纳;旧选择修订拒绝;已发出通话维持旧快照 | +| Agent 入站 | 旧格式、无效配置、禁用或未授权线路在记录尝试/发出前拒绝 | +| AI 地址 | 火山/百炼 ASR、百炼 TTS 与 OpenAI 兼容 LLM 地址分工及无回退测试通过 | +| `make check` | 通过格式、合同/Proto 来源、`go vet ./...`、`go test -race ./...`、构建及本地集成检查 | +| RabbitMQ 及完整 Mock 链路 | Broker、Dispatcher runtime、正式命令入口与两个 SaaS Mock 测试实际运行,未跳过;正式命令测试包含双向 TLS Agent 与 HTTPS Mock OSS 上传及结果回报 | +| `make coverage` | 总语句覆盖率 71.6%,高于 65% 要求 | +| `make release-check-local` | 通过;共享合同为上述提交;`source dirty=true`、`production_approval=false` | + +新测试按失败再修复:读取测试先暴露旧 `[]string` 无法解析新对象;占用测试先暴露 0 被接纳、未知占用未受到任务线路限制、20 个竞争请求全部被接纳。实现调整后通过,并进行全量回归。 + +A01–A12 / K01–K16 的权威项目内对照仍见 [`saas-dispatcher-p08-acceptance.md`](saas-dispatcher-p08-acceptance.md);本轮只验证新结构影响的配置、调度、占用及入站校验,并回归现有测试,不重新签收已确认业务规则。 + +## 未授权或未签收 + +没有更新测试机,没有真实拨号、真实 AI/OSS 请求或生产发布。RabbitMQ、本地 HTTPS 和双向 TLS 测试不代表真实 SaaS、线路、Asterisk 媒体或 AI 服务签收。旧快照/旧数据库的处置仍须使用者单独确认,不自动迁移或清理。本记录生成时,共享合同已在前一阶段推送,主项目业务改动仍未提交;后续提交、合并与推送状态以 Git 记录为准。 diff --git a/go.mod b/go.mod index dd66cf1..71b3def 100644 --- a/go.mod +++ b/go.mod @@ -3,10 +3,14 @@ module git.ipao.vip/rogee/go-sip go 1.27.1 require ( + git.ipao.vip/rogee/doubao-speech-go v0.0.0-20261008141055-ae4a740a8538 github.com/CyCoreSystems/ari/v5 v5.3.1 - github.com/GizClaw/doubao-speech-go v0.0.0-20260915022405-e38c14802696 github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.6.0 + github.com/coder/websocket v1.8.15 + github.com/cyberphone/json-canonicalization v0.0.0-20241213102144-19d51d7fe467 + github.com/emiago/sipgo v1.6.0 github.com/google/uuid v1.6.0 + github.com/gorilla/websocket v1.5.3 github.com/openai/openai-go/v3 v3.62.0 github.com/pion/rtp v1.10.5 github.com/pion/rtp/v2 v2.0.0 @@ -15,27 +19,22 @@ require ( github.com/shirou/gopsutil/v4 v4.26.8 github.com/spf13/cobra v1.10.1 github.com/zaf/g711 v1.4.0 - golang.org/x/text v0.41.0 - golang.org/x/time v0.4.0 + golang.org/x/sys v0.47.0 + google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa google.golang.org/grpc v1.83.2 google.golang.org/protobuf v1.36.12 modernc.org/sqlite v1.59.0 ) require ( - github.com/coder/websocket v1.8.15 // indirect - github.com/cyberphone/json-canonicalization v0.0.0-20241213102144-19d51d7fe467 // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/ebitengine/purego v0.10.2 // indirect - github.com/emiago/sipgo v1.6.0 // indirect github.com/go-ole/go-ole v1.2.6 // indirect github.com/go-stack/stack v1.8.0 // indirect github.com/gobwas/httphead v0.1.0 // indirect github.com/gobwas/pool v0.2.1 // indirect github.com/gobwas/ws v1.3.2 // indirect github.com/gogo/protobuf v1.3.2 // indirect - github.com/gorilla/websocket v1.5.3 // indirect - github.com/icholy/digest v1.1.0 // indirect github.com/inconshreveable/log15 v0.0.0-20201112154412-8562bdadbbac // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 // indirect @@ -57,8 +56,8 @@ require ( github.com/yusufpapurcu/wmi v1.2.4 // indirect golang.org/x/net v0.58.0 // indirect golang.org/x/sync v0.22.0 // indirect - golang.org/x/sys v0.47.0 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect + golang.org/x/text v0.41.0 // indirect + golang.org/x/time v0.4.0 // indirect modernc.org/libc v1.75.7 // indirect modernc.org/mathutil v1.7.1 // indirect modernc.org/memory v1.12.1 // indirect diff --git a/go.sum b/go.sum index c7b7728..0476f84 100644 --- a/go.sum +++ b/go.sum @@ -1,7 +1,7 @@ +git.ipao.vip/rogee/doubao-speech-go v0.0.0-20261008141055-ae4a740a8538 h1:AuSI71TFfPLU/IUGa4izRtI58eGl56u7ME0b2k2X3kc= +git.ipao.vip/rogee/doubao-speech-go v0.0.0-20261008141055-ae4a740a8538/go.mod h1:NlDUTh0C+cYDnYUxfq2eZDk7gDMxGxF9roxv+B3P/Og= github.com/CyCoreSystems/ari/v5 v5.3.1 h1:S+NHG1+uMwoAIl0hMBnRUGNZsQKQQQFq7XCRSE2c2mg= github.com/CyCoreSystems/ari/v5 v5.3.1/go.mod h1:8cn9pshP+OAcmAh1y+G2hrGBS1NSF3QmvrARXyhvXxs= -github.com/GizClaw/doubao-speech-go v0.0.0-20260915022405-e38c14802696 h1:FL/2Z4hIT4gaVWz0kCVFSZP1O4RrvpWTZWwLnuVmNO0= -github.com/GizClaw/doubao-speech-go v0.0.0-20260915022405-e38c14802696/go.mod h1:4R3wUAZkYk1BSgzv+QVF+2SZPu3VDy8NvaQdNRYGgBs= github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.6.0 h1:uWzn3io54f9L9mvwsQQSv1KpkkFA06hBxI++RvIyvpI= github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.6.0/go.mod h1:FTzydeQVmR24FI0D6XWUOMKckjXehM/jgMn1xC+DA9M= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= @@ -52,8 +52,6 @@ github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aN github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= -github.com/icholy/digest v1.1.0 h1:HfGg9Irj7i+IX1o1QAmPfIBNu/Q5A5Tu3n/MED9k9H4= -github.com/icholy/digest v1.1.0/go.mod h1:QNrsSGQ5v7v9cReDI0+eyjsXGUoRSUZQHeQ5C4XLa0Y= github.com/inconshreveable/log15 v0.0.0-20201112154412-8562bdadbbac h1:n1DqxAo4oWPMvH1+v+DLYlMCecgumhhgnxAPdqDIFHI= github.com/inconshreveable/log15 v0.0.0-20201112154412-8562bdadbbac/go.mod h1:cOaXtrgN4ScfRrD9Bre7U1thNq5RtJ8ZoP4iXVGRj6o= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= diff --git a/internal/ai/asr_sdk_test.go b/internal/ai/asr_sdk_test.go index f9b874d..28ff169 100644 --- a/internal/ai/asr_sdk_test.go +++ b/internal/ai/asr_sdk_test.go @@ -6,6 +6,7 @@ import ( "encoding/json" "net/http" "net/http/httptest" + "strings" "testing" "time" @@ -40,9 +41,12 @@ func TestASRParametersReachDoubaoSDKWebSocket(t *testing.T) { })) defer server.Close() task, providers := currentFixture(t, tc.mode) - asrProvider := providers["asr-example"] - asrProvider.Endpoint = server.URL - providers[asrProvider.ProviderRef] = asrProvider + task = changeCurrentAgent(t, task, func(a map[string]any) { + a["asr"].(map[string]any)["params"] = map[string]any{"result_type": tc.resultType, "enable_nonstream": tc.nonstream, "unknown_flag": false, "nullable": nil, "opaque": "vendor-only", "large_number": json.Number("9007199254740993")} + }) + asrProvider := providers["volcengine"] + asrProvider.WSEndpoint = "ws" + strings.TrimPrefix(server.URL, "http") + providers[asrProvider.Code] = asrProvider bound, err := Bind(task, providers) if err != nil { t.Fatal(err) @@ -63,13 +67,16 @@ func TestASRParametersReachDoubaoSDKWebSocket(t *testing.T) { if start < 0 { t.Fatal("SDK sent no ASR start JSON frame") } + if !bytes.Contains(got.frame[start:], []byte(`"large_number":9007199254740993`)) { + t.Fatal("large integer changed") + } var payload map[string]any if err := json.Unmarshal(got.frame[start:], &payload); err != nil { t.Fatalf("SDK ASR start JSON cannot be decoded: %v", err) } audio := payload["audio"].(map[string]any) request := payload["request"].(map[string]any) - if audio["format"] != "pcm_s16le" || audio["sample_rate"] != float64(16000) || audio["channel"] != float64(1) || audio["bits"] != float64(16) || audio["language"] != "zh-CN" || request["result_type"] != tc.resultType || request["enable_nonstream"] != tc.nonstream { + if len(audio) != 4 || len(request) != 9 || audio["format"] != "pcm_s16le" || audio["sample_rate"] != float64(16000) || audio["channel"] != float64(1) || audio["bits"] != float64(16) || request["result_type"] != tc.resultType || request["enable_nonstream"] != tc.nonstream || request["unknown_flag"] != false || request["nullable"] != nil || request["opaque"] != "vendor-only" { t.Fatalf("SDK ASR request omitted approved input/interim fields: audio=%v request=%v", audio, request) } case <-ctx.Done(): diff --git a/internal/ai/bailian_models_test.go b/internal/ai/bailian_models_test.go index 5d1b85e..1540101 100644 --- a/internal/ai/bailian_models_test.go +++ b/internal/ai/bailian_models_test.go @@ -1,51 +1,35 @@ package ai import ( - "encoding/json" "testing" - - "git.ipao.vip/rogee/go-sip/internal/configread" ) func bailianFixture(t *testing.T) Binding { t.Helper() - task, _ := currentFixture(t, "full_ai") + task, ps := currentFixture(t, "full_ai") task = changeCurrentAgent(t, task, func(a map[string]any) { - asr := a["asr"].(map[string]any) - delete(asr, "interim") - asr["provider_id"] = "bailian" - delete(asr, "provider_ref") - asr["model"] = "fun-asr-flash-8k-realtime-2026-01-28" - asr["input"].(map[string]any)["sample_rate_hz"] = 8000 - llm := a["llm"].(map[string]any) - llm["provider_id"] = "bailian" - llm["model"] = "Qwen3.8-Flash" - delete(llm, "provider_ref") - tts := a["tts"].(map[string]any) - tts["provider_id"] = "bailian" - delete(tts, "provider_ref") - tts["protocol"] = TTSProtocolDashScopeTask - tts["model"] = "cosyvoice-v3-flash" - tts["voice"] = "longanyang" - tts["speed"] = 1.25 + for _, role := range []string{"asr", "llm", "tts"} { + a[role].(map[string]any)["provider_ref"] = "ali_bailian" + } + a["asr"].(map[string]any)["model"] = "fun-asr-flash-8k-realtime-2026-01-28" + a["asr"].(map[string]any)["params"] = map[string]any{"sample_rate": 16000, "format": "pcm", "language": "zh"} + a["llm"].(map[string]any)["model"] = "Qwen3.8-Flash" + a["tts"].(map[string]any)["model"] = "cosyvoice-v3-flash" + a["tts"].(map[string]any)["voice"] = "longanyang" + a["tts"].(map[string]any)["params"] = map[string]any{"rate": 1.25, "sample_rate": 16000, "format": "pcm"} }) - var provider configread.Provider - if err := json.Unmarshal([]byte(`{"provider_id":"bailian","provider_code":"ali_bailian","api_endpoint":"https://dashscope.aliyuncs.com/compatible-mode/v1","ws_endpoint":"wss://dashscope.aliyuncs.com/api-ws/v1/inference","api_key":"mock-key","extra_config":{}}`), &provider); err != nil { - t.Fatal(err) - } - bound, err := Bind(task, map[string]configread.Provider{"bailian": provider}) + p := ps["ali_bailian"] + p.Credential = "mock-key" + ps["ali_bailian"] = p + b, err := Bind(task, ps) if err != nil { t.Fatal(err) } - return bound + return b } - func TestTaskModelsShareConnectionOnlyProvider(t *testing.T) { - bound := bailianFixture(t) - if bound.ASR.Provider.ProviderRef != "bailian" || bound.LLM.Provider.ProviderRef != "bailian" || bound.TTS.Provider.ProviderRef != "bailian" { - t.Fatal("task provider_id binding was lost") - } - if bound.TTS.Model != "cosyvoice-v3-flash" || bound.TTS.Voice != "longanyang" || bound.TTS.Speed != 1.25 { - t.Fatal("task TTS model/voice/speed must remain exact") + b := bailianFixture(t) + if b.ASR.Provider.Code != "ali_bailian" || b.LLM.Provider.Code != "ali_bailian" || b.TTS.Provider.Code != "ali_bailian" || string(b.TTS.Params["rate"]) != "1.25" { + t.Fatal("shared connection or opaque params lost") } } diff --git a/internal/ai/bailian_realtime.go b/internal/ai/bailian_realtime.go index ef4dbf4..c0dab41 100644 --- a/internal/ai/bailian_realtime.go +++ b/internal/ai/bailian_realtime.go @@ -5,7 +5,6 @@ package ai // per-call approved endpoints and strict started/finished/error boundaries. // Keep this adapter small; WebSocket framing is handled by the existing library. import ( - "bytes" "context" "encoding/json" "errors" @@ -13,7 +12,6 @@ import ( "log/slog" "net/http" "net/url" - "os/exec" "strconv" "strings" @@ -150,41 +148,18 @@ func runBailianTask(ctx context.Context, endpoint, key string, payload any, send } func recognizeBailian(ctx context.Context, approved ASRConfig, pcm16 []byte) (string, error) { - if strings.TrimSpace(approved.Model) == "" || (approved.Request.SampleRate != 8000 && approved.Request.SampleRate != 16000) { - return "", errors.New("Bailian task ASR needs an explicit model and 8000/16000 Hz PCM16") - } - hint, err := bailianLanguageHint(approved.Language) - if err != nil { - return "", err + if strings.TrimSpace(approved.Model) == "" { + return "", errors.New("Bailian task ASR needs an explicit model") } if len(pcm16) == 0 || len(pcm16)%2 != 0 || len(pcm16) > maxBailianPCMBytes { return "", errors.New("Bailian task ASR input is empty, incomplete or oversized PCM16") } - // Agent media is PCM16/16000. Only convert when the task explicitly selects - // 8000 Hz; a 16000 Hz request sends the original bytes, without relabeling. input := pcm16 - if approved.Request.SampleRate == 8000 { - cmd := exec.CommandContext(ctx, "ffmpeg", "-hide_banner", "-loglevel", "error", "-f", "s16le", "-ar", "16000", "-ac", "1", "-i", "pipe:0", "-f", "s16le", "-ar", "8000", "-ac", "1", "pipe:1") - cmd.Stdin = bytes.NewReader(pcm16) - input, err = cmd.Output() - if err != nil { - exitCode := -1 - var exitErr *exec.ExitError - if errors.As(err, &exitErr) { - exitCode = exitErr.ExitCode() - } - return "", fmt.Errorf("Bailian task ASR PCM16 conversion failed (cause=%T, exit_code=%d, input_bytes=%d)", err, exitCode, len(pcm16)) - } - if len(input) == 0 || len(input)%2 != 0 { - return "", errors.New("Bailian task ASR resampler returned incomplete PCM16") - } - slog.Debug("Bailian task ASR PCM resampling completed", "input_sample_rate", 16000, "output_sample_rate", 8000, "input_bytes", len(pcm16), "output_bytes", len(input)) - } - payload := map[string]any{"task_group": "audio", "task": "asr", "function": "recognition", "model": approved.Model, "parameters": map[string]any{"format": "pcm", "sample_rate": int(approved.Request.SampleRate), "language_hints": []string{hint}}, "input": map[string]any{}} + payload := map[string]any{"task_group": "audio", "task": "asr", "function": "recognition", "model": approved.Model, "parameters": approved.Params, "input": map[string]any{}} var finals []string finalBytes := 0 seen := map[string]string{} - err = runBailianTask(ctx, approved.Provider.Endpoint, approved.Provider.Credential, payload, func(t bailianTask) error { + err := runBailianTask(ctx, approved.Provider.Endpoint, approved.Provider.Credential, payload, func(t bailianTask) error { for off := 0; off < len(input); off += 3200 { if err := t.conn.Write(ctx, websocket.MessageBinary, input[off:min(off+3200, len(input))]); err != nil { return errors.New("Bailian task ASR audio delivery failed") @@ -227,20 +202,14 @@ func recognizeBailian(ctx context.Context, approved ASRConfig, pcm16 []byte) (st } func synthesizeBailianTaskTTS(ctx context.Context, approved TTSConfig, text string) ([]byte, error) { - if approved.Protocol != TTSProtocolDashScopeTask || strings.TrimSpace(approved.Model) == "" || strings.TrimSpace(approved.Voice) == "" || approved.SampleRate != 16000 || approved.Speed < 0.5 || approved.Speed > 2 || strings.TrimSpace(text) == "" { + if approved.Protocol != TTSProtocolDashScopeTask || strings.TrimSpace(approved.Model) == "" || strings.TrimSpace(approved.Voice) == "" || strings.TrimSpace(text) == "" { return nil, errors.New("unapproved Bailian task TTS settings or empty text") } - hints, err := bailianTTSLanguageHints(approved.LanguageType) - if err != nil { - return nil, err - } - parameters := map[string]any{"text_type": "PlainText", "voice": approved.Voice, "format": "pcm", "sample_rate": approved.SampleRate, "rate": approved.Speed} - if len(hints) != 0 { - parameters["language_hints"] = hints - } + parameters := cloneParams(approved.Params) + parameters["voice"], _ = json.Marshal(approved.Voice) payload := map[string]any{"task_group": "audio", "task": "tts", "function": "SpeechSynthesizer", "model": approved.Model, "parameters": parameters, "input": map[string]any{}} var audio []byte - err = runBailianTask(ctx, approved.Provider.WSEndpoint, approved.Provider.Credential, payload, func(t bailianTask) error { + err := runBailianTask(ctx, approved.Provider.WSEndpoint, approved.Provider.Credential, payload, func(t bailianTask) error { if err := t.command(ctx, "continue-task", map[string]any{"input": map[string]any{"text": text}}); err != nil { return err } @@ -260,6 +229,6 @@ func synthesizeBailianTaskTTS(ctx context.Context, approved TTSConfig, text stri if len(audio) == 0 || len(audio)%2 != 0 { return nil, errors.New("Bailian task TTS finished without complete PCM16 audio") } - slog.Debug("Bailian task TTS complete PCM16 received", "sample_rate", approved.SampleRate, "bytes", len(audio)) + slog.Debug("Bailian task TTS complete PCM16 received", "bytes", len(audio)) return audio, nil } diff --git a/internal/ai/bailian_realtime_test.go b/internal/ai/bailian_realtime_test.go index 84631bf..57d5688 100644 --- a/internal/ai/bailian_realtime_test.go +++ b/internal/ai/bailian_realtime_test.go @@ -85,7 +85,7 @@ func TestCosyVoiceTaskUsesExactModelVoiceRateAndCompletePCM(t *testing.T) { t.Fatalf("synthesis: %v %v requests=%d", audio, err, requests.Load()) } } -func TestFunASRResamplesAndReturnsOnlyConfirmedFinalSentences(t *testing.T) { +func TestFunASRPreservesPCMAndReturnsOnlyConfirmedFinalSentences(t *testing.T) { bound := bailianFixture(t) var bytesReceived atomic.Int64 server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -98,7 +98,7 @@ func TestFunASRResamplesAndReturnsOnlyConfirmedFinalSentences(t *testing.T) { ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second) defer cancel() cmd := mockRead(t, c, ctx) - if cmd.Payload["model"] != "fun-asr-flash-8k-realtime-2026-01-28" || cmd.Payload["parameters"].(map[string]any)["sample_rate"] != float64(8000) { + if cmd.Payload["model"] != "fun-asr-flash-8k-realtime-2026-01-28" || cmd.Payload["parameters"].(map[string]any)["sample_rate"] != float64(16000) { t.Error("ASR model/sample rate changed") } mockEvent(c, ctx, cmd.Header.TaskID, "task-started", map[string]any{}) @@ -129,7 +129,7 @@ func TestFunASRResamplesAndReturnsOnlyConfirmedFinalSentences(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) defer cancel() text, err := bound.Recognize(ctx, make([]byte, 32000)) - if err != nil || text != "第一句。第二句。" || bytesReceived.Load() != 16000 { + if err != nil || text != "第一句。第二句。" || bytesReceived.Load() != 32000 { t.Fatalf("ASR: text=%q err=%v sent=%d", text, err, bytesReceived.Load()) } } diff --git a/internal/ai/bailian_tts.go b/internal/ai/bailian_tts.go index eb62f9b..7911548 100644 --- a/internal/ai/bailian_tts.go +++ b/internal/ai/bailian_tts.go @@ -1,194 +1,22 @@ package ai import ( - "bytes" "context" - "encoding/json" "errors" - "fmt" - "io" - "net" - "net/http" - "net/url" - "os/exec" - "strconv" - "strings" ) -const ( - maxBailianReplyBytes = 64 << 10 - maxBailianAudioBytes = 8 << 20 - maxBailianPCMBytes = 2 << 20 // ponytail: enough for about one minute; raise only with an approved longer sentence. -) +const maxBailianPCMBytes = 2 << 20 // Per-turn media allocation bound; unrelated to prompt length. -func synthesizeBailianTTS(ctx context.Context, approved TTSConfig, text string) ([]byte, error) { - if _, err := exec.LookPath("ffmpeg"); err != nil { - return nil, errors.New("Bailian TTS requires the configured ffmpeg audio converter") - } - if err := bailianURL(approved.Provider.Endpoint, false); err != nil { - return nil, fmt.Errorf("approved Bailian endpoint: %w", err) - } - if approved.Protocol != TTSProtocolDashScopeHTTP || strings.TrimSpace(approved.Model) == "" || strings.TrimSpace(approved.Voice) == "" || strings.TrimSpace(approved.LanguageType) == "" || approved.Speed != 1 || approved.SampleRate != 16000 || approved.Provider.Credential == "" || text == "" { - return nil, errors.New("approved Bailian TTS settings or text are incomplete") - } - // Never follow a provider or audio redirect into an unapproved second request. - client := http.Client{CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }} - request := struct { - Model string `json:"model"` - Input struct { - Text string `json:"text"` - Voice string `json:"voice"` - LanguageType string `json:"language_type"` - } `json:"input"` - }{Model: approved.Model} - request.Input.Text, request.Input.Voice, request.Input.LanguageType = text, approved.Voice, approved.LanguageType - body, err := json.Marshal(request) - if err != nil { - return nil, errors.New("encode approved Bailian request") - } - req, err := http.NewRequestWithContext(ctx, http.MethodPost, approved.Provider.Endpoint, bytes.NewReader(body)) - if err != nil { - return nil, errors.New("create approved Bailian request") - } - req.Header.Set("Authorization", "Bearer "+approved.Provider.Credential) - req.Header.Set("Content-Type", "application/json") - response, err := client.Do(req) - if err != nil { - return nil, fmt.Errorf("Bailian TTS request failed: type=%T", err) - } - defer response.Body.Close() - if response.StatusCode/100 != 2 { - return nil, fmt.Errorf("Bailian TTS rejected request: HTTP %d", response.StatusCode) - } - raw, err := io.ReadAll(io.LimitReader(response.Body, maxBailianReplyBytes+1)) - if err != nil || len(raw) > maxBailianReplyBytes { - return nil, errors.New("Bailian TTS response is unreadable or oversized") - } - var result struct { - Output struct { - Audio struct { - URL string `json:"url"` - } `json:"audio"` - } `json:"output"` - } - if err := json.Unmarshal(raw, &result); err != nil || result.Output.Audio.URL == "" { - return nil, errors.New("Bailian TTS response has no complete audio reference") - } - audioURL, err := bailianAudioDownloadURL(result.Output.Audio.URL) - if err != nil { - return nil, errors.New("Bailian TTS audio reference is invalid") - } - audioRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, audioURL, nil) - if err != nil { - return nil, errors.New("create Bailian audio request") - } - audioResponse, err := client.Do(audioRequest) - if err != nil { - return nil, fmt.Errorf("Bailian audio download failed: type=%T", err) - } - defer audioResponse.Body.Close() - if audioResponse.StatusCode/100 != 2 { - return nil, fmt.Errorf("Bailian audio download rejected: HTTP %d", audioResponse.StatusCode) - } - encoded, err := io.ReadAll(io.LimitReader(audioResponse.Body, maxBailianAudioBytes+1)) - if err != nil || len(encoded) == 0 || len(encoded) > maxBailianAudioBytes { - return nil, errors.New("Bailian audio is empty, unreadable or oversized") - } - cmd := exec.CommandContext(ctx, "ffmpeg", "-hide_banner", "-loglevel", "error", "-nostdin", "-i", "pipe:0", "-f", "s16le", "-ac", "1", "-ar", "16000", "pipe:1") - cmd.Stdin = bytes.NewReader(encoded) - cmd.Stderr = io.Discard - var pcm boundedPCM - cmd.Stdout = &pcm - if err := cmd.Run(); err != nil { - return nil, fmt.Errorf("Bailian audio conversion failed: type=%T", err) - } - if err := ctx.Err(); err != nil { - return nil, fmt.Errorf("Bailian TTS deadline or cancellation: %w", err) - } - if pcm.Len() == 0 || pcm.Len()%2 != 0 { - return nil, errors.New("Bailian TTS returned incomplete PCM16 audio") - } - return pcm.Bytes(), nil -} - -// bailianFailureSummary emits only fixed categories and HTTP status codes; -// provider responses, signed URLs, credentials and prompt text stay private. +// Never log vendor payloads, prompt text or credentials. func bailianFailureSummary(err error) (string, int) { if err == nil { return "none", 0 } - message := err.Error() - for _, item := range []struct{ prefix, kind string }{ - {"Bailian TTS rejected request: HTTP ", "generation_http"}, - {"Bailian audio download rejected: HTTP ", "audio_http"}, - } { - if status, err := strconv.Atoi(strings.TrimPrefix(message, item.prefix)); strings.HasPrefix(message, item.prefix) && err == nil && status >= 100 && status <= 599 { - return item.kind, status - } + if errors.Is(err, context.DeadlineExceeded) { + return "deadline", 0 } - for _, item := range []struct{ prefix, kind string }{ - {"Bailian TTS requires the configured ffmpeg audio converter", "missing_converter"}, - {"approved Bailian TTS settings or text are incomplete", "invalid_settings"}, - {"encode approved Bailian request", "request_encoding"}, - {"create approved Bailian request", "request_construction"}, - {"approved Bailian endpoint:", "invalid_endpoint"}, - {"Bailian TTS response is unreadable or oversized", "invalid_response"}, - {"Bailian TTS response has no complete audio reference", "missing_audio"}, - {"Bailian TTS audio reference is invalid", "invalid_audio_reference"}, - {"create Bailian audio request", "audio_request_construction"}, - {"Bailian audio is empty, unreadable or oversized", "invalid_audio"}, - {"Bailian TTS PCM exceeds approved media ceiling", "pcm_ceiling"}, - {"Bailian TTS request failed:", "generation_transport"}, - {"Bailian audio download failed:", "audio_transport"}, - {"Bailian audio conversion failed:", "audio_conversion"}, - {"Bailian TTS deadline or cancellation:", "deadline"}, - {"Bailian TTS returned incomplete PCM16 audio", "incomplete_audio"}, - } { - if strings.HasPrefix(message, item.prefix) { - return item.kind, 0 - } + if errors.Is(err, context.Canceled) { + return "canceled", 0 } - return "unknown", 0 -} - -type boundedPCM struct{ bytes.Buffer } - -func (w *boundedPCM) Write(data []byte) (int, error) { - if len(data) > maxBailianPCMBytes-w.Len() { - return 0, errors.New("Bailian TTS PCM exceeds approved media ceiling") - } - return w.Buffer.Write(data) -} - -// Bailian documents HTTP signed OSS links. Preserve the path and signature -// byte-for-byte, but fetch only over HTTPS; never fall back to HTTP. -func bailianAudioDownloadURL(raw string) (string, error) { - parsed, err := url.Parse(raw) - if err != nil { - return "", errors.New("invalid audio address") - } - if strings.HasPrefix(raw, "http://") && parsed.Host == "dashscope-result-bj.oss-cn-beijing.aliyuncs.com" && parsed.User == nil && parsed.Fragment == "" { - raw = "https" + raw[len("http"):] - } - if err := bailianURL(raw, true); err != nil { - return "", err - } - return raw, nil -} - -func bailianURL(raw string, signedAudio bool) error { - parsed, err := url.Parse(raw) - if err != nil || parsed.Host == "" || parsed.User != nil || parsed.Fragment != "" { - return errors.New("HTTPS URL without embedded identity required") - } - if parsed.Scheme != "https" { - host := parsed.Hostname() - if parsed.Scheme != "http" || (host != "localhost" && !net.ParseIP(host).IsLoopback()) { - return errors.New("HTTPS URL required outside isolated local Mock") - } - } - if !signedAudio && (parsed.RawQuery != "" || !strings.HasSuffix(parsed.Path, "/api/v1/services/aigc/multimodal-generation/generation")) { - return errors.New("approved Bailian generation endpoint required") - } - return nil + return "task_websocket", 0 } diff --git a/internal/ai/bailian_tts_test.go b/internal/ai/bailian_tts_test.go index 1656dbd..dbb15e6 100644 --- a/internal/ai/bailian_tts_test.go +++ b/internal/ai/bailian_tts_test.go @@ -2,158 +2,47 @@ package ai import ( "context" - "encoding/json" - "errors" + "github.com/coder/websocket" "net/http" "net/http/httptest" - "os/exec" "strings" "sync/atomic" "testing" - - "git.ipao.vip/rogee/go-sip/internal/media" + "time" ) func mockBailianTTS(t *testing.T, pcm []byte, observe func(string)) (string, *atomic.Int32) { t.Helper() - wav, _, err := media.EncodeMonoWAV(pcm, 1024) - if err != nil { - t.Fatal(err) - } calls := &atomic.Int32{} - var server *httptest.Server - server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/audio" && r.Method == http.MethodGet { - _, _ = w.Write(wav) - return - } - if r.URL.Path != "/api/v1/services/aigc/multimodal-generation/generation" || r.Method != http.MethodPost { - http.Error(w, "unexpected endpoint", http.StatusNotFound) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + c, err := websocket.Accept(w, r, nil) + if err != nil { return } + defer c.CloseNow() + ctx, cancel := context.WithTimeout(r.Context(), 3*time.Second) + defer cancel() + start := mockRead(t, c, ctx) calls.Add(1) - var request struct { - Input struct { - Text string `json:"text"` - } `json:"input"` + id := start.Header.TaskID + mockEvent(c, ctx, id, "task-started", map[string]any{}) + next := mockRead(t, c, ctx) + text := next.Payload["input"].(map[string]any)["text"].(string) + if observe != nil { + observe(text) } - if err := json.NewDecoder(r.Body).Decode(&request); err != nil { - http.Error(w, "invalid request", http.StatusBadRequest) + mockRead(t, c, ctx) + if err := c.Write(ctx, websocket.MessageBinary, pcm); err != nil { return } - if observe != nil { - observe(request.Input.Text) - } - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"output": map[string]any{"audio": map[string]any{"url": server.URL + "/audio?sig=fixture"}}}) + mockEvent(c, ctx, id, "task-finished", map[string]any{}) })) t.Cleanup(server.Close) - return server.URL + "/api/v1/services/aigc/multimodal-generation/generation", calls + return "ws" + strings.TrimPrefix(server.URL, "http"), calls } - -func TestBailianAudioDownloadURLUpgradesDocumentedOSSHostOnly(t *testing.T) { - const from = "http://dashscope-result-bj.oss-cn-beijing.aliyuncs.com/ab%2Fcd.wav?Expires=123&Signature=fake%2Bvalue" - const want = "https://dashscope-result-bj.oss-cn-beijing.aliyuncs.com/ab%2Fcd.wav?Expires=123&Signature=fake%2Bvalue" - got, err := bailianAudioDownloadURL(from) - if err != nil || got != want { - t.Fatalf("audio address = %q, %v; want exact HTTPS upgrade", got, err) - } - for _, address := range []string{ - "http://other.oss-cn-beijing.aliyuncs.com/audio.wav?Signature=fake", - "http://dashscope-result-bj.oss-cn-beijing.aliyuncs.com:80/audio.wav?Signature=fake", - "http://user@dashscope-result-bj.oss-cn-beijing.aliyuncs.com/audio.wav?Signature=fake", - "http://dashscope-result-bj.oss-cn-beijing.aliyuncs.com/audio.wav#fragment", - } { - if _, err := bailianAudioDownloadURL(address); err == nil { - t.Fatalf("unapproved HTTP audio address accepted") - } - } - const alreadyHTTPS = "https://dashscope-result-bj.oss-cn-beijing.aliyuncs.com/audio.wav?Signature=fake" - if got, err := bailianAudioDownloadURL(alreadyHTTPS); err != nil || got != alreadyHTTPS { - t.Fatalf("existing HTTPS address changed: %v", err) - } - if got, err := bailianAudioDownloadURL("http://127.0.0.1/audio.wav"); err != nil || got != "http://127.0.0.1/audio.wav" { - t.Fatalf("isolated local Mock address changed: %v", err) - } -} - func TestBailianFailureSummaryNeverLogsProviderPayload(t *testing.T) { - for _, tc := range []struct { - message, kind string - status int - }{ - {"Bailian TTS rejected request: HTTP 403", "generation_http", 403}, - {"Bailian audio download rejected: HTTP 502", "audio_http", 502}, - {"Bailian TTS response has no complete audio reference", "missing_audio", 0}, - {"Bailian TTS requires the configured ffmpeg audio converter", "missing_converter", 0}, - {"approved Bailian TTS settings or text are incomplete", "invalid_settings", 0}, - {"Bailian TTS response is unreadable or oversized", "invalid_response", 0}, - {"Bailian TTS audio reference is invalid", "invalid_audio_reference", 0}, - {"Bailian audio is empty, unreadable or oversized", "invalid_audio", 0}, - {"approved Bailian endpoint: HTTPS URL required outside isolated local Mock", "invalid_endpoint", 0}, - {"Authorization: Bearer private-signature and prompt", "unknown", 0}, - } { - kind, status := bailianFailureSummary(errors.New(tc.message)) - if kind != tc.kind || status != tc.status || strings.Contains(kind, "private") { - t.Fatalf("unsafe failure summary: kind=%q status=%d", kind, status) - } - } -} - -func TestBailianTTSDownloadFailureDoesNotLeakSignedURLOrRetry(t *testing.T) { - if _, err := exec.LookPath("ffmpeg"); err != nil { - t.Skip("Bailian TTS conversion requires ffmpeg") - } - var calls, downloads atomic.Int32 - var server *httptest.Server - server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/audio" { - downloads.Add(1) - http.Error(w, "private-data", http.StatusForbidden) - return - } - calls.Add(1) - _ = json.NewEncoder(w).Encode(map[string]any{"output": map[string]any{"audio": map[string]any{"url": server.URL + "/audio?sig=private-signature"}}}) - })) - defer server.Close() - task, providers := currentFixture(t, "full_ai") - p := providers["tts-example"] - p.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation" - providers[p.ProviderRef] = p - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) - } - _, err = bound.Synthesize(context.Background(), "批准的回复") - if err == nil || !strings.Contains(err.Error(), "HTTP 403") || strings.Contains(err.Error(), "private-signature") || strings.Contains(err.Error(), "private-data") || strings.Contains(err.Error(), p.Credential) || calls.Load() != 1 || downloads.Load() != 1 { - t.Fatalf("download failure must be explicit without leaking/retrying: calls=%d downloads=%d err=%v", calls.Load(), downloads.Load(), err) - } -} - -func TestBailianTTSDoesNotFollowChargeableRedirect(t *testing.T) { - if _, err := exec.LookPath("ffmpeg"); err != nil { - t.Skip("Bailian TTS conversion requires ffmpeg") - } - var calls, redirected atomic.Int32 - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/charged-again" { - redirected.Add(1) - return - } - calls.Add(1) - http.Redirect(w, r, "/charged-again", http.StatusTemporaryRedirect) - })) - defer server.Close() - task, providers := currentFixture(t, "full_ai") - p := providers["tts-example"] - p.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation" - providers[p.ProviderRef] = p - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) - } - _, err = bound.Synthesize(context.Background(), "批准的回复") - if err == nil || !strings.Contains(err.Error(), "HTTP 307") || calls.Load() != 1 || redirected.Load() != 0 { - t.Fatalf("TTS redirect cannot create a second provider request: calls=%d redirected=%d err=%v", calls.Load(), redirected.Load(), err) + kind, status := bailianFailureSummary(context.DeadlineExceeded) + if kind != "deadline" || status != 0 { + t.Fatal("context failure lost") } } diff --git a/internal/ai/binding.go b/internal/ai/binding.go index 0ab5c0f..b98a495 100644 --- a/internal/ai/binding.go +++ b/internal/ai/binding.go @@ -4,26 +4,22 @@ import ( "encoding/json" "errors" "fmt" - "math" "net/url" "strings" "time" "git.ipao.vip/rogee/go-sip/internal/configread" "git.ipao.vip/rogee/go-sip/internal/contract" - doubaospeech "github.com/GizClaw/doubao-speech-go" ) -// Binding is a per-call, immutable binding of the approved task and -// provider settings. It must be built before dispatching any side effects. -// Credential values are never included in errors or logs. +// Binding freezes the task, credentials and raw scalar parameters before any side effect. +// Errors never contain credentials, prompt text or parameter values. type Binding struct { Mode string ASR ASRConfig LLM *LLMConfig TTS *TTSConfig Prompt string - PromptMaxBytes int AllowedVariables []string Opening string HangupKeywords []HangupKeyword @@ -32,33 +28,21 @@ type Binding struct { type ASRConfig struct { Provider configread.Provider - Request doubaospeech.ASRV2Config Model string - Language string - Timeout time.Duration + Params map[string]json.RawMessage } - type LLMConfig struct { - Provider configread.Provider - Model string - Temperature *float64 - MaxTokens *int64 - Timeout time.Duration + Provider configread.Provider + Model string + Params map[string]json.RawMessage } - type TTSConfig struct { - Provider configread.Provider - Protocol string - Model string - Voice string - LanguageType string - Speed float64 - SampleRate int - Timeout time.Duration + Provider configread.Provider + Protocol string + Model string + Voice string + Params map[string]json.RawMessage } - -// ConversationConfig preserves the task's explicit dialogue limits. An -// absent positive limit remains zero; an explicit false remains non-nil. type ConversationConfig struct { AllowInterrupt *bool SilenceTimeout time.Duration @@ -67,47 +51,20 @@ type ConversationConfig struct { SentenceMaxChars int MaxPendingAudioChunks int } - +type modelSettings struct { + ProviderRef string `json:"provider_ref"` + Model string `json:"model"` + Params map[string]json.RawMessage `json:"params"` + Voice string `json:"voice,omitempty"` +} type currentAgentSettings struct { - Mode string `json:"mode"` - ASR struct { - ProviderRef string `json:"provider_id"` - Model string `json:"model"` - Language string `json:"language"` - Interim *bool `json:"interim"` - TimeoutMS *int64 `json:"timeout_ms"` - Input struct { - Encoding string `json:"encoding"` - SampleRateHz int `json:"sample_rate_hz"` - Channels int `json:"channels"` - SampleWidthBytes int `json:"sample_width_bytes"` - } `json:"input"` - } `json:"asr"` - LLM *struct { - ProviderRef string `json:"provider_id"` - Model string `json:"model"` - Temperature *float64 `json:"temperature"` - MaxTokens *int64 `json:"max_tokens"` - TimeoutMS *int64 `json:"timeout_ms"` - } `json:"llm"` - TTS *struct { - ProviderRef string `json:"provider_id"` - Protocol string `json:"protocol"` - Model string `json:"model"` - Voice string `json:"voice"` - LanguageType string `json:"language_type"` - Speed *float64 `json:"speed"` - TimeoutMS *int64 `json:"timeout_ms"` - Format struct { - Encoding string `json:"encoding"` - SampleRateHz int `json:"sample_rate_hz"` - Channels int `json:"channels"` - } `json:"format"` - } `json:"tts"` + Mode string `json:"mode"` + ASR modelSettings `json:"asr"` + LLM *modelSettings `json:"llm"` + TTS *modelSettings `json:"tts"` Prompt *struct { Text string `json:"text"` AllowedVariables []string `json:"allowed_variables"` - MaxBytes *int `json:"max_bytes"` } `json:"prompt"` Conversation *struct { Opening string `json:"opening"` @@ -121,8 +78,6 @@ type currentAgentSettings struct { } `json:"conversation"` } -// Bind rejects schema-valid settings which the approved provider cannot -// express; neither defaults nor lossy format/speed conversions are permitted. func Bind(task configread.Task, providers map[string]configread.Provider) (Binding, error) { if len(task.Raw) == 0 { return Binding{}, errors.New("approved task snapshot is missing") @@ -132,181 +87,91 @@ func Bind(task configread.Task, providers map[string]configread.Provider) (Bindi } var frozen configread.Task if err := json.Unmarshal(task.Raw, &frozen); err != nil { - return Binding{}, fmt.Errorf("decode task snapshot: %w", err) + return Binding{}, errors.New("decode task snapshot") } if frozen.Resource != "task_config" || frozen.TaskID != task.TaskID || frozen.TenantID != task.TenantID || frozen.DispatcherID != task.DispatcherID { return Binding{}, errors.New("approved task snapshot identity mismatch") } - var settings currentAgentSettings - if err := json.Unmarshal(frozen.Agent.Raw, &settings); err != nil { - return Binding{}, fmt.Errorf("decode immutable AI settings: %w", err) + var s currentAgentSettings + if err := json.Unmarshal(frozen.Agent.Raw, &s); err != nil { + return Binding{}, errors.New("decode immutable AI settings") } - bound := Binding{Mode: settings.Mode} - asrProvider, err := currentProvider(providers, settings.ASR.ProviderRef, "asr", "") + b := Binding{Mode: s.Mode} + p, err := currentProvider(providers, s.ASR.ProviderRef, "asr") if err != nil { return Binding{}, err } - if settings.ASR.Input.Encoding != "pcm_s16le" || settings.ASR.Input.Channels != 1 || settings.ASR.Input.SampleWidthBytes != 2 { - return Binding{}, errors.New("ASR input format is unsupported by selected SDK") + if err := reservedParams(s.ASR.Params, "reqid", "sequence", "model_name"); err != nil { + return Binding{}, err } - sampleRate, err := currentSampleRate(settings.ASR.Input.SampleRateHz) - if err != nil { - return Binding{}, fmt.Errorf("ASR input: %w", err) + b.ASR = ASRConfig{Provider: p, Model: s.ASR.Model, Params: s.ASR.Params} + if s.Mode == "asr_only" { + return b, nil } - lang := doubaospeech.Language(settings.ASR.Language) - if asrProvider.Code == "ali_bailian" { - if settings.ASR.Interim != nil { - return Binding{}, errors.New("Bailian task ASR protocol cannot express interim selection") - } - if strings.TrimSpace(settings.ASR.Model) == "" || (sampleRate != 8000 && sampleRate != 16000) { - return Binding{}, errors.New("Bailian task ASR requires an explicit model and 8000/16000 Hz PCM16") - } - if _, err := bailianLanguageHint(settings.ASR.Language); err != nil { - return Binding{}, fmt.Errorf("ASR language: %w", err) - } - } else { - if sampleRate != 16000 { - return Binding{}, errors.New("ASR media requires 16000 Hz PCM16") - } - switch lang { - case doubaospeech.LanguageZhCN, doubaospeech.LanguageEnUS, doubaospeech.LanguageJaJP, doubaospeech.LanguageKoKR: - default: - return Binding{}, errors.New("ASR language is unsupported by selected SDK") - } + if s.Mode != "full_ai" || s.LLM == nil || s.TTS == nil || s.Prompt == nil { + return Binding{}, errors.New("full AI settings are incomplete") } - asrTimeout, err := currentTimeout(settings.ASR.TimeoutMS) - if err != nil { - return Binding{}, fmt.Errorf("ASR timeout: %w", err) - } - asrRequest := doubaospeech.ASRV2Config{ - Format: doubaospeech.FormatPCMS16LE, SampleRate: sampleRate, - Channel: 1, Bits: 16, Language: lang, - Request: &doubaospeech.ASRV2RequestConfig{ModelName: settings.ASR.Model}, - } - if settings.ASR.Interim != nil { - asrRequest.Request.EnableNonstream = new(bool) - *asrRequest.Request.EnableNonstream = !*settings.ASR.Interim - if *settings.ASR.Interim { - asrRequest.ResultType = "full" - } else { - asrRequest.ResultType = "single" - } - asrRequest.Request.ResultType = asrRequest.ResultType - } - bound.ASR = ASRConfig{Provider: asrProvider, Request: asrRequest, Model: settings.ASR.Model, Language: settings.ASR.Language, Timeout: asrTimeout} - - switch settings.Mode { - case "asr_only": - return bound, nil - case "full_ai": - if settings.LLM == nil || settings.TTS == nil || settings.Prompt == nil || settings.Conversation == nil { - return Binding{}, errors.New("full AI settings are incomplete") - } - default: - return Binding{}, errors.New("AI mode is unsupported") - } - llmProvider, err := currentProvider(providers, settings.LLM.ProviderRef, "llm", "") + p, err = currentProvider(providers, s.LLM.ProviderRef, "llm") if err != nil { return Binding{}, err } - llmTimeout, err := currentTimeout(settings.LLM.TimeoutMS) - if err != nil { - return Binding{}, fmt.Errorf("LLM timeout: %w", err) + if err := reservedParams(s.LLM.Params, "model", "messages"); err != nil { + return Binding{}, err } - bound.LLM = &LLMConfig{Provider: llmProvider, Model: settings.LLM.Model, Temperature: settings.LLM.Temperature, MaxTokens: settings.LLM.MaxTokens, Timeout: llmTimeout} - ttsProvider, err := currentProvider(providers, settings.TTS.ProviderRef, "tts", settings.TTS.Protocol) + b.LLM = &LLMConfig{Provider: p, Model: s.LLM.Model, Params: s.LLM.Params} + p, err = currentProvider(providers, s.TTS.ProviderRef, "tts") if err != nil { return Binding{}, err } - if settings.TTS.Format.Encoding != "pcm_s16le" || settings.TTS.Format.Channels != 1 { - return Binding{}, errors.New("TTS output encoding/channel is unsupported by selected SDK") - } - ttsSampleRate, err := currentSampleRate(settings.TTS.Format.SampleRateHz) - if err != nil { - return Binding{}, fmt.Errorf("TTS output: %w", err) - } - if ttsSampleRate != doubaospeech.SampleRate(16000) { - return Binding{}, errors.New("TTS media requires 16000 Hz PCM16") - } - if strings.TrimSpace(settings.TTS.Model) == "" || strings.TrimSpace(settings.TTS.Voice) == "" || strings.TrimSpace(settings.TTS.LanguageType) == "" || settings.TTS.Speed == nil { - return Binding{}, errors.New("TTS model, voice, language and speed must be explicit") - } - switch settings.TTS.Protocol { - case TTSProtocolDashScopeHTTP: - if *settings.TTS.Speed != 1 { - return Binding{}, errors.New("TTS HTTP protocol cannot express the requested speed") - } - case TTSProtocolDashScopeTask: - if *settings.TTS.Speed < 0.5 || *settings.TTS.Speed > 2 { - return Binding{}, errors.New("TTS task protocol speed is outside 0.5–2") - } - if _, err := bailianTTSLanguageHints(settings.TTS.LanguageType); err != nil { - return Binding{}, fmt.Errorf("TTS language: %w", err) - } - default: - return Binding{}, errors.New("TTS protocol is unsupported") - } - ttsTimeout, err := currentTimeout(settings.TTS.TimeoutMS) - if err != nil { - return Binding{}, fmt.Errorf("TTS timeout: %w", err) - } - bound.TTS = &TTSConfig{ - Provider: ttsProvider, Protocol: settings.TTS.Protocol, Model: settings.TTS.Model, Voice: settings.TTS.Voice, LanguageType: settings.TTS.LanguageType, - Speed: *settings.TTS.Speed, SampleRate: int(ttsSampleRate), Timeout: ttsTimeout, - } - bound.Prompt = settings.Prompt.Text - bound.AllowedVariables = append([]string(nil), settings.Prompt.AllowedVariables...) - if settings.Prompt.MaxBytes != nil { - bound.PromptMaxBytes = *settings.Prompt.MaxBytes - if len(bound.Prompt) > bound.PromptMaxBytes { - return Binding{}, errors.New("prompt exceeds its approved byte limit") - } - } - if settings.Conversation.AllowInterrupt != nil && *settings.Conversation.AllowInterrupt { - return Binding{}, errors.New("conversation interrupt is unsupported by current media controller") - } - silence, err := currentTimeout(settings.Conversation.SilenceTimeoutMS) - if err != nil { - return Binding{}, fmt.Errorf("conversation silence timeout: %w", err) - } - maxDuration, err := currentTimeout(settings.Conversation.MaxDurationMS) - if err != nil { - return Binding{}, fmt.Errorf("conversation duration: %w", err) - } - bound.Conversation = ConversationConfig{ - AllowInterrupt: settings.Conversation.AllowInterrupt, - SilenceTimeout: silence, MaxDuration: maxDuration, - MaxTurns: settings.Conversation.MaxTurns, - SentenceMaxChars: settings.Conversation.SentenceMaxChars, - MaxPendingAudioChunks: settings.Conversation.MaxPendingAudioChunks, - } - bound.Opening = settings.Conversation.Opening - bound.HangupKeywords = cloneHangupKeywords(settings.Conversation.HangupKeywords) - if err := validateHangupKeywords(bound.HangupKeywords); err != nil { + if err := reservedParams(s.TTS.Params, "voice"); err != nil { return Binding{}, err } - return bound, nil + b.TTS = &TTSConfig{Provider: p, Protocol: TTSProtocolDashScopeTask, Model: s.TTS.Model, Voice: s.TTS.Voice, Params: s.TTS.Params} + b.Prompt = s.Prompt.Text + b.AllowedVariables = append([]string(nil), s.Prompt.AllowedVariables...) + if s.Conversation != nil { + c := s.Conversation + b.Opening = c.Opening + b.HangupKeywords = cloneHangupKeywords(c.HangupKeywords) + if err := validateHangupKeywords(b.HangupKeywords); err != nil { + return Binding{}, err + } + b.Conversation = ConversationConfig{AllowInterrupt: c.AllowInterrupt, MaxTurns: c.MaxTurns, SentenceMaxChars: c.SentenceMaxChars, MaxPendingAudioChunks: c.MaxPendingAudioChunks} + if c.AllowInterrupt != nil && *c.AllowInterrupt { + return Binding{}, errors.New("allow_interrupt=true is unsupported by current media controller") + } + if c.SilenceTimeoutMS != nil { + b.Conversation.SilenceTimeout = time.Duration(*c.SilenceTimeoutMS) * time.Millisecond + } + if c.MaxDurationMS != nil { + b.Conversation.MaxDuration = time.Duration(*c.MaxDurationMS) * time.Millisecond + } + } + return b, nil } -func currentProvider(providers map[string]configread.Provider, ref, role, protocol string) (configread.Provider, error) { - p, found := providers[ref] - if !found || ref == "" || p.ProviderRef != ref || p.Credential == "" { - return configread.Provider{}, fmt.Errorf("%s provider connection is missing or incomplete", role) - } - if len(p.ExtraConfig) != 0 { - var extra map[string]json.RawMessage - if err := json.Unmarshal(p.ExtraConfig, &extra); err != nil || len(extra) != 0 { - return configread.Provider{}, fmt.Errorf("%s provider has unsupported connection parameters", role) +// Protocol identity belongs to the task/session, not vendor business parameters. +// Reject collisions explicitly instead of overwriting or dropping supplied params. +func reservedParams(params map[string]json.RawMessage, keys ...string) error { + for _, key := range keys { + if _, found := params[key]; found { + return fmt.Errorf("params conflicts with protocol field %q", key) } } + return nil +} + +func currentProvider(providers map[string]configread.Provider, ref, role string) (configread.Provider, error) { + p, ok := providers[ref] + if !ok || p.Code != ref || p.Credential == "" { + return configread.Provider{}, fmt.Errorf("%s provider connection is missing or incomplete", role) + } switch role { case "asr": if p.Code != "volcengine" && p.Code != "ali_bailian" { return configread.Provider{}, errors.New("ASR provider is unsupported") } - if p.Code == "ali_bailian" { - p.Endpoint = p.WSEndpoint - } + p.Endpoint = p.WSEndpoint case "llm": if p.Code != "openai_compatible" && p.Code != "ali_bailian" { return configread.Provider{}, errors.New("LLM provider is unsupported") @@ -315,44 +180,20 @@ func currentProvider(providers map[string]configread.Provider, ref, role, protoc if p.Code != "ali_bailian" { return configread.Provider{}, errors.New("TTS provider is unsupported") } - switch protocol { - case TTSProtocolDashScopeTask: - p.Endpoint = p.WSEndpoint - case TTSProtocolDashScopeHTTP: - default: - return configread.Provider{}, errors.New("TTS protocol is unsupported") - } + p.Endpoint = p.WSEndpoint default: return configread.Provider{}, errors.New("AI purpose is unsupported") } u, err := url.Parse(p.Endpoint) - if err != nil || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" || (u.Scheme != "https" && u.Scheme != "http" && u.Scheme != "wss" && u.Scheme != "ws") || strings.ContainsAny(p.Credential, "\r\n") { + if err != nil || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" || strings.ContainsAny(p.Credential, "\r\n") { return configread.Provider{}, fmt.Errorf("%s provider endpoint or credential is invalid", role) } - if (role == "asr" && p.Code == "ali_bailian" || role == "tts" && protocol == TTSProtocolDashScopeTask) && u.Scheme != "wss" && u.Scheme != "ws" { - return configread.Provider{}, errors.New("selected speech model requires a WebSocket connection") - } - if (role == "llm" || role == "tts" && protocol == TTSProtocolDashScopeHTTP) && u.Scheme != "http" && u.Scheme != "https" { - return configread.Provider{}, errors.New("selected protocol requires an HTTP connection") + if role == "llm" { + if u.Scheme != "https" && u.Scheme != "http" { + return configread.Provider{}, errors.New("LLM requires an HTTP connection") + } + } else if u.Scheme != "wss" && u.Scheme != "ws" { + return configread.Provider{}, errors.New("speech model requires a WebSocket connection") } return p, nil } - -func currentSampleRate(rate int) (doubaospeech.SampleRate, error) { - switch rate { - case 8000, 16000, 22050, 24000, 32000, 44100, 48000: - return doubaospeech.SampleRate(rate), nil - default: - return 0, errors.New("sample rate is unsupported by selected SDK") - } -} - -func currentTimeout(ms *int64) (time.Duration, error) { - if ms == nil { - return 0, nil - } - if *ms <= 0 || *ms > int64(math.MaxInt64/int64(time.Millisecond)) { - return 0, errors.New("timeout is outside representable range") - } - return time.Duration(*ms) * time.Millisecond, nil -} diff --git a/internal/ai/binding_test.go b/internal/ai/binding_test.go index 22b2233..7163225 100644 --- a/internal/ai/binding_test.go +++ b/internal/ai/binding_test.go @@ -5,10 +5,8 @@ import ( "os" "strings" "testing" - "time" "git.ipao.vip/rogee/go-sip/internal/configread" - doubaospeech "github.com/GizClaw/doubao-speech-go" ) func currentFixture(t *testing.T, mode string) (configread.Task, map[string]configread.Provider) { @@ -25,31 +23,27 @@ func currentFixture(t *testing.T, mode string) (configread.Task, map[string]conf if err := json.Unmarshal(raw, &task); err != nil { t.Fatal(err) } - raw, err = os.ReadFile("../../contracts/schema/examples/config-read-providers.json") - if err != nil { - t.Fatal(err) - } - var list struct { - Providers []configread.Provider `json:"providers"` - } - if err := json.Unmarshal(raw, &list); err != nil { - t.Fatal(err) - } - providers := make(map[string]configread.Provider, len(list.Providers)) - for _, p := range list.Providers { - providers[p.ProviderRef] = p + providers := map[string]configread.Provider{} + for _, code := range []string{"volcengine", "openai_compatible", "ali_bailian"} { + providers[code] = configread.Provider{ProviderRef: "catalog-" + code, Code: code, Name: "Example only", Endpoint: "https://example.invalid", WSEndpoint: "wss://example.invalid", Credential: "example-only-not-a-real-secret"} } + task = changeCurrentAgent(t, task, func(a map[string]any) { + a["asr"] = map[string]any{"provider_ref": "volcengine", "model": "example-asr", "params": map[string]any{"result_type": "full", "enable_nonstream": false, "enable_itn": false}} + if mode == "full_ai" { + a["llm"] = map[string]any{"provider_ref": "openai_compatible", "model": "example-chat", "params": map[string]any{"temperature": 0, "max_tokens": 256}} + a["tts"] = map[string]any{"provider_ref": "ali_bailian", "model": "qwen3-tts-flash", "voice": "Cherry", "params": map[string]any{"format": "pcm", "sample_rate": 16000, "rate": 1, "language": "Chinese"}} + a["prompt"].(map[string]any)["text"] = "Example only" + } + }) return task, providers } - func changeCurrentAgent(t *testing.T, task configread.Task, change func(map[string]any)) configread.Task { t.Helper() var body map[string]any if err := json.Unmarshal(task.Raw, &body); err != nil { t.Fatal(err) } - agent := body["agent"].(map[string]any) - change(agent) + change(body["agent"].(map[string]any)) raw, err := json.Marshal(body) if err != nil { t.Fatal(err) @@ -60,107 +54,48 @@ func changeCurrentAgent(t *testing.T, task configread.Task, change func(map[stri } return updated } - func TestBindCurrentFullAIUsesApprovedSDKFields(t *testing.T) { - task, providers := currentFixture(t, "full_ai") - bound, err := Bind(task, providers) + task, ps := currentFixture(t, "full_ai") + b, err := Bind(task, ps) if err != nil { t.Fatal(err) } - if bound.Mode != "full_ai" || bound.LLM == nil || bound.TTS == nil { - t.Fatalf("full mode needs all three providers: mode=%q LLM=%t TTS=%t", bound.Mode, bound.LLM != nil, bound.TTS != nil) + if b.LLM == nil || b.TTS == nil || b.ASR.Provider.Code != "volcengine" || b.LLM.Model != "example-chat" || b.TTS.Voice != "Cherry" || b.TTS.Protocol != TTSProtocolDashScopeTask || string(b.LLM.Params["temperature"]) != "0" { + t.Fatal("approved settings changed") } - if bound.ASR.Provider.Credential != providers["asr-example"].Credential || bound.ASR.Request.Format != doubaospeech.FormatPCMS16LE || bound.ASR.Request.SampleRate != 16000 || bound.ASR.Request.Channel != 1 || bound.ASR.Request.Bits != 16 || bound.ASR.Request.Language != doubaospeech.LanguageZhCN || bound.ASR.Request.ResultType != "full" || bound.ASR.Timeout != 5*time.Second { - t.Fatal("ASR approved input/interim/credential/timeout not bound to SDK request") - } - if bound.LLM.Provider.Credential != providers["llm-example"].Credential || bound.LLM.Model != "example-chat" || bound.LLM.Temperature == nil || *bound.LLM.Temperature != 0 || bound.LLM.MaxTokens == nil || *bound.LLM.MaxTokens != 256 || bound.LLM.Timeout != 5*time.Second { - t.Fatal("LLM model/explicit zero/limit/credential/timeout not bound") - } - if bound.TTS.Provider.Credential != providers["tts-example"].Credential || bound.TTS.Model != "qwen3-tts-flash" || bound.TTS.Voice != "Cherry" || bound.TTS.LanguageType != "Chinese" || bound.TTS.Speed != 1 || bound.TTS.SampleRate != 16000 || bound.TTS.Timeout != 5*time.Second { - t.Fatal("Bailian TTS model/voice/neutral speed/PCM16 output/credential/timeout not bound") - } - if len(bound.HangupKeywords) != 1 || bound.HangupKeywords[0].Name != "结束通话" || len(bound.HangupKeywords[0].Triggers) != 1 || bound.HangupKeywords[0].Triggers[0] != "不用了" || bound.HangupKeywords[0].ClosingRemark != "好的,祝您生活愉快。" || bound.Prompt != "Example only" || bound.Opening != "Example greeting" { - t.Fatal("immutable prompt and keyword behavior not bound") + if len(b.HangupKeywords) != 1 || b.HangupKeywords[0].ClosingRemark != "好的,祝您生活愉快。" || b.Prompt != "Example only" || b.Opening != "Example greeting" { + t.Fatal("conversation settings changed") } } - func TestBindCurrentASROnlyDoesNotBindOtherProviders(t *testing.T) { - task, providers := currentFixture(t, "asr_only") - delete(providers, "llm-example") - delete(providers, "tts-example") - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) - } - if bound.Mode != "asr_only" || bound.LLM != nil || bound.TTS != nil || bound.ASR.Request.ResultType != "single" { - t.Fatal("ASR-only mode must not inherit LLM/TTS configuration or interim results") + task, ps := currentFixture(t, "asr_only") + delete(ps, "openai_compatible") + delete(ps, "ali_bailian") + b, err := Bind(task, ps) + if err != nil || b.LLM != nil || b.TTS != nil { + t.Fatal("ASR-only inherited another provider", err) } } - -func TestBindCurrentRejectsSDKUnsupportedTTSWithoutChangingSchema(t *testing.T) { - for _, tc := range []struct { - name string - edit func(map[string]any) - }{ - {"pcma", func(tts map[string]any) { tts["format"].(map[string]any)["encoding"] = "pcma" }}, - {"speed-below", func(tts map[string]any) { tts["speed"] = 0.25 }}, - {"speed-above", func(tts map[string]any) { tts["speed"] = 3.0 }}, - {"speed-unrepresentable", func(tts map[string]any) { tts["speed"] = 1.005 }}, - {"sample-rate", func(tts map[string]any) { tts["format"].(map[string]any)["sample_rate_hz"] = 12345 }}, - } { - t.Run(tc.name, func(t *testing.T) { - task, providers := currentFixture(t, "full_ai") - task = changeCurrentAgent(t, task, func(agent map[string]any) { tc.edit(agent["tts"].(map[string]any)) }) - _, err := Bind(task, providers) - if err == nil || !strings.Contains(err.Error(), "TTS") || strings.Contains(err.Error(), providers["tts-example"].Credential) { - t.Fatalf("expected explicit non-secret TTS capability error, got %v", err) - } - }) - } -} - -func TestBindCurrentRejectsRemovedVolcTTSAdapter(t *testing.T) { - task, providers := currentFixture(t, "full_ai") - provider := providers["tts-example"] - provider.Code = "volcengine" - providers[provider.ProviderRef] = provider - if _, err := Bind(task, providers); err == nil || !strings.Contains(err.Error(), "provider") { - t.Fatalf("removed TTS adapter must fail closed, got %v", err) - } -} - func TestBindCurrentRejectsUnauthorizedProviderBeforeCall(t *testing.T) { - for _, tc := range []struct { - name string - mutate func(map[string]configread.Provider) - }{ - {"missing", func(ps map[string]configread.Provider) { delete(ps, "asr-example") }}, - {"identity-mismatch", func(ps map[string]configread.Provider) { - p := ps["asr-example"] - p.ProviderRef = "another-provider" - ps["asr-example"] = p - }}, - {"unsupported-provider", func(ps map[string]configread.Provider) { - p := ps["asr-example"] - p.Code = "unknown" - ps[p.ProviderRef] = p - }}, - {"invalid-endpoint", func(ps map[string]configread.Provider) { - p := ps["asr-example"] - p.Endpoint = "" - ps[p.ProviderRef] = p - }}, - {"missing-credential", func(ps map[string]configread.Provider) { - p := ps["asr-example"] - p.Credential = "" - ps[p.ProviderRef] = p - }}, - } { - t.Run(tc.name, func(t *testing.T) { - task, providers := currentFixture(t, "full_ai") - tc.mutate(providers) - if _, err := Bind(task, providers); err == nil { - t.Fatal("unavailable AI provider cannot authorize execution") + for _, kind := range []string{"missing", "code", "endpoint", "credential"} { + t.Run(kind, func(t *testing.T) { + task, ps := currentFixture(t, "full_ai") + p := ps["volcengine"] + switch kind { + case "missing": + delete(ps, "volcengine") + case "code": + p.Code = "unknown" + ps["volcengine"] = p + case "endpoint": + p.WSEndpoint = "" + ps["volcengine"] = p + case "credential": + p.Credential = "" + ps["volcengine"] = p + } + if _, err := Bind(task, ps); err == nil || strings.Contains(err.Error(), "example-only-not-a-real-secret") { + t.Fatal("missing/redacted provider validation", err) } }) } diff --git a/internal/ai/controls_test.go b/internal/ai/controls_test.go index 07530a4..93409fe 100644 --- a/internal/ai/controls_test.go +++ b/internal/ai/controls_test.go @@ -12,7 +12,7 @@ func TestBindCurrentPreservesPromptAndConversationControls(t *testing.T) { if err != nil { t.Fatal(err) } - if bound.PromptMaxBytes != 32768 || len(bound.AllowedVariables) != 0 { + if len(bound.AllowedVariables) != 0 { t.Fatal("approved prompt controls not preserved") } c := bound.Conversation @@ -21,14 +21,13 @@ func TestBindCurrentPreservesPromptAndConversationControls(t *testing.T) { } } -func TestBindCurrentRejectsPromptOverConfiguredByteLimit(t *testing.T) { +func TestBindDoesNotLimitPromptBytes(t *testing.T) { task, providers := currentFixture(t, "full_ai") - task = changeCurrentAgent(t, task, func(agent map[string]any) { - agent["prompt"].(map[string]any)["max_bytes"] = 3 - }) - _, err := Bind(task, providers) - if err == nil || !strings.Contains(err.Error(), "prompt") { - t.Fatalf("prompt limit was ignored: %v", err) + text := strings.Repeat("长", 40000) + task = changeCurrentAgent(t, task, func(agent map[string]any) { agent["prompt"].(map[string]any)["text"] = text }) + bound, err := Bind(task, providers) + if err != nil || bound.Prompt != text { + t.Fatalf("prompt was limited or changed: %v", err) } } diff --git a/internal/ai/dependency_test.go b/internal/ai/dependency_test.go new file mode 100644 index 0000000..022b579 --- /dev/null +++ b/internal/ai/dependency_test.go @@ -0,0 +1,24 @@ +package ai + +import ( + "os" + "reflect" + "strings" + "testing" + + doubaospeech "git.ipao.vip/rogee/doubao-speech-go" +) + +func TestSpeechSDKIsOurPinnedIndependentModule(t *testing.T) { + if got := reflect.TypeOf(doubaospeech.Client{}).PkgPath(); got != "git.ipao.vip/rogee/doubao-speech-go" { + t.Fatalf("speech SDK is not our maintained module: %s", got) + } + raw, err := os.ReadFile("../../go.mod") + if err != nil { + t.Fatal(err) + } + text := string(raw) + if !strings.Contains(text, "git.ipao.vip/rogee/doubao-speech-go v") || strings.Contains(text, "github.com/GizClaw/doubao-speech-go") || strings.Contains(text, "./third_party/doubao-speech-go") { + t.Fatal("speech SDK is not pinned to its independent repository") + } +} diff --git a/internal/ai/limits_test.go b/internal/ai/limits_test.go index e4754fa..23f21fc 100644 --- a/internal/ai/limits_test.go +++ b/internal/ai/limits_test.go @@ -18,9 +18,9 @@ func TestTTSRejectsExcessPendingAudioWithoutRetry(t *testing.T) { conversation["max_pending_audio_chunks"] = 1 }) endpoint, requests := mockBailianTTS(t, []byte{1, 0}, nil) - p := providers["tts-example"] - p.Endpoint = endpoint - providers[p.ProviderRef] = p + p := providers["ali_bailian"] + p.WSEndpoint = endpoint + providers[p.Code] = p bound, err := Bind(task, providers) if err != nil { t.Fatal(err) @@ -45,9 +45,9 @@ func TestReplyChunksUnicodeByApprovedSentenceLimit(t *testing.T) { }) texts := make(chan string, 3) endpoint, requests := mockBailianTTS(t, []byte{1, 0}, func(text string) { texts <- text }) - p := providers["tts-example"] - p.Endpoint = endpoint - providers[p.ProviderRef] = p + p := providers["ali_bailian"] + p.WSEndpoint = endpoint + providers[p.Code] = p bound, err := Bind(task, providers) if err != nil { t.Fatal(err) diff --git a/internal/ai/llm_params_test.go b/internal/ai/llm_params_test.go new file mode 100644 index 0000000..07a6f33 --- /dev/null +++ b/internal/ai/llm_params_test.go @@ -0,0 +1,48 @@ +package ai + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" +) + +func TestLLMWirePreservesOpaqueScalarParamsWithoutDefaults(t *testing.T) { + task, ps := currentFixture(t, "full_ai") + task = changeCurrentAgent(t, task, func(a map[string]any) { + a["llm"].(map[string]any)["params"] = map[string]any{"opaque": "vendor-value", "unknown_flag": false, "nullable": nil, "large_number": json.Number("9007199254740993"), "stream": false} + }) + b, e := Bind(task, ps) + if e != nil { + t.Fatal(e) + } + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + raw, e := io.ReadAll(r.Body) + if e != nil { + t.Error(e) + return + } + if !bytes.Contains(raw, []byte(`"large_number":9007199254740993`)) { + t.Error("large integer rewritten") + } + var body map[string]any + if json.Unmarshal(raw, &body) != nil { + t.Error("invalid request") + return + } + if len(body) != 7 || body["stream"] != false || body["model"] != "example-chat" || body["opaque"] != "vendor-value" || body["unknown_flag"] != false || body["nullable"] != nil { + t.Error("params dropped, rewritten or defaulted") + } + w.Header().Set("Content-Type", "application/json") + io.WriteString(w, `{"choices":[{"message":{"role":"assistant","content":"ok"}}]}`) + })) + defer server.Close() + b.LLM.Provider.Endpoint = server.URL + "/v1" + result, e := b.Complete(context.Background(), "Final user text") + if e != nil || result != "ok" { + t.Fatal("LLM passthrough failed", e) + } +} diff --git a/internal/ai/media_profile_test.go b/internal/ai/media_profile_test.go index 0fbd295..dd98bae 100644 --- a/internal/ai/media_profile_test.go +++ b/internal/ai/media_profile_test.go @@ -1,30 +1,15 @@ package ai -import ( - "strings" - "testing" -) +import "testing" -func TestBindRejectsSampleRatesUnsupportedByAgentMedia(t *testing.T) { - for _, tc := range []struct { - name, mode, endpoint string - }{ - {name: "ASR-only capture", mode: "asr_only", endpoint: "asr"}, - {name: "full-AI capture", mode: "full_ai", endpoint: "asr"}, - {name: "full-AI playback", mode: "full_ai", endpoint: "tts"}, - } { - t.Run(tc.name, func(t *testing.T) { - task, providers := currentFixture(t, tc.mode) - task = changeCurrentAgent(t, task, func(agent map[string]any) { - if tc.endpoint == "asr" { - agent["asr"].(map[string]any)["input"].(map[string]any)["sample_rate_hz"] = 24000 - } else { - agent["tts"].(map[string]any)["format"].(map[string]any)["sample_rate_hz"] = 24000 - } - }) - if _, err := Bind(task, providers); err == nil || !strings.Contains(err.Error(), "media") { - t.Fatalf("SDK-supported sample rate cannot silently mismatch 16 kHz Agent media: %v", err) - } +func TestVendorParameterValuesAreNotLocallyValidated(t *testing.T) { + for _, role := range []string{"asr", "tts"} { + task, ps := currentFixture(t, "full_ai") + task = changeCurrentAgent(t, task, func(a map[string]any) { + a[role].(map[string]any)["params"] = map[string]any{"sample_rate": 24000, "format": "vendor-format", "speed": 9.5} }) + if _, err := Bind(task, ps); err != nil { + t.Fatal("vendor parameter was locally rejected", err) + } } } diff --git a/internal/ai/opening_test.go b/internal/ai/opening_test.go index 5bb6b59..8e264a4 100644 --- a/internal/ai/opening_test.go +++ b/internal/ai/opening_test.go @@ -17,9 +17,9 @@ func TestCallOpeningUsesApprovedTTSOnce(t *testing.T) { task, providers := currentFixture(t, "full_ai") text := make(chan string, 1) endpoint, calls := mockBailianTTS(t, []byte{1, 0, 2, 0}, func(value string) { text <- value }) - p := providers["tts-example"] - p.Endpoint = endpoint - providers[p.ProviderRef] = p + p := providers["ali_bailian"] + p.WSEndpoint = endpoint + providers[p.Code] = p bound, err := Bind(task, providers) if err != nil { t.Fatal(err) @@ -45,9 +45,9 @@ func TestCallOpeningFailureIsVisibleAndNeverRetried(t *testing.T) { http.Error(w, "mock provider unavailable", http.StatusServiceUnavailable) })) defer server.Close() - p := providers["tts-example"] - p.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation" - providers[p.ProviderRef] = p + p := providers["ali_bailian"] + p.WSEndpoint = "ws" + strings.TrimPrefix(server.URL, "http") + providers[p.Code] = p bound, err := Bind(task, providers) if err != nil { t.Fatal(err) diff --git a/internal/ai/pipeline.go b/internal/ai/pipeline.go index ef83fc1..281a5bd 100644 --- a/internal/ai/pipeline.go +++ b/internal/ai/pipeline.go @@ -2,19 +2,17 @@ package ai import ( "context" + "encoding/json" "errors" "fmt" "log/slog" "net/url" "strings" "sync" - "time" - "github.com/GizClaw/doubao-speech-go" + "git.ipao.vip/rogee/doubao-speech-go" "github.com/openai/openai-go/v3" "github.com/openai/openai-go/v3/option" - "github.com/openai/openai-go/v3/packages/param" - "github.com/openai/openai-go/v3/shared" ) // Call owns the keyword action for exactly one call. A failed or @@ -139,26 +137,23 @@ func (b Binding) Recognize(ctx context.Context, pcm16 []byte) (string, error) { if len(pcm16) == 0 || len(pcm16)%2 != 0 { return "", errors.New("ASR requires nonempty signed 16-bit PCM") } - ctx, cancel := currentDeadline(ctx, b.ASR.Timeout) - defer cancel() if b.ASR.Provider.Code == "ali_bailian" { return recognizeBailian(ctx, b.ASR, pcm16) } u, err := url.Parse(b.ASR.Provider.Endpoint) - if err != nil { - return "", errors.New("ASR endpoint is invalid") - } - if u.Scheme == "https" { - u.Scheme = "wss" - } else if u.Scheme == "http" { - u.Scheme = "ws" + if err != nil || u.Host == "" || (u.Scheme != "wss" && u.Scheme != "ws") { + return "", errors.New("ASR requires a WebSocket endpoint") } client := doubaospeech.NewClient("", doubaospeech.WithAPIKey(b.ASR.Provider.Credential), doubaospeech.WithResourceID(doubaospeech.ResourceASRStreamV2), doubaospeech.WithWebSocketURL(u.String()), ) - request := b.ASR.Request + request := doubaospeech.ASRV2Config{ + Format: doubaospeech.FormatPCMS16LE, SampleRate: 16000, Channel: 1, Bits: 16, + Request: &doubaospeech.ASRV2RequestConfig{ModelName: b.ASR.Model}, + Parameters: b.ASR.Params, + } session, err := client.ASRV2.OpenStreamSession(ctx, &request) if err != nil { return "", err @@ -190,22 +185,19 @@ func (b Binding) Complete(ctx context.Context, finalUserText string) (string, er if strings.TrimSpace(finalUserText) == "" { return "", errors.New("LLM requires final user text") } - ctx, cancel := currentDeadline(ctx, b.LLM.Timeout) - defer cancel() client := openai.NewClient( option.WithAPIKey(b.LLM.Provider.Credential), option.WithBaseURL(strings.TrimRight(b.LLM.Provider.Endpoint, "/")), option.WithMaxRetries(0), ) - messages := []openai.ChatCompletionMessageParamUnion{openai.SystemMessage(b.Prompt), openai.UserMessage(finalUserText)} - params := openai.ChatCompletionNewParams{Model: shared.ChatModel(b.LLM.Model), Messages: messages} - if b.LLM.Temperature != nil { - params.Temperature = param.NewOpt(*b.LLM.Temperature) + body := cloneParams(b.LLM.Params) + body["model"], _ = json.Marshal(b.LLM.Model) + body["messages"], _ = json.Marshal([]map[string]string{{"role": "system", "content": b.Prompt}, {"role": "user", "content": finalUserText}}) + raw, err := json.Marshal(body) + if err != nil { + return "", errors.New("encode LLM request") } - if b.LLM.MaxTokens != nil { - params.MaxTokens = param.NewOpt(*b.LLM.MaxTokens) - } - result, err := client.Chat.Completions.New(ctx, params) + result, err := client.Chat.Completions.New(ctx, openai.ChatCompletionNewParams{}, option.WithRequestBody("application/json", raw)) if err != nil { return "", err } @@ -256,8 +248,6 @@ func (b Binding) synthesize(ctx context.Context, text string, alreadyPending int if limit > 0 && alreadyPending >= limit { return nil, 0, errors.New("TTS exceeded approved pending audio chunk limit") } - ctx, cancel := currentDeadline(ctx, b.TTS.Timeout) - defer cancel() // Bailian produces one complete response for this approved sentence; a // missing or failed response never counts as queued audio. var audio []byte @@ -265,8 +255,6 @@ func (b Binding) synthesize(ctx context.Context, text string, alreadyPending int switch b.TTS.Protocol { case TTSProtocolDashScopeTask: audio, err = synthesizeBailianTaskTTS(ctx, *b.TTS, text) - case TTSProtocolDashScopeHTTP: - audio, err = synthesizeBailianTTS(ctx, *b.TTS, text) default: return nil, 0, errors.New("TTS protocol is missing or unsupported") } @@ -278,11 +266,12 @@ func (b Binding) synthesize(ctx context.Context, text string, alreadyPending int return audio, 1, nil } -func currentDeadline(ctx context.Context, timeout time.Duration) (context.Context, context.CancelFunc) { - if timeout > 0 { - return context.WithTimeout(ctx, timeout) +func cloneParams(params map[string]json.RawMessage) map[string]json.RawMessage { + out := make(map[string]json.RawMessage, len(params)) + for key, value := range params { + out[key] = append(json.RawMessage(nil), value...) } - return context.WithCancel(ctx) + return out } type finalASRAccumulator struct { diff --git a/internal/ai/pipeline_test.go b/internal/ai/pipeline_test.go index a599490..22729e8 100644 --- a/internal/ai/pipeline_test.go +++ b/internal/ai/pipeline_test.go @@ -6,63 +6,10 @@ import ( "fmt" "net/http" "net/http/httptest" - "os/exec" "strings" "testing" - - "git.ipao.vip/rogee/go-sip/internal/media" ) -func TestTTSPassesApprovedParametersToSDK(t *testing.T) { - if _, err := exec.LookPath("ffmpeg"); err != nil { - t.Skip("Bailian audio conversion requires ffmpeg") - } - want := []byte{1, 0, 2, 0} - wav, _, err := media.EncodeMonoWAV(want, 1024) - if err != nil { - t.Fatal(err) - } - task, providers := currentFixture(t, "full_ai") - var captured map[string]any - var credential string - postCount, getCount := 0, 0 - var server *httptest.Server - server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - switch r.URL.Path { - case "/api/v1/services/aigc/multimodal-generation/generation": - postCount++ - credential = r.Header.Get("Authorization") - if r.Method != http.MethodPost || json.NewDecoder(r.Body).Decode(&captured) != nil { - http.Error(w, "invalid generation request", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _, _ = fmt.Fprintf(w, `{"output":{"audio":{"url":%q}}}`, server.URL+"/audio?test-token=redacted") - case "/audio": - getCount++ - _, _ = w.Write(wav) - default: - http.Error(w, "unexpected endpoint", http.StatusNotFound) - } - })) - defer server.Close() - p := providers["tts-example"] - p.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation" - providers[p.ProviderRef] = p - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) - } - audio, err := bound.Synthesize(context.Background(), "批准的回复") - if err != nil || string(audio) != string(want) { - t.Fatalf("Bailian TTS response: length=%d err=%v", len(audio), err) - } - input, ok := captured["input"].(map[string]any) - if !ok || captured["model"] != "qwen3-tts-flash" || input["voice"] != "Cherry" || input["language_type"] != "Chinese" || input["text"] != "批准的回复" || credential != "Bearer "+p.Credential || postCount != 1 || getCount != 1 { - t.Fatal("approved Bailian model/voice/text/credential or single-request bound was lost") - } -} - func TestLLMPassesExplicitZeroAndDoesNotRetry(t *testing.T) { task, providers := currentFixture(t, "full_ai") calls := 0 @@ -83,9 +30,9 @@ func TestLLMPassesExplicitZeroAndDoesNotRetry(t *testing.T) { _, _ = fmt.Fprintln(w, `{"id":"mock","object":"chat.completion","created":0,"model":"example-chat","choices":[{"index":0,"message":{"role":"assistant","content":"收到"},"finish_reason":"stop"}]}`) })) defer server.Close() - p := providers["llm-example"] + p := providers["openai_compatible"] p.Endpoint = server.URL + "/v1" - providers[p.ProviderRef] = p + providers[p.Code] = p bound, err := Bind(task, providers) if err != nil { t.Fatal(err) @@ -122,9 +69,9 @@ func TestLLMFailureAndEmptyChoicesAreNotRetried(t *testing.T) { _, _ = fmt.Fprintln(w, tc.body) })) defer server.Close() - provider := providers["llm-example"] + provider := providers["openai_compatible"] provider.Endpoint = server.URL + "/v1" - providers[provider.ProviderRef] = provider + providers[provider.Code] = provider bound, err := Bind(task, providers) if err != nil { t.Fatal(err) diff --git a/internal/ai/protocol_capabilities_test.go b/internal/ai/protocol_capabilities_test.go index 77ceb59..dc597ff 100644 --- a/internal/ai/protocol_capabilities_test.go +++ b/internal/ai/protocol_capabilities_test.go @@ -1,73 +1,21 @@ package ai -import ( - "git.ipao.vip/rogee/go-sip/internal/configread" - "testing" -) +import "testing" -func TestProtocolCapabilitiesRejectUnexpressibleSettingsWithoutFallback(t *testing.T) { - for _, tc := range []struct { - name, protocol string - change func(map[string]any) - connection func(map[string]configread.Provider) - }{ - {name: "HTTP-speed", protocol: TTSProtocolDashScopeHTTP, change: func(a map[string]any) { a["tts"].(map[string]any)["speed"] = 1.5 }}, - {name: "HTTP-missing-endpoint", protocol: TTSProtocolDashScopeHTTP, connection: func(p map[string]configread.Provider) { x := p["tts-example"]; x.Endpoint = ""; p[x.ProviderRef] = x }}, - {name: "WS-missing-endpoint", protocol: TTSProtocolDashScopeTask, connection: func(p map[string]configread.Provider) { x := p["tts-example"]; x.WSEndpoint = ""; p[x.ProviderRef] = x }}, - {name: "WS-speed", protocol: TTSProtocolDashScopeTask, change: func(a map[string]any) { a["tts"].(map[string]any)["speed"] = 2.5 }}, - {name: "WS-language", protocol: TTSProtocolDashScopeTask, change: func(a map[string]any) { a["tts"].(map[string]any)["language_type"] = "not a language" }}, - {name: "WS-unspecified-language", protocol: TTSProtocolDashScopeTask, change: func(a map[string]any) { a["tts"].(map[string]any)["language_type"] = "und-US" }}, - {name: "ASR-interim", protocol: TTSProtocolDashScopeTask, change: func(a map[string]any) { a["asr"].(map[string]any)["interim"] = false }}, - {name: "missing-ASR-model", protocol: TTSProtocolDashScopeTask, change: func(a map[string]any) { delete(a["asr"].(map[string]any), "model") }}, - {name: "ASR-rate", protocol: TTSProtocolDashScopeTask, change: func(a map[string]any) { a["asr"].(map[string]any)["input"].(map[string]any)["sample_rate_hz"] = 48000 }}, - {name: "ASR-language", protocol: TTSProtocolDashScopeTask, change: func(a map[string]any) { a["asr"].(map[string]any)["language"] = "bad language" }}, - {name: "unknown-parameter", protocol: TTSProtocolDashScopeTask, change: func(a map[string]any) { a["tts"].(map[string]any)["extra_parameter"] = true }}, - } { - t.Run(tc.name, func(t *testing.T) { - task, providers := protocolTask(t, tc.protocol) - if tc.change != nil { - task = changeCurrentAgent(t, task, tc.change) - } - if tc.connection != nil { - tc.connection(providers) - } - if _, err := Bind(task, providers); err == nil { - t.Fatal("unsupported request silently altered or accepted") - } - }) - } -} - -func TestKnownModelNamesDoNotSelectTheProtocol(t *testing.T) { - for _, tc := range []struct{ protocol, model string }{{TTSProtocolDashScopeHTTP, "cosyvoice-v3-flash"}, {TTSProtocolDashScopeTask, "qwen3-tts-flash"}} { - t.Run(tc.protocol, func(t *testing.T) { - task, providers := protocolTask(t, tc.protocol) - task = changeCurrentAgent(t, task, func(a map[string]any) { a["tts"].(map[string]any)["model"] = tc.model }) - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) - } - if bound.TTS.Protocol != tc.protocol || bound.TTS.Model != tc.model { - t.Fatal("explicit selection replaced by a model-name mapping") - } - }) - } -} - -func TestLanguageHintsKeepExplicitLanguageAndNeverGuessFromRegion(t *testing.T) { - for _, tc := range []struct{ value, want string }{{"Chinese", "zh"}, {"English", "en"}, {"Japanese", "ja"}, {"Korean", "ko"}, {"German", "de"}, {"French", "fr"}, {"Russian", "ru"}, {"Italian", "it"}, {"Spanish", "es"}, {"Portuguese", "pt"}, {"en-US", "en"}, {"zh-CN", "zh"}, {"fr-FR", "fr"}} { - got, err := bailianLanguageHint(tc.value) - if err != nil || got != tc.want { - t.Fatalf("hint %s = %s %v", tc.value, got, err) +func TestTTSAlwaysUsesExplicitlyApprovedWebSocketProtocol(t *testing.T) { + for _, model := range []string{"qwen3-tts-flash", "cosyvoice-v3-flash", "unknown-model"} { + task, ps := currentFixture(t, "full_ai") + task = changeCurrentAgent(t, task, func(a map[string]any) { a["tts"].(map[string]any)["model"] = model }) + b, err := Bind(task, ps) + if err != nil || b.TTS.Protocol != TTSProtocolDashScopeTask { + t.Fatal("model selected the protocol", err) } } - for _, value := range []string{"", "und", "und-US", "not a language"} { - if _, err := bailianLanguageHint(value); err == nil { - t.Fatalf("language inferred from %q", value) - } - } - hints, err := bailianTTSLanguageHints("Auto") - if err != nil || hints != nil { - t.Fatal("explicit Auto must omit the hint") + task, ps := currentFixture(t, "full_ai") + p := ps["ali_bailian"] + p.WSEndpoint = "" + ps[p.Code] = p + if _, err := Bind(task, ps); err == nil { + t.Fatal("missing WS endpoint silently fell back to HTTP") } } diff --git a/internal/ai/protocol_models_test.go b/internal/ai/protocol_models_test.go index dc94767..c2ae663 100644 --- a/internal/ai/protocol_models_test.go +++ b/internal/ai/protocol_models_test.go @@ -7,73 +7,22 @@ import ( func protocolTask(t *testing.T, protocol string) (configread.Task, map[string]configread.Provider) { t.Helper() - task, providers := currentFixture(t, "full_ai") + task, ps := currentFixture(t, "full_ai") task = changeCurrentAgent(t, task, func(a map[string]any) { - asr := a["asr"].(map[string]any) - delete(asr, "interim") - asr["model"] = "unlisted-asr-model-2029" - asr["language"] = "en-US" - asr["input"].(map[string]any)["sample_rate_hz"] = 16000 - llm := a["llm"].(map[string]any) - llm["model"] = "unlisted-chat-model-2029" - tts := a["tts"].(map[string]any) - tts["protocol"] = protocol - tts["model"] = "unlisted-tts-model-2029" - tts["voice"] = "unlisted-voice" - tts["language_type"] = "English" - tts["speed"] = 1.0 - if protocol == "dashscope_task_websocket" { - tts["speed"] = 1.75 - } + a["asr"].(map[string]any)["provider_ref"] = "ali_bailian" + a["asr"].(map[string]any)["model"] = "unlisted-asr-model-2029" + a["asr"].(map[string]any)["params"] = map[string]any{"sample_rate": 16000, "language": "English", "format": "pcm"} + a["llm"].(map[string]any)["model"] = "unlisted-chat-model-2029" + a["tts"].(map[string]any)["model"] = "unlisted-tts-model-2029" + a["tts"].(map[string]any)["voice"] = "unlisted-voice" + a["tts"].(map[string]any)["params"] = map[string]any{"sample_rate": 16000, "format": "pcm", "rate": 1.75, "language": "English"} }) - asr := providers["asr-example"] - asr.Code = "ali_bailian" - asr.WSEndpoint = "wss://speech.example.invalid/inference" - providers[asr.ProviderRef] = asr - tts := providers["tts-example"] - tts.WSEndpoint = asr.WSEndpoint - providers[tts.ProviderRef] = tts - return task, providers + return task, ps } - -func TestModelsAndVoicesAreDataWithinExplicitTTSProtocol(t *testing.T) { - for _, protocol := range []string{"dashscope_tts_http", "dashscope_task_websocket"} { - t.Run(protocol, func(t *testing.T) { - task, providers := protocolTask(t, protocol) - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) - } - if bound.ASR.Model != "unlisted-asr-model-2029" || bound.ASR.Request.SampleRate != 16000 || bound.ASR.Language != "en-US" || bound.LLM.Model != "unlisted-chat-model-2029" || bound.TTS.Model != "unlisted-tts-model-2029" || bound.TTS.Voice != "unlisted-voice" || bound.TTS.LanguageType != "English" { - t.Fatal("task-selected model/voice/language changed") - } - expected := providers["tts-example"].Endpoint - if protocol == "dashscope_task_websocket" { - expected = providers["tts-example"].WSEndpoint - } - if bound.TTS.Provider.Endpoint != expected { - t.Fatal("connection selected from model name instead of explicit protocol") - } - }) - } -} - -func TestMissingOrUnsupportedProtocolNeverFallsBack(t *testing.T) { - for _, protocol := range []string{"", "unknown-protocol"} { - t.Run(protocol, func(t *testing.T) { - task, providers := protocolTask(t, "dashscope_task_websocket") - task = changeCurrentAgent(t, task, func(a map[string]any) { - tts := a["tts"].(map[string]any) - if protocol == "" { - delete(tts, "protocol") - } else { - tts["protocol"] = protocol - } - tts["model"] = "cosyvoice-v3-flash" - }) - if _, err := Bind(task, providers); err == nil { - t.Fatal("protocol guessed from a familiar model") - } - }) +func TestProtocolModelsAndVoicesAreTaskSelectedNotHardcoded(t *testing.T) { + task, ps := protocolTask(t, TTSProtocolDashScopeTask) + b, err := Bind(task, ps) + if err != nil || b.ASR.Model != "unlisted-asr-model-2029" || b.LLM.Model != "unlisted-chat-model-2029" || b.TTS.Model != "unlisted-tts-model-2029" || b.TTS.Voice != "unlisted-voice" { + t.Fatal("task-selected model rejected", err) } } diff --git a/internal/ai/protocol_runtime_test.go b/internal/ai/protocol_runtime_test.go index d28d19c..afacd52 100644 --- a/internal/ai/protocol_runtime_test.go +++ b/internal/ai/protocol_runtime_test.go @@ -4,209 +4,107 @@ import ( "bytes" "context" "encoding/json" - "fmt" + "github.com/coder/websocket" "net/http" "net/http/httptest" "strings" - "sync/atomic" "testing" "time" - - "git.ipao.vip/rogee/go-sip/internal/media" - doubaospeech "github.com/GizClaw/doubao-speech-go" - "github.com/coder/websocket" ) -func TestTaskTTSProtocolForwardsOpaqueModelVoiceSpeedAndLanguage(t *testing.T) { - for _, language := range []string{"English", "fr-FR", "Auto"} { - t.Run(language, func(t *testing.T) { - task, providers := protocolTask(t, TTSProtocolDashScopeTask) - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) - } - bound.TTS.LanguageType = language - var calls atomic.Int32 - expected := []byte{1, 0, 2, 0} - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - calls.Add(1) - c, err := websocket.Accept(w, r, nil) - if err != nil { - t.Error(err) - return - } - defer c.CloseNow() - ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second) - defer cancel() - cmd := mockRead(t, c, ctx) - p := cmd.Payload["parameters"].(map[string]any) - if cmd.Payload["model"] != "unlisted-tts-model-2029" || p["voice"] != "unlisted-voice" || p["rate"] != 1.75 || p["sample_rate"] != float64(16000) || p["format"] != "pcm" { - t.Errorf("approved parameters changed: %+v", cmd.Payload) - } - if language == "Auto" { - if _, ok := p["language_hints"]; ok { - t.Error("explicit Auto must not insert a language hint") - } - } else { - expectedHint := "en" - if language == "fr-FR" { - expectedHint = "fr" - } - hints, ok := p["language_hints"].([]any) - if !ok || len(hints) != 1 || hints[0] != expectedHint { - t.Error("language was replaced with a model-specific default") - } - } - mockEvent(c, ctx, cmd.Header.TaskID, "task-started", nil) - text := mockRead(t, c, ctx) - finish := mockRead(t, c, ctx) - if text.Payload["input"].(map[string]any)["text"] != "Protocol fixture." || finish.Header.Action != "finish-task" { - t.Error("TTS text/lifecycle changed") - } - c.Write(ctx, websocket.MessageBinary, expected) - mockEvent(c, ctx, cmd.Header.TaskID, "task-finished", nil) - })) - defer server.Close() - bound.TTS.Provider.WSEndpoint = "ws" + strings.TrimPrefix(server.URL, "http") - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - audio, err := bound.Synthesize(ctx, "Protocol fixture.") - if err != nil || !bytes.Equal(audio, expected) || calls.Load() != 1 { - t.Fatalf("task TTS: audio=%d err=%v requests=%d", len(audio), err, calls.Load()) - } - }) - } -} - -func TestTaskASRProtocolForwardsOpaqueModelAndSelectedSampleRate(t *testing.T) { - for _, rate := range []int{8000, 16000} { - t.Run(fmt.Sprint(rate), func(t *testing.T) { - task, providers := protocolTask(t, TTSProtocolDashScopeTask) - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) - } - bound.ASR.Request.SampleRate = doubaospeech.SampleRate(rate) - pcm := make([]byte, 32000) - for i := range pcm { - pcm[i] = byte(i % 251) - } - var calls atomic.Int32 - captured := make(chan []byte, 1) - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - var received []byte - defer func() { captured <- received }() - calls.Add(1) - c, err := websocket.Accept(w, r, nil) - if err != nil { - t.Error(err) - return - } - defer c.CloseNow() - ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second) - defer cancel() - cmd := mockRead(t, c, ctx) - p := cmd.Payload["parameters"].(map[string]any) - if cmd.Payload["model"] != "unlisted-asr-model-2029" || p["sample_rate"] != float64(rate) || p["language_hints"].([]any)[0] != "en" { - t.Error("ASR task model/rate/language changed") - } - mockEvent(c, ctx, cmd.Header.TaskID, "task-started", nil) - for { - kind, raw, err := c.Read(ctx) - if err != nil { - t.Error(err) - return - } - if kind == websocket.MessageBinary { - received = append(received, raw...) - continue - } - var finish mockCommand - json.Unmarshal(raw, &finish) - if finish.Header.Action != "finish-task" { - t.Error("ASR not finished explicitly") - } - break - } - mockEvent(c, ctx, cmd.Header.TaskID, "result-generated", map[string]any{"output": map[string]any{"sentence": map[string]any{"text": "Final user text.", "sentence_end": true, "begin_time": 0, "end_time": 1000}}}) - mockEvent(c, ctx, cmd.Header.TaskID, "task-finished", nil) - })) - defer server.Close() - bound.ASR.Provider.Endpoint = "ws" + strings.TrimPrefix(server.URL, "http") - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - text, err := bound.Recognize(ctx, pcm) - received := <-captured - if err != nil || text != "Final user text." || len(received) != rate*2 || calls.Load() != 1 { - t.Fatalf("ASR: text=%q err=%v sent=%d calls=%d", text, err, len(received), calls.Load()) - } - if rate == 16000 && !bytes.Equal(received, pcm) { - t.Fatal("16k PCM was transformed despite the selected 16k protocol rate") - } - }) - } -} - -func TestHTTPProtocolForwardsOpaqueModelVoiceAndLanguage(t *testing.T) { - wav, _, err := media.EncodeMonoWAV([]byte{0, 0, 2, 0, 3, 0, 4, 0}, 1024) +func TestTaskTTSProtocolForwardsOpaqueModelVoiceAndParams(t *testing.T) { + task, ps := protocolTask(t, TTSProtocolDashScopeTask) + b, err := Bind(task, ps) if err != nil { t.Fatal(err) } - var generation, downloads atomic.Int32 - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/audio" { - downloads.Add(1) - w.Header().Set("Content-Type", "audio/wav") - w.Write(wav) + b.TTS.Params = map[string]json.RawMessage{"opaque": json.RawMessage(`"vendor-only"`), "rate": json.RawMessage("9.5"), "unknown_flag": json.RawMessage("false"), "nullable": json.RawMessage("null"), "large_number": json.RawMessage("9007199254740993")} + expected := []byte{1, 0, 2, 0} + s := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + c, e := websocket.Accept(w, r, nil) + if e != nil { + t.Error(e) return } - generation.Add(1) - var req struct { - Model string `json:"model"` - Input struct { - Text, Voice string - LanguageType string `json:"language_type"` - } `json:"input"` + defer c.CloseNow() + ctx, cancel := context.WithTimeout(r.Context(), time.Second) + defer cancel() + kind, raw, err := c.Read(ctx) + if err != nil || kind != websocket.MessageText { + t.Error("missing TTS start payload") + return } - if json.NewDecoder(r.Body).Decode(&req) != nil || req.Model != "unlisted-tts-model-2029" || req.Input.Voice != "unlisted-voice" || req.Input.LanguageType != "English" || req.Input.Text != "Protocol fixture." { - t.Error("HTTP task parameters changed") + if !bytes.Contains(raw, []byte(`"large_number":9007199254740993`)) { + t.Error("large TTS integer changed") } - w.Header().Set("Content-Type", "application/json") - fmt.Fprintf(w, `{"output":{"audio":{"url":%q}}}`, "http://"+r.Host+"/audio") + var cmd mockCommand + if err := json.Unmarshal(raw, &cmd); err != nil { + t.Error(err) + return + } + p := cmd.Payload["parameters"].(map[string]any) + if cmd.Payload["model"] != "unlisted-tts-model-2029" || p["voice"] != "unlisted-voice" || p["opaque"] != "vendor-only" || p["rate"] != 9.5 || p["unknown_flag"] != false || p["nullable"] != nil || len(p) != 6 { + t.Error("opaque parameters changed") + } + mockEvent(c, ctx, cmd.Header.TaskID, "task-started", nil) + mockRead(t, c, ctx) + mockRead(t, c, ctx) + c.Write(ctx, websocket.MessageBinary, expected) + mockEvent(c, ctx, cmd.Header.TaskID, "task-finished", nil) })) - defer server.Close() - task, providers := protocolTask(t, TTSProtocolDashScopeHTTP) - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) - } - bound.TTS.Provider.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation" - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer s.Close() + b.TTS.Provider.WSEndpoint = "ws" + strings.TrimPrefix(s.URL, "http") + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() - audio, err := bound.Synthesize(ctx, "Protocol fixture.") - if err != nil || len(audio) == 0 || generation.Load() != 1 || downloads.Load() != 1 { - t.Fatalf("HTTP TTS: audio=%d err=%v requests=%d downloads=%d", len(audio), err, generation.Load(), downloads.Load()) + audio, e := b.Synthesize(ctx, "Protocol fixture.") + if e != nil || !bytes.Equal(audio, expected) { + t.Fatal("TTS protocol failure", e) } } -func TestExplicitHTTPProtocolNeverGuessesWebSocketFromFamiliarModel(t *testing.T) { - var httpCalls, wsCalls atomic.Int32 - ws := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { wsCalls.Add(1); w.WriteHeader(500) })) - defer ws.Close() - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { httpCalls.Add(1); w.WriteHeader(400) })) - defer server.Close() - task, providers := protocolTask(t, TTSProtocolDashScopeHTTP) - bound, err := Bind(task, providers) - if err != nil { - t.Fatal(err) +func TestTaskASRForwardsParamsWithoutRewritingAudio(t *testing.T) { + task, ps := protocolTask(t, TTSProtocolDashScopeTask) + b, e := Bind(task, ps) + if e != nil { + t.Fatal(e) } - bound.TTS.Model = "cosyvoice-v3-flash" - bound.TTS.Provider.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation" - bound.TTS.Provider.WSEndpoint = "ws" + strings.TrimPrefix(ws.URL, "http") - ctx, cancel := context.WithTimeout(context.Background(), time.Second) + b.ASR.Params = map[string]json.RawMessage{"sample_rate": json.RawMessage("24000"), "opaque": json.RawMessage(`"unchanged"`), "nullable": json.RawMessage("null")} + pcm := []byte{1, 0, 2, 0} + captured := make(chan []byte, 1) + s := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + c, e := websocket.Accept(w, r, nil) + if e != nil { + t.Error(e) + return + } + defer c.CloseNow() + ctx, cancel := context.WithTimeout(r.Context(), time.Second) + defer cancel() + cmd := mockRead(t, c, ctx) + p := cmd.Payload["parameters"].(map[string]any) + if len(p) != 3 || p["sample_rate"] != float64(24000) || p["opaque"] != "unchanged" || p["nullable"] != nil { + t.Error("ASR params rewritten") + } + mockEvent(c, ctx, cmd.Header.TaskID, "task-started", nil) + _, body, e := c.Read(ctx) + if e != nil { + t.Error(e) + return + } + captured <- body + mockRead(t, c, ctx) + mockEvent(c, ctx, cmd.Header.TaskID, "result-generated", map[string]any{"output": map[string]any{"sentence": map[string]any{"text": "Final.", "sentence_end": true, "begin_time": 0, "end_time": 1}}}) + mockEvent(c, ctx, cmd.Header.TaskID, "task-finished", nil) + })) + defer s.Close() + b.ASR.Provider.Endpoint = "ws" + strings.TrimPrefix(s.URL, "http") + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() - audio, err := bound.Synthesize(ctx, "No fallback.") - if err == nil || len(audio) != 0 || httpCalls.Load() != 1 || wsCalls.Load() != 0 { - t.Fatalf("model/protocol fallback: audio=%d err=%v HTTP=%d WS=%d", len(audio), err, httpCalls.Load(), wsCalls.Load()) + text, e := b.Recognize(ctx, pcm) + if e != nil || text != "Final." { + t.Fatal(e) + } + if !bytes.Equal(<-captured, pcm) { + t.Fatal("input media rewritten from opaque params") } } diff --git a/internal/ai/protocols.go b/internal/ai/protocols.go index d746474..91891e5 100644 --- a/internal/ai/protocols.go +++ b/internal/ai/protocols.go @@ -1,40 +1,4 @@ package ai -import ( - "errors" - "golang.org/x/text/language" -) - -const ( - TTSProtocolDashScopeHTTP = "dashscope_tts_http" - TTSProtocolDashScopeTask = "dashscope_task_websocket" -) - -// Task language labels/locales express the language; DashScope task parameters -// use ISO language hints. Do not infer an unspecified language from a region. -func bailianLanguageHint(value string) (string, error) { - names := map[string]string{"Chinese": "zh", "English": "en", "Japanese": "ja", "Korean": "ko", "German": "de", "French": "fr", "Russian": "ru", "Italian": "it", "Spanish": "es", "Portuguese": "pt"} - if code, ok := names[value]; ok { - return code, nil - } - tag, err := language.Parse(value) - if err != nil { - return "", errors.New("language cannot be expressed by the selected task protocol") - } - base, _, _ := tag.Raw() - if base.String() == "und" { - return "", errors.New("task language must be explicit") - } - return base.String(), nil -} - -func bailianTTSLanguageHints(value string) ([]string, error) { - if value == "Auto" { - return nil, nil - } // Explicit user selection, not a missing-value default. - hint, err := bailianLanguageHint(value) - if err != nil { - return nil, err - } - return []string{hint}, nil -} +// The current task contract uses task-based WebSocket TTS only. +const TTSProtocolDashScopeTask = "dashscope_task" diff --git a/internal/ai/provider_endpoints_test.go b/internal/ai/provider_endpoints_test.go new file mode 100644 index 0000000..b97346b --- /dev/null +++ b/internal/ai/provider_endpoints_test.go @@ -0,0 +1,55 @@ +package ai + +import "testing" + +func TestBindUsesWSEndpointForSpeechAndAPIEndpointForLLM(t *testing.T) { + for _, asrProvider := range []string{"volcengine", "ali_bailian"} { + t.Run(asrProvider, func(t *testing.T) { + task, providers := currentFixture(t, "full_ai") + if asrProvider == "ali_bailian" { + task = changeCurrentAgent(t, task, func(agent map[string]any) { + agent["asr"].(map[string]any)["provider_ref"] = "ali_bailian" + }) + } + for code, provider := range providers { + provider.Endpoint = "https://http-" + code + ".example.invalid/v1" + provider.WSEndpoint = "wss://stream-" + code + ".example.invalid/v1" + providers[code] = provider + } + binding, err := Bind(task, providers) + if err != nil { + t.Fatal(err) + } + if binding.ASR.Provider.Endpoint != providers[asrProvider].WSEndpoint || binding.TTS.Provider.Endpoint != providers["ali_bailian"].WSEndpoint || binding.LLM.Provider.Endpoint != providers["openai_compatible"].Endpoint { + t.Fatal("AI role selected the wrong connection endpoint") + } + if providers["ali_bailian"].Endpoint != "https://http-ali_bailian.example.invalid/v1" { + t.Fatal("binding mutated the reusable provider catalog") + } + }) + } +} + +func TestBindNeverFallsBackBetweenAPIAndWSEndpoints(t *testing.T) { + for _, role := range []string{"asr", "tts", "llm"} { + t.Run(role, func(t *testing.T) { + task, providers := currentFixture(t, "full_ai") + code := "volcengine" + if role == "tts" { + code = "ali_bailian" + } else if role == "llm" { + code = "openai_compatible" + } + provider := providers[code] + if role == "llm" { + provider.Endpoint = "" + } else { + provider.WSEndpoint = "" + } + providers[code] = provider + if _, err := Bind(task, providers); err == nil { + t.Fatal("missing endpoint was hidden by a fallback") + } + }) + } +} diff --git a/internal/ai/real_provider_integration_test.go b/internal/ai/real_provider_integration_test.go index 39a891a..78f4482 100644 --- a/internal/ai/real_provider_integration_test.go +++ b/internal/ai/real_provider_integration_test.go @@ -3,6 +3,7 @@ package ai import ( "context" "crypto/sha256" + "encoding/json" "net/url" "os" "strings" @@ -24,14 +25,12 @@ func TestRealBailianDiagnosticOnce(t *testing.T) { if err != nil || parsed.Scheme != "https" || parsed.Host == "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" || credential == "" { t.Fatal("real AI diagnostic requires a private API key and plain HTTPS endpoint") } - maxTokens := int64(16) - temperature := 0.0 bound := Binding{ Mode: "full_ai", Prompt: "This is a harmless connectivity test. Respond briefly.", LLM: &LLMConfig{ Provider: configread.Provider{Endpoint: endpoint, Credential: credential}, - Model: "qwen-plus", Temperature: &temperature, MaxTokens: &maxTokens, Timeout: 30 * time.Second, + Model: "qwen-plus", Params: map[string]json.RawMessage{"temperature": json.RawMessage("0"), "max_tokens": json.RawMessage("16")}, }, } ctx, cancel := context.WithTimeout(context.Background(), 35*time.Second) diff --git a/internal/ai/schema_params_test.go b/internal/ai/schema_params_test.go new file mode 100644 index 0000000..b9bd15e --- /dev/null +++ b/internal/ai/schema_params_test.go @@ -0,0 +1,40 @@ +package ai + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestNewSchemaBindsCodeReferenceAndOpaqueParams(t *testing.T) { + task, providers := currentFixture(t, "full_ai") + task = changeCurrentAgent(t, task, func(a map[string]any) { + a["asr"].(map[string]any)["params"] = map[string]any{"unknown_flag": false, "nullable": nil, "opaque": "vendor-value"} + a["llm"].(map[string]any)["params"] = map[string]any{"temperature": 0, "unknown_limit": json.Number("9007199254740993")} + a["tts"].(map[string]any)["params"] = map[string]any{"rate": 9.5, "language": "vendor-only"} + a["prompt"].(map[string]any)["text"] = strings.Repeat("长", 40000) + }) + b, err := Bind(task, providers) + if err != nil { + t.Fatal(err) + } + if b.ASR.Provider.Code != task.Agent.ASR.ProviderRef || b.TTS.Protocol != TTSProtocolDashScopeTask { + t.Fatal("provider code or fixed WebSocket protocol lost") + } + if string(b.ASR.Params["unknown_flag"]) != "false" || string(b.ASR.Params["nullable"]) != "null" || string(b.LLM.Params["unknown_limit"]) != "9007199254740993" || string(b.TTS.Params["rate"]) != "9.5" { + t.Fatal("opaque parameter value changed") + } + if len(b.Prompt) != 120000 { + t.Fatal("prompt was limited or truncated") + } +} + +func TestParamsRejectOnlyMalformedShapeAndReservedProtocolFields(t *testing.T) { + for _, value := range []any{nil, []any{}, map[string]any{"nested": map[string]any{}}, map[string]any{"messages": "override"}} { + task, providers := currentFixture(t, "full_ai") + task = changeCurrentAgent(t, task, func(a map[string]any) { a["llm"].(map[string]any)["params"] = value }) + if _, err := Bind(task, providers); err == nil { + t.Fatal("invalid parameter shape or protocol identity override accepted") + } + } +} diff --git a/internal/configread/snapshots.go b/internal/configread/snapshots.go index 388cb1a..663f573 100644 --- a/internal/configread/snapshots.go +++ b/internal/configread/snapshots.go @@ -47,13 +47,13 @@ type providerResponse struct { type Agent struct { Mode string `json:"mode"` ASR struct { - ProviderRef string `json:"provider_id"` + ProviderRef string `json:"provider_ref"` } `json:"asr"` LLM struct { - ProviderRef string `json:"provider_id"` + ProviderRef string `json:"provider_ref"` } `json:"llm"` TTS struct { - ProviderRef string `json:"provider_id"` + ProviderRef string `json:"provider_ref"` } `json:"tts"` Raw json.RawMessage `json:"-"` } @@ -69,6 +69,11 @@ func (a *Agent) UnmarshalJSON(raw []byte) error { return nil } +type AllowedTrunk struct { + TrunkID string `json:"trunk_id"` + Concurrency int64 `json:"concurrency"` +} + type Task struct { Resource string `json:"resource"` DispatcherID string `json:"dispatcher_id"` @@ -81,7 +86,7 @@ type Task struct { RingTimeoutMS int64 `json:"ring_timeout_ms"` MaxCallDurationMS int64 `json:"max_call_duration_ms"` RoutePolicyID string `json:"route_policy_id"` - AllowedTrunkIDs []string `json:"allowed_trunk_ids"` + AllowedTrunks []AllowedTrunk `json:"allowed_trunk_ids"` Schedule json.RawMessage `json:"schedule"` Agent Agent `json:"agent"` Raw json.RawMessage `json:"-"` @@ -93,11 +98,35 @@ func (t *Task) UnmarshalJSON(raw []byte) error { if err := json.Unmarshal(raw, &parsed); err != nil { return err } - *t = Task(parsed) + decoded := Task(parsed) + if _, err := decoded.TrunkLimits(); err != nil { + return err + } + *t = decoded t.Raw = append([]byte(nil), raw...) return nil } +// TrunkLimits validates the task's per-line limits without changing them. +// Zero disables that line for this task; duplicate IDs are ambiguous even +// when the complete objects differ and therefore pass JSON Schema uniqueness. +func (t Task) TrunkLimits() (map[string]int64, error) { + if len(t.AllowedTrunks) == 0 { + return nil, errors.New("task allowed_trunk_ids is empty") + } + limits := make(map[string]int64, len(t.AllowedTrunks)) + for _, trunk := range t.AllowedTrunks { + if trunk.TrunkID == "" || trunk.Concurrency < 0 { + return nil, errors.New("task trunk ID or concurrency is invalid") + } + if _, exists := limits[trunk.TrunkID]; exists { + return nil, errors.New("task allowed_trunk_ids contains duplicate trunk IDs") + } + limits[trunk.TrunkID] = trunk.Concurrency + } + return limits, nil +} + type Quota struct { Resource string `json:"resource"` DispatcherID string `json:"dispatcher_id"` @@ -132,10 +161,10 @@ func (c *Client) ReadProviders(ctx context.Context) (map[string]Provider, error) } providers := make(map[string]Provider, len(response.Providers)) for _, provider := range response.Providers { - if _, exists := providers[provider.ProviderRef]; exists { - return nil, fmt.Errorf("duplicate AI provider reference %q", provider.ProviderRef) + if _, exists := providers[provider.Code]; exists { + return nil, fmt.Errorf("duplicate AI provider reference %q", provider.Code) } - providers[provider.ProviderRef] = provider + providers[provider.Code] = provider } return providers, nil } @@ -168,8 +197,8 @@ func (c *Client) ReadTask(ctx context.Context, taskID string, tenantID int64, si continue } provider, found := result.Providers[expected.ref] - if !found || provider.ProviderRef != expected.ref { - return Snapshot{}, fmt.Errorf("%s provider_id is missing from the approved connection catalog", expected.role) + if !found || provider.Code != expected.ref { + return Snapshot{}, fmt.Errorf("%s provider_ref is missing from the approved connection catalog", expected.role) } } return result, nil diff --git a/internal/configread/snapshots_test.go b/internal/configread/snapshots_test.go index d4d2fd3..7b9d854 100644 --- a/internal/configread/snapshots_test.go +++ b/internal/configread/snapshots_test.go @@ -41,7 +41,7 @@ func TestReadTaskSnapshot(t *testing.T) { responses := map[string][]byte{ "/internal/v1/dispatcher/sip": example(t, "config-read-sip"), "/internal/v1/dispatcher/ai-providers": example(t, "config-read-providers"), - "/internal/v1/dispatcher/task/task-" + mode: example(t, "config-read-task-"+mode), + "/internal/v1/dispatcher/task/task-" + mode: []byte(strings.Replace(string(example(t, "config-read-task-"+mode)), `"temperature":0.7`, `"temperature":0`, 1)), "/internal/v1/dispatcher/tenant/1001/quota": example(t, "config-read-quota"), } calls := 0 @@ -77,7 +77,7 @@ func TestReadTaskSnapshot(t *testing.T) { if calls != 3 || snapshot.SIP.Revision != 8 || snapshot.Task.TenantID != 1001 || snapshot.Quota.TenantID != 1001 { t.Fatalf("unexpected request count or identity: calls=%d snapshot=%+v", calls, snapshot) } - if got := snapshot.Providers["asr-example"].Credential; got != "example-only-not-a-real-secret" { + if got := snapshot.Providers["volcengine"].Credential; got != "sk-******" { t.Fatalf("credential not passed unchanged: %q", got) } if mode == "asr" && (snapshot.Task.Agent.Mode != "asr_only" || strings.Contains(string(snapshot.Task.Agent.Raw), `"llm"`)) { @@ -86,13 +86,15 @@ func TestReadTaskSnapshot(t *testing.T) { if mode == "full" { var agent struct { LLM struct { - Temperature *float64 `json:"temperature"` + Params struct { + Temperature *float64 `json:"temperature"` + } `json:"params"` } `json:"llm"` Conversation struct { AllowInterrupt *bool `json:"allow_interrupt"` } `json:"conversation"` } - if err := json.Unmarshal(snapshot.Task.Agent.Raw, &agent); err != nil || agent.LLM.Temperature == nil || *agent.LLM.Temperature != 0 || agent.Conversation.AllowInterrupt == nil || *agent.Conversation.AllowInterrupt { + if err := json.Unmarshal(snapshot.Task.Agent.Raw, &agent); err != nil || agent.LLM.Params.Temperature == nil || *agent.LLM.Params.Temperature != 0 || agent.Conversation.AllowInterrupt == nil || *agent.Conversation.AllowInterrupt { t.Fatalf("explicit zero/false lost: %v, %s", err, snapshot.Task.Agent.Raw) } } @@ -232,9 +234,9 @@ func TestReadTaskRejectsMismatchedOwnerAndMalformedSIP(t *testing.T) { case "provider-ref": responses["/internal/v1/dispatcher/ai-providers"] = example(t, "invalid/config-read-provider-ref") case "provider-empty-key": - responses["/internal/v1/dispatcher/ai-providers"] = []byte(strings.Replace(string(responses["/internal/v1/dispatcher/ai-providers"]), `"api_key":"example-only-not-a-real-secret"`, `"api_key":""`, 1)) + responses["/internal/v1/dispatcher/ai-providers"] = []byte(strings.Replace(string(responses["/internal/v1/dispatcher/ai-providers"]), `"api_key":"sk-******"`, `"api_key":""`, 1)) case "provider-missing-id": - responses["/internal/v1/dispatcher/ai-providers"] = []byte(strings.Replace(string(responses["/internal/v1/dispatcher/ai-providers"]), `"provider_id":"asr-example",`, ``, 1)) + responses["/internal/v1/dispatcher/ai-providers"] = []byte(strings.Replace(string(responses["/internal/v1/dispatcher/ai-providers"]), `"provider_id":"3",`, ``, 1)) } server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") diff --git a/internal/configread/trunk_concurrency_test.go b/internal/configread/trunk_concurrency_test.go new file mode 100644 index 0000000..7b5c37a --- /dev/null +++ b/internal/configread/trunk_concurrency_test.go @@ -0,0 +1,75 @@ +package configread + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" +) + +func TestReadTaskTrunkConcurrency(t *testing.T) { + for _, tc := range []struct { + name, trunks string + valid bool + }{ + {"objects", `[{"trunk_id":"12","concurrency":5},{"trunk_id":"15","concurrency":5}]`, true}, + {"disabled", `[{"trunk_id":"12","concurrency":0}]`, true}, + {"old strings", `["12"]`, false}, + {"missing concurrency", `[{"trunk_id":"12"}]`, false}, + {"negative", `[{"trunk_id":"12","concurrency":-1}]`, false}, + {"duplicate IDs", `[{"trunk_id":"12","concurrency":1},{"trunk_id":"12","concurrency":2}]`, false}, + } { + t.Run(tc.name, func(t *testing.T) { + var task map[string]json.RawMessage + if err := json.Unmarshal(example(t, "config-read-task-asr"), &task); err != nil { + t.Fatal(err) + } + task["allowed_trunk_ids"] = json.RawMessage(tc.trunks) + body, err := json.Marshal(task) + if err != nil { + t.Fatal(err) + } + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/internal/v1/dispatcher/task/task-asr": + _, _ = w.Write(body) + case "/internal/v1/dispatcher/tenant/1001/quota": + _, _ = w.Write(example(t, "config-read-quota")) + case "/internal/v1/dispatcher/ai-providers": + _, _ = w.Write(example(t, "config-read-providers")) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + client, err := NewClient(server.URL, "c046b893-8628-4589-ae50-619d049248a6", "isolated-secret", server.Client()) + if err != nil { + t.Fatal(err) + } + providers, err := client.ReadProviders(context.Background()) + if err != nil { + t.Fatal(err) + } + snapshot, err := client.ReadTask(context.Background(), "task-asr", 1001, approvedSIP(t), providers) + if (err == nil) != tc.valid { + t.Fatalf("accepted=%t, want %t: %v", err == nil, tc.valid, err) + } + if !tc.valid { + return + } + encoded, err := json.Marshal(snapshot.Task) + if err != nil { + t.Fatal(err) + } + var roundTrip map[string]json.RawMessage + if err := json.Unmarshal(encoded, &roundTrip); err != nil { + t.Fatal(err) + } + if string(roundTrip["allowed_trunk_ids"]) != tc.trunks { + t.Fatalf("trunk limits changed: %s", roundTrip["allowed_trunk_ids"]) + } + }) + } +} diff --git a/internal/dispatcher/ai_test.go b/internal/dispatcher/ai_test.go index 9ff3e2b..55a11b2 100644 --- a/internal/dispatcher/ai_test.go +++ b/internal/dispatcher/ai_test.go @@ -17,14 +17,14 @@ import ( func unsupportedASRTask(t *testing.T) []byte { t.Helper() original := configExample(t, "config-read-task-asr") - modified := bytes.Replace(original, []byte(`"language":"zh-CN"`), []byte(`"language":"fr-FR"`), 1) + modified := bytes.Replace(original, []byte(`"params":{}`), []byte(`"params":{"reqid":"override"}`), 1) if bytes.Equal(original, modified) { - t.Fatal("ASR fixture did not contain the expected language") + t.Fatal("ASR fixture did not contain params") } return modified } -func TestBootstrapRejectsSDKUnsupportedTaskBeforeOpeningAdmission(t *testing.T) { +func TestBootstrapRejectsProtocolOverrideBeforeOpeningAdmission(t *testing.T) { id := "c046b893-8628-4589-ae50-619d049248a6" server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { var body []byte @@ -67,7 +67,7 @@ func TestBootstrapRejectsSDKUnsupportedTaskBeforeOpeningAdmission(t *testing.T) DrainControls: func(context.Context) error { drained = true; return nil }, } err = b.Run(context.Background()) - if err == nil || !strings.Contains(err.Error(), "ASR language") || drained { + if err == nil || !strings.Contains(err.Error(), "protocol field") || drained { t.Fatalf("SDK-unsupported task must not be admitted or drained: %v, drained=%t", err, drained) } if admitted, err := db.CanAdmit(id, 1001, "task-asr"); err != nil || admitted { @@ -104,7 +104,7 @@ func TestExecuteCorruptAISnapshotFailsClosedBeforeReservation(t *testing.T) { Now: func() time.Time { return monday(9, 30) }, } err = controller.ProcessExecute(context.Background(), executeBody(t, "bad-ai", "15003164745")) - if err == nil || !strings.Contains(err.Error(), "ASR language") || len(originator.calls) != 0 { + if err == nil || !strings.Contains(err.Error(), "protocol field") || len(originator.calls) != 0 { t.Fatalf("corrupt AI config must not reach Agent: err=%v calls=%d", err, len(originator.calls)) } if admitted, err := db.CanAdmit(id, 1001, "task-asr"); err != nil || admitted { diff --git a/internal/dispatcher/execute.go b/internal/dispatcher/execute.go index 4eb6cfb..cdd2480 100644 --- a/internal/dispatcher/execute.go +++ b/internal/dispatcher/execute.go @@ -167,7 +167,11 @@ func (c *ExecuteController) dispatchPending(ctx context.Context, cmd store.Execu if err != nil { return err } - choice, err := SelectTrunk(snapshot, cmd.Callee, at, occupied, loaded) + taskOccupied, err := c.Store.TaskTrunkOccupancy(c.DispatcherID, cmd.TenantID, cmd.TaskID) + if err != nil { + return fmt.Errorf("read task trunk occupancy: %w", err) + } + choice, err := SelectTrunk(snapshot, cmd.Callee, at, occupied, taskOccupied, loaded) if errors.Is(err, ErrRuleWait) { return nil } @@ -175,7 +179,7 @@ func (c *ExecuteController) dispatchPending(ctx context.Context, cmd store.Execu return fmt.Errorf("task %q rule validation: %w", cmd.TaskID, err) } reservation := store.CallReservation{ - TrunkID: choice.TrunkID, SIPRevision: snapshot.SIP.Revision, + TrunkID: choice.TrunkID, SIPRevision: snapshot.SIP.Revision, TaskRevision: snapshot.Task.TaskRevision, CallerID: choice.CallerID, DialedCallee: choice.DialedCallee, Deadline: choice.Deadline, } @@ -193,7 +197,7 @@ func (c *ExecuteController) dispatchPending(ctx context.Context, cmd store.Execu cancelErr := c.Store.CancelReservationBeforeOrigin(cmd.DispatcherID, cmd.EventID) return errors.Join(fmt.Errorf("recheck applied SIP before originate: %w", err), cancelErr) } - fresh, ruleErr := SelectTrunk(snapshot, cmd.Callee, actualAt, occupied, freshLoaded) + fresh, ruleErr := SelectTrunk(snapshot, cmd.Callee, actualAt, occupied, taskOccupied, freshLoaded) if !actualAt.Before(choice.Deadline) || ruleErr != nil || fresh.TrunkID != choice.TrunkID || fresh.CallerID != choice.CallerID { cancelErr := c.Store.CancelReservationBeforeOrigin(cmd.DispatcherID, cmd.EventID) if cancelErr != nil { @@ -220,6 +224,11 @@ func (c *ExecuteController) dispatchPending(ctx context.Context, cmd store.Execu MaxCallDurationMS: choice.MaxCallDurationMS, Deadline: choice.Deadline, Snapshot: snapshot, } + limits, err := snapshot.Task.TrunkLimits() + if err != nil { + return errors.Join(fmt.Errorf("read reserved task trunk limit: %w", err), c.Store.CancelReservationBeforeOrigin(cmd.DispatcherID, cmd.EventID)) + } + log.Printf("Dispatcher call reserved before originate: event_id=%q task_revision=%d trunk_id=%q task_trunk_limit=%d task_trunk_used_before=%d", cmd.EventID, snapshot.Task.TaskRevision, choice.TrunkID, limits[choice.TrunkID], taskOccupied[choice.TrunkID]) if err := c.Originator.Originate(ctx, spec); err != nil { if agentProvedCallNotIssued(err) { rejectErr := c.Store.RejectUnissuedExecute(cmd.DispatcherID, cmd.EventID, "Agent refused before outbound call") diff --git a/internal/dispatcher/policy.go b/internal/dispatcher/policy.go index e772ff6..25553cf 100644 --- a/internal/dispatcher/policy.go +++ b/internal/dispatcher/policy.go @@ -56,7 +56,7 @@ func validRawCallee(callee string) bool { // SelectTrunk makes one ordered choice before originate. Its answer is // frozen with the accepted command; callers never silently reselect on an // unknown execution or after a failed originate. -func SelectTrunk(snapshot configread.Snapshot, callee string, at time.Time, trunkOccupancy, loadedRevisions map[string]int64) (SelectedTrunk, error) { +func SelectTrunk(snapshot configread.Snapshot, callee string, at time.Time, trunkOccupancy, taskTrunkOccupancy, loadedRevisions map[string]int64) (SelectedTrunk, error) { if !validRawCallee(callee) { return SelectedTrunk{}, ErrCalleeRejected } @@ -81,8 +81,9 @@ func SelectTrunk(snapshot configread.Snapshot, callee string, at time.Time, trun } byID[trunk.TrunkID] = trunk } - if len(snapshot.Task.AllowedTrunkIDs) == 0 { - return SelectedTrunk{}, fmt.Errorf("%w: task has no allowed trunks", ErrRuleInvalid) + limits, err := snapshot.Task.TrunkLimits() + if err != nil { + return SelectedTrunk{}, fmt.Errorf("%w: %v", ErrRuleInvalid, err) } maxDurationMS := snapshot.Task.MaxCallDurationMS if snapshot.Task.Agent.Mode == "full_ai" { @@ -107,7 +108,8 @@ func SelectTrunk(snapshot configread.Snapshot, callee string, at time.Time, trun return SelectedTrunk{}, fmt.Errorf("%w: call duration exceeds time range", ErrRuleInvalid) } var invalid error - for _, id := range snapshot.Task.AllowedTrunkIDs { + for _, allowed := range snapshot.Task.AllowedTrunks { + id := allowed.TrunkID trunk, found := byID[id] if !found { invalid = fmt.Errorf("%w: allowed trunk %q is missing", ErrRuleInvalid, id) @@ -129,7 +131,7 @@ func SelectTrunk(snapshot configread.Snapshot, callee string, at time.Time, trun invalid = fmt.Errorf("%w: schedule of trunk %q: %v", ErrRuleInvalid, id, err) continue } - if !trunk.Enabled || loadedRevisions[id] != snapshot.SIP.Revision || trunkOccupancy[id] >= *trunk.MaxConcurrentCalls { + if !trunk.Enabled || limits[id] == 0 || taskTrunkOccupancy[id] >= limits[id] || loadedRevisions[id] != snapshot.SIP.Revision || trunkOccupancy[id] >= *trunk.MaxConcurrentCalls { continue } deadline := at.Add(time.Duration(maxDurationMS) * time.Millisecond) diff --git a/internal/dispatcher/policy_test.go b/internal/dispatcher/policy_test.go index 4faaf34..0a06da5 100644 --- a/internal/dispatcher/policy_test.go +++ b/internal/dispatcher/policy_test.go @@ -34,7 +34,7 @@ func policySnapshot(t *testing.T) configread.Snapshot { } snapshot.Providers = make(map[string]configread.Provider, len(providerList.Providers)) for _, provider := range providerList.Providers { - snapshot.Providers[provider.ProviderRef] = provider + snapshot.Providers[provider.Code] = provider } return snapshot } @@ -46,7 +46,7 @@ func monday(hour, minute int) time.Time { func TestSelectTrunkSaaSEventCalleeScheduleAndCaller(t *testing.T) { snapshot := policySnapshot(t) at := monday(9, 30) - selected, err := SelectTrunk(snapshot, "15003164745", at, map[string]int64{}, map[string]int64{"trunk-mock": 8}) + selected, err := SelectTrunk(snapshot, "15003164745", at, map[string]int64{}, nil, map[string]int64{"trunk-mock": 8}) if err != nil { t.Fatal(err) } @@ -54,18 +54,18 @@ func TestSelectTrunkSaaSEventCalleeScheduleAndCaller(t *testing.T) { t.Fatalf("selected route/prefix/caller/deadline invalid: %+v", selected) } for _, number := range []string{"15803300952", "13900000000", "15003164746"} { - choice, err := SelectTrunk(snapshot, number, at, nil, map[string]int64{"trunk-mock": 8}) + choice, err := SelectTrunk(snapshot, number, at, nil, nil, map[string]int64{"trunk-mock": 8}) if err != nil || choice.Callee != number || choice.DialedCallee != number { t.Fatalf("SaaS event number changed or rejected: number=%q choice=%+v err=%v", number, choice, err) } } for _, number := range []string{"", "abc", "15803300952x", "15803300952\n", strings.Repeat("7", 33)} { - if _, err := SelectTrunk(snapshot, number, at, nil, map[string]int64{"trunk-mock": 8}); !errors.Is(err, ErrCalleeRejected) { + if _, err := SelectTrunk(snapshot, number, at, nil, nil, map[string]int64{"trunk-mock": 8}); !errors.Is(err, ErrCalleeRejected) { t.Fatalf("invalid raw dial route %q not rejected: %v", number, err) } } for _, at := range []time.Time{monday(8, 59), monday(20, 0)} { - if _, err := SelectTrunk(snapshot, "15003164745", at, nil, map[string]int64{"trunk-mock": 8}); !errors.Is(err, ErrRuleWait) { + if _, err := SelectTrunk(snapshot, "15003164745", at, nil, nil, map[string]int64{"trunk-mock": 8}); !errors.Is(err, ErrRuleWait) { t.Fatalf("outside half-open schedule did not wait: %v", err) } } @@ -85,8 +85,8 @@ func TestSelectTrunkUsesEachLinesOwnCaller(t *testing.T) { if err != nil { t.Fatal(err) } - snapshot.Task.AllowedTrunkIDs = []string{second.TrunkID} - chosen, err := SelectTrunk(snapshot, "15003164745", monday(9, 30), nil, map[string]int64{second.TrunkID: 8}) + snapshot.Task.AllowedTrunks = []configread.AllowedTrunk{{TrunkID: second.TrunkID, Concurrency: 2}} + chosen, err := SelectTrunk(snapshot, "15003164745", monday(9, 30), nil, nil, map[string]int64{second.TrunkID: 8}) if err != nil || chosen.TrunkID != second.TrunkID || chosen.CallerID != second.CallerID { t.Fatalf("second line did not keep its own approved caller: %+v, %v", chosen, err) } @@ -95,7 +95,7 @@ func TestSelectTrunkUsesEachLinesOwnCaller(t *testing.T) { if err != nil { t.Fatal(err) } - if _, err := SelectTrunk(snapshot, "15003164745", monday(9, 30), nil, map[string]int64{second.TrunkID: 8}); !errors.Is(err, ErrRuleInvalid) { + if _, err := SelectTrunk(snapshot, "15003164745", monday(9, 30), nil, nil, map[string]int64{second.TrunkID: 8}); !errors.Is(err, ErrRuleInvalid) { t.Fatalf("second line with missing caller did not fail closed: %v", err) } } @@ -103,17 +103,17 @@ func TestSelectTrunkUsesEachLinesOwnCaller(t *testing.T) { func TestSelectTrunkFailsClosedForUnknownLineAndLimits(t *testing.T) { snapshot := policySnapshot(t) at := monday(9, 30) - if _, err := SelectTrunk(snapshot, "15003164745", at, nil, nil); !errors.Is(err, ErrRuleWait) { + if _, err := SelectTrunk(snapshot, "15003164745", at, nil, nil, nil); !errors.Is(err, ErrRuleWait) { t.Fatalf("unloaded line did not hold command: %v", err) } - if _, err := SelectTrunk(snapshot, "15003164745", at, map[string]int64{"trunk-mock": 2}, map[string]int64{"trunk-mock": 8}); !errors.Is(err, ErrRuleWait) { + if _, err := SelectTrunk(snapshot, "15003164745", at, map[string]int64{"trunk-mock": 2}, nil, map[string]int64{"trunk-mock": 8}); !errors.Is(err, ErrRuleWait) { t.Fatalf("full trunk did not hold command: %v", err) } sip := string(configExample(t, "config-read-sip")) if err := json.Unmarshal([]byte(sip), &snapshot.SIP); err != nil { t.Fatal(err) } - if _, err := SelectTrunk(snapshot, "15003164745", at, nil, map[string]int64{"trunk-mock": 8}); !errors.Is(err, ErrRuleInvalid) { + if _, err := SelectTrunk(snapshot, "15003164745", at, nil, nil, map[string]int64{"trunk-mock": 8}); !errors.Is(err, ErrRuleInvalid) { t.Fatalf("unknown transport/auth/quota was not fail-closed: %v", err) } } @@ -124,7 +124,7 @@ func TestSelectTrunkCapsCallByAIAndPreservesPrefix(t *testing.T) { snapshot.Task.Agent.Raw = []byte(`{"mode":"full_ai","conversation":{"max_duration_ms":60000}}`) snapshot.SIP.Trunks = []byte(strings.Replace(string(snapshot.SIP.Trunks), `"dial_prefix":""`, `"dial_prefix":"7089"`, 1)) at := monday(9, 30) - selected, err := SelectTrunk(snapshot, "15830461047", at, nil, map[string]int64{"trunk-mock": 8}) + selected, err := SelectTrunk(snapshot, "15830461047", at, nil, nil, map[string]int64{"trunk-mock": 8}) if err != nil { t.Fatal(err) } diff --git a/internal/dispatcher/task_trunk_concurrency_test.go b/internal/dispatcher/task_trunk_concurrency_test.go new file mode 100644 index 0000000..3f04306 --- /dev/null +++ b/internal/dispatcher/task_trunk_concurrency_test.go @@ -0,0 +1,164 @@ +package dispatcher + +import ( + "context" + "encoding/json" + "errors" + "testing" + + "git.ipao.vip/rogee/go-sip/internal/configread" +) + +func TestSelectTrunkHonorsTaskLineLimit(t *testing.T) { + for _, tc := range []struct { + name string + limit, task, total int64 + wantErr error + }{ + {"available", 1, 0, 0, nil}, + {"other task occupancy", 1, 0, 1, nil}, + {"disabled", 0, 0, 0, ErrRuleWait}, + {"task line full", 1, 1, 1, ErrRuleWait}, + {"unknown line occupied", 1, 2, 0, ErrRuleWait}, + {"global line full", 5, 0, 2, ErrRuleWait}, + {"negative limit", -1, 0, 0, ErrRuleInvalid}, + } { + t.Run(tc.name, func(t *testing.T) { + snapshot := policySnapshot(t) + snapshot.Task.AllowedTrunks = []configread.AllowedTrunk{{TrunkID: "trunk-mock", Concurrency: tc.limit}} + choice, err := SelectTrunk(snapshot, "15003164745", monday(9, 30), map[string]int64{"trunk-mock": tc.total}, map[string]int64{"trunk-mock": tc.task}, map[string]int64{"trunk-mock": 8}) + if !errors.Is(err, tc.wantErr) { + t.Fatalf("selected=%q error=%v, want %v", choice.TrunkID, err, tc.wantErr) + } + }) + } +} + +func TestSelectTrunkSkipsDisabledAndSaturatedTaskLinesInConfiguredOrder(t *testing.T) { + for _, limit := range []int64{0, 1} { + snapshot := policySnapshot(t) + var trunks []trunkConfig + if err := json.Unmarshal(snapshot.SIP.Trunks, &trunks); err != nil { + t.Fatal(err) + } + second := trunks[0] + second.TrunkID = "trunk-second" + trunks = append(trunks, second) + var err error + snapshot.SIP.Trunks, err = json.Marshal(trunks) + if err != nil { + t.Fatal(err) + } + snapshot.Task.AllowedTrunks = []configread.AllowedTrunk{{TrunkID: "trunk-mock", Concurrency: limit}, {TrunkID: "trunk-second", Concurrency: 1}} + choice, err := SelectTrunk(snapshot, "15003164745", monday(9, 30), map[string]int64{"trunk-mock": 1}, map[string]int64{"trunk-mock": 1}, map[string]int64{"trunk-mock": 8, "trunk-second": 8}) + if err != nil || choice.TrunkID != "trunk-second" { + t.Fatalf("limit=%d selected=%q err=%v", limit, choice.TrunkID, err) + } + } +} + +func TestSelectTrunkDoesNotHideMissingLineBehindZeroConcurrency(t *testing.T) { + snapshot := policySnapshot(t) + snapshot.Task.AllowedTrunks = []configread.AllowedTrunk{{TrunkID: "missing", Concurrency: 0}} + if _, err := SelectTrunk(snapshot, "15003164745", monday(9, 30), nil, nil, nil); !errors.Is(err, ErrRuleInvalid) { + t.Fatalf("disabled line hid a broken reference: %v", err) + } +} + +func TestSelectTrunkRejectsDuplicateTaskLineIDs(t *testing.T) { + snapshot := policySnapshot(t) + snapshot.Task.AllowedTrunks = []configread.AllowedTrunk{{TrunkID: "trunk-mock", Concurrency: 1}, {TrunkID: "trunk-mock", Concurrency: 2}} + _, err := SelectTrunk(snapshot, "15003164745", monday(9, 30), nil, nil, map[string]int64{"trunk-mock": 8}) + if !errors.Is(err, ErrRuleInvalid) { + t.Fatalf("duplicate task line accepted: %v", err) + } +} + +func editedTaskLine(t *testing.T, snapshot configread.Snapshot, revision, limit int64) configread.Snapshot { + t.Helper() + var task map[string]json.RawMessage + if err := json.Unmarshal(snapshot.Task.Raw, &task); err != nil { + t.Fatal(err) + } + for key, value := range map[string]any{ + "task_revision": revision, + "allowed_trunk_ids": []configread.AllowedTrunk{{TrunkID: "trunk-mock", Concurrency: limit}}, + } { + raw, err := json.Marshal(value) + if err != nil { + t.Fatal(err) + } + task[key] = raw + } + raw, err := json.Marshal(task) + if err != nil { + t.Fatal(err) + } + if err := json.Unmarshal(raw, &snapshot.Task); err != nil { + t.Fatal(err) + } + return snapshot +} + +func TestExecuteControllerKeepsUnknownTaskLineSlotUntilConfirmedEnd(t *testing.T) { + c, origin, _, _ := newExecuteFixture(t) + snapshot := editedTaskLine(t, policySnapshot(t), 2, 1) + if state, err := c.Store.CompleteEdit(snapshot, "edit-line-limit"); err != nil || state != "applied" { + t.Fatalf("edit=%q err=%v", state, err) + } + ctx := context.Background() + lostReply := errors.New("simulated originate reply lost") + origin.err = lostReply + if err := c.ProcessExecute(ctx, executeBody(t, "call-first", "15003164745")); !errors.Is(err, lostReply) { + t.Fatalf("unknown originate outcome=%v", err) + } + origin.err = nil + if err := c.ProcessExecute(ctx, executeBody(t, "call-second", "15003164745")); err != nil { + t.Fatal(err) + } + if len(origin.calls) != 1 { + t.Fatalf("task-line limit exceeded: attempts=%d", len(origin.calls)) + } + if err := c.ProcessPending(ctx); err != nil || len(origin.calls) != 1 { + t.Fatalf("unknown call was retried or lost its slot: calls=%d err=%v", len(origin.calls), err) + } + if err := c.Store.FinishExecute(c.DispatcherID, "call-first"); err != nil { + t.Fatal(err) + } + if err := c.ProcessPending(ctx); err != nil || len(origin.calls) != 2 { + t.Fatalf("confirmed end did not admit the pending call: calls=%d err=%v", len(origin.calls), err) + } + if origin.calls[1].Snapshot.Task.TaskRevision != 2 || origin.calls[1].Snapshot.Task.AllowedTrunks[0].Concurrency != 1 { + t.Fatal("new task-line limit was not frozen into the approved call") + } +} + +type editingTaskOriginator struct { + *fakeOriginator + edit func() +} + +func (o *editingTaskOriginator) LoadedTrunks(ctx context.Context) (map[string]int64, error) { + if o.edit != nil { + edit := o.edit + o.edit = nil + edit() + } + return o.fakeOriginator.LoadedTrunks(ctx) +} + +func TestExecuteControllerRejectsSelectionIfTaskChangesBeforeReservation(t *testing.T) { + c, origin, _, _ := newExecuteFixture(t) + c.Originator = &editingTaskOriginator{fakeOriginator: origin, edit: func() { + snapshot := editedTaskLine(t, policySnapshot(t), 2, 0) + if state, err := c.Store.CompleteEdit(snapshot, "disable-before-reserve"); err != nil || state != "applied" { + t.Fatalf("edit=%q err=%v", state, err) + } + }} + if err := c.ProcessExecute(context.Background(), executeBody(t, "call-stale-choice", "15003164745")); err != nil { + t.Fatal(err) + } + if len(origin.calls) != 0 { + t.Fatal("stale selection originated after the task line was disabled") + } +} diff --git a/internal/rpc/approved_execution.go b/internal/rpc/approved_execution.go index 8c2e7ab..7180333 100644 --- a/internal/rpc/approved_execution.go +++ b/internal/rpc/approved_execution.go @@ -16,6 +16,7 @@ import ( "git.ipao.vip/rogee/go-sip/internal/agent" "git.ipao.vip/rogee/go-sip/internal/ai" "git.ipao.vip/rogee/go-sip/internal/configread" + "git.ipao.vip/rogee/go-sip/internal/contract" "google.golang.org/genproto/googleapis/rpc/errdetails" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" @@ -120,6 +121,9 @@ func (s *Server) ExecuteApproved(ctx context.Context, req *agentpb.ExecuteApprov if err != nil || hash != req.BindingSha256 { return nil, status.Error(codes.FailedPrecondition, "bound execution snapshot does not match") } + if err := contract.ValidateCurrent("config-read", req.TaskConfigJson); err != nil { + return nil, status.Error(codes.InvalidArgument, "approved task schema is invalid") + } var task configread.Task if err := json.Unmarshal(req.TaskConfigJson, &task); err != nil { return nil, status.Error(codes.InvalidArgument, "approved task JSON is invalid") @@ -127,6 +131,10 @@ func (s *Server) ExecuteApproved(ctx context.Context, req *agentpb.ExecuteApprov if task.DispatcherID != dispatcherID || task.TenantID != req.TenantId || task.TaskID != req.TaskId { return nil, status.Error(codes.FailedPrecondition, "approved task identity does not match") } + limits, err := task.TrunkLimits() + if err != nil || limits[req.SelectedTrunkId] <= 0 { + return nil, status.Error(codes.FailedPrecondition, "selected trunk is not enabled for this task") + } var providers map[string]configread.Provider if err := json.Unmarshal(req.ProvidersJson, &providers); err != nil { return nil, status.Error(codes.InvalidArgument, "approved provider JSON is invalid") diff --git a/internal/rpc/approved_execution_test.go b/internal/rpc/approved_execution_test.go index 6411567..72daf36 100644 --- a/internal/rpc/approved_execution_test.go +++ b/internal/rpc/approved_execution_test.go @@ -54,7 +54,10 @@ func approvedTestRequest(t *testing.T, now time.Time) *agentpb.ExecuteApprovedRe } providers := make(map[string]configread.Provider) for _, p := range payload.Providers { - providers[p.ProviderRef] = p + providers[p.Code] = p + } + for _, code := range []string{"openai_compatible", "ali_bailian"} { + providers[code] = configread.Provider{ProviderRef: "test-" + code, Code: code, Name: "Example only", Endpoint: "https://example.invalid", WSEndpoint: "wss://example.invalid", Credential: "example-only-not-a-real-secret"} } providerJSON, err := json.Marshal(providers) if err != nil { @@ -67,7 +70,7 @@ func approvedTestRequest(t *testing.T, now time.Time) *agentpb.ExecuteApprovedRe return &agentpb.ExecuteApprovedRequest{ Meta: testMeta("event-1", "event-1", 1), DispatcherId: task.DispatcherID, TenantId: task.TenantID, TaskId: task.TaskID, - SourceEventId: "event-1", CallId: "event-1", SelectedTrunkId: task.AllowedTrunkIDs[0], + SourceEventId: "event-1", CallId: "event-1", SelectedTrunkId: task.AllowedTrunks[0].TrunkID, CallerId: "BD93205882", Callee: "15003164745", DialedCallee: "708915003164745", RingTimeoutMs: 10000, MaxCallDurationMs: 20000, DialBeforeUnixMs: now.Add(time.Minute).UnixMilli(), SipRevision: 8, TaskConfigJson: taskJSON, ProvidersJson: providerJSON, BindingSha256: hash, diff --git a/internal/rpc/approved_fixture_test.go b/internal/rpc/approved_fixture_test.go new file mode 100644 index 0000000..ae276db --- /dev/null +++ b/internal/rpc/approved_fixture_test.go @@ -0,0 +1,83 @@ +package rpc + +import ( + "context" + "encoding/json" + "github.com/coder/websocket" + "net/http" + "os" + "testing" + "time" +) + +func approvedFullTestJSON(t *testing.T) []byte { + t.Helper() + raw, err := os.ReadFile("../../contracts/schema/examples/config-read-task-full.json") + if err != nil { + t.Fatal(err) + } + var body map[string]any + if json.Unmarshal(raw, &body) != nil { + t.Fatal("decode full task fixture") + } + a := body["agent"].(map[string]any) + a["llm"] = map[string]any{"provider_ref": "openai_compatible", "model": "example-chat", "params": map[string]any{"temperature": 0, "max_tokens": 256}} + a["tts"] = map[string]any{"provider_ref": "ali_bailian", "model": "qwen3-tts-flash", "voice": "Cherry", "params": map[string]any{"format": "pcm", "sample_rate": 16000}} + raw, err = json.Marshal(body) + if err != nil { + t.Fatal(err) + } + return raw +} + +func serveApprovedTTSMock(t *testing.T, w http.ResponseWriter, r *http.Request, observe func(map[string]any)) { + t.Helper() + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.CloseNow() + ctx, cancel := context.WithTimeout(r.Context(), 3*time.Second) + defer cancel() + read := func() map[string]any { + _, raw, err := conn.Read(ctx) + if err != nil { + t.Error(err) + return nil + } + var command map[string]any + if json.Unmarshal(raw, &command) != nil { + t.Error("invalid task command") + } + return command + } + start := read() + if start == nil { + return + } + id := start["header"].(map[string]any)["task_id"].(string) + event := func(name string) { + raw, _ := json.Marshal(map[string]any{"header": map[string]any{"task_id": id, "event": name}, "payload": map[string]any{}}) + if err := conn.Write(ctx, websocket.MessageText, raw); err != nil { + t.Error(err) + } + } + event("task-started") + next := read() + if next == nil { + return + } + finish := read() + if finish == nil { + return + } + p := start["payload"].(map[string]any) + input := next["payload"].(map[string]any)["input"] + observe(map[string]any{"model": p["model"], "parameters": p["parameters"], "input": input}) + if err := conn.Write(ctx, websocket.MessageBinary, []byte{1, 0, 2, 0}); err != nil { + t.Error(err) + return + } + event("task-finished") +} diff --git a/internal/rpc/approved_full_ai_integration_test.go b/internal/rpc/approved_full_ai_integration_test.go index ec1a4af..6af4f53 100644 --- a/internal/rpc/approved_full_ai_integration_test.go +++ b/internal/rpc/approved_full_ai_integration_test.go @@ -8,8 +8,8 @@ import ( "net/http" "net/http/httptest" "os" - "os/exec" "path/filepath" + "strings" "testing" "time" @@ -18,7 +18,6 @@ import ( "git.ipao.vip/rogee/go-sip/internal/callflow" "git.ipao.vip/rogee/go-sip/internal/configread" "git.ipao.vip/rogee/go-sip/internal/dispatcher" - "git.ipao.vip/rogee/go-sip/internal/media" ) type approvedSDKObservation struct { @@ -35,15 +34,9 @@ func (p finalUserSDKPipeline) RunTurn(ctx context.Context, _ []byte) (ai.TurnRes } func TestApprovedFullAIUsesBoundProviderSDKsAndOnlyFinalUserKeyword(t *testing.T) { - if _, err := exec.LookPath("ffmpeg"); err != nil { - t.Skip("Bailian TTS conversion requires ffmpeg") - } now := time.Date(2026, 9, 20, 10, 0, 0, 0, time.UTC) req := approvedTestRequest(t, now) - fullJSON, err := os.ReadFile("../../contracts/schema/examples/config-read-task-full.json") - if err != nil { - t.Fatal(err) - } + fullJSON := approvedFullTestJSON(t) req.TaskConfigJson = fullJSON var task configread.Task if err := json.Unmarshal(fullJSON, &task); err != nil { @@ -53,14 +46,12 @@ func TestApprovedFullAIUsesBoundProviderSDKsAndOnlyFinalUserKeyword(t *testing.T req.SourceEventId, req.CallId = "event-full", "event-full" llmObserved := make(chan approvedSDKObservation, 1) ttsObserved := make(chan approvedSDKObservation, 3) - wav, _, err := media.EncodeMonoWAV([]byte{1, 0, 2, 0}, 1024) - if err != nil { - t.Fatal(err) - } var mockSDKs *httptest.Server mockSDKs = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/audio" && r.Method == http.MethodGet { - _, _ = w.Write(wav) + if r.URL.Path == "/tts" { + serveApprovedTTSMock(t, w, r, func(body map[string]any) { + ttsObserved <- approvedSDKObservation{Body: body, Authorization: r.Header.Get("Authorization")} + }) return } var body map[string]any @@ -73,10 +64,6 @@ func TestApprovedFullAIUsesBoundProviderSDKsAndOnlyFinalUserKeyword(t *testing.T llmObserved <- approvedSDKObservation{Body: body, Authorization: r.Header.Get("Authorization")} w.Header().Set("Content-Type", "application/json") _, _ = fmt.Fprintln(w, `{"id":"mock","object":"chat.completion","created":0,"model":"example-chat","choices":[{"index":0,"message":{"role":"assistant","content":"回答"},"finish_reason":"stop"}]}`) - case "/api/v1/services/aigc/multimodal-generation/generation": - ttsObserved <- approvedSDKObservation{Body: body, Authorization: r.Header.Get("Authorization")} - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"output": map[string]any{"audio": map[string]any{"url": mockSDKs.URL + "/audio?sig=fixture"}}}) default: http.Error(w, "unexpected SDK request", http.StatusNotFound) } @@ -86,12 +73,13 @@ func TestApprovedFullAIUsesBoundProviderSDKsAndOnlyFinalUserKeyword(t *testing.T if err := json.Unmarshal(req.ProvidersJson, &providers); err != nil { t.Fatal(err) } - llm := providers["llm-example"] + llm := providers["openai_compatible"] llm.Endpoint = mockSDKs.URL + "/v1" - providers[llm.ProviderRef] = llm - tts := providers["tts-example"] - tts.Endpoint = mockSDKs.URL + "/api/v1/services/aigc/multimodal-generation/generation" - providers[tts.ProviderRef] = tts + providers[llm.Code] = llm + tts := providers["ali_bailian"] + tts.WSEndpoint = "ws" + strings.TrimPrefix(mockSDKs.URL, "http") + "/tts" + providers[tts.Code] = tts + var err error req.ProvidersJson, err = json.Marshal(providers) if err != nil { t.Fatal(err) @@ -104,7 +92,7 @@ func TestApprovedFullAIUsesBoundProviderSDKsAndOnlyFinalUserKeyword(t *testing.T agent := activatedApprovedServer(t, now, filepath.Join(t.TempDir(), "agent-session.json"), req.DispatcherId, 1, func(ctx context.Context, approved ApprovedExecution) error { calls++ - if approved.AI.Mode != "full_ai" || approved.AI.LLM == nil || approved.AI.TTS == nil || approved.AI.Conversation.MaxTurns != 20 || approved.AI.LLM.Temperature == nil || *approved.AI.LLM.Temperature != 0 { + if approved.AI.Mode != "full_ai" || approved.AI.LLM == nil || approved.AI.TTS == nil || approved.AI.Conversation.MaxTurns != 20 || string(approved.AI.LLM.Params["temperature"]) != "0" { t.Fatal("approved full-AI settings were not delivered to the Agent") } session := callflow.NewMemorySession(bytes.Repeat([]byte{1, 0}, 1600)) @@ -180,8 +168,8 @@ func TestApprovedFullAIUsesBoundProviderSDKsAndOnlyFinalUserKeyword(t *testing.T if ttsCall.Authorization != "Bearer "+tts.Credential || ttsCall.Body["model"] != "qwen3-tts-flash" { t.Fatal("Agent did not use approved Bailian TTS model and credential") } - params := ttsCall.Body["input"].(map[string]any) - if openingCall.Body["input"].(map[string]any)["text"] != "Example greeting" || params["voice"] != "Cherry" || params["text"] != "回答" { + params := ttsCall.Body["parameters"].(map[string]any) + if openingCall.Body["input"].(map[string]any)["text"] != "Example greeting" || params["voice"] != "Cherry" || ttsCall.Body["input"].(map[string]any)["text"] != "回答" { t.Fatal("Agent did not send the approved opening then reply via Bailian TTS") } } diff --git a/internal/rpc/approved_integration_test.go b/internal/rpc/approved_integration_test.go index 8366e8c..eabc674 100644 --- a/internal/rpc/approved_integration_test.go +++ b/internal/rpc/approved_integration_test.go @@ -43,7 +43,7 @@ func TestApprovedDispatcherToAgentUnaryMockRetainsSnapshotAndOneShotCall(t *test agent := activatedApprovedServer(t, now, filepath.Join(t.TempDir(), "agent-session.json"), req.DispatcherId, 1, func(_ context.Context, call ApprovedExecution) error { calls++ - if call.SIPRevision != 8 || call.AI.Mode != "asr_only" || call.AI.ASR.Provider.Credential != "example-only-not-a-real-secret" || call.MaxCallDuration != 20*time.Second { + if call.SIPRevision != 8 || call.AI.Mode != "asr_only" || call.AI.ASR.Provider.Credential != "sk-******" || call.MaxCallDuration != 20*time.Second { t.Fatal("Agent received the wrong frozen execution configuration") } return nil diff --git a/internal/rpc/approved_runner_test.go b/internal/rpc/approved_runner_test.go index 4bc140a..a19c433 100644 --- a/internal/rpc/approved_runner_test.go +++ b/internal/rpc/approved_runner_test.go @@ -6,8 +6,6 @@ import ( "encoding/json" "net/http" "net/http/httptest" - "os" - "os/exec" "strings" "sync/atomic" "testing" @@ -65,13 +63,7 @@ func TestApprovedCallRunnerASROnlyHonorsSignedTimeout(t *testing.T) { } func TestApprovedCallRunnerSynthesizesOpeningBeforeMediaCapture(t *testing.T) { - if _, err := exec.LookPath("ffmpeg"); err != nil { - t.Skip("Bailian TTS conversion requires ffmpeg") - } - raw, err := os.ReadFile("../../contracts/schema/examples/config-read-task-full.json") - if err != nil { - t.Fatal(err) - } + raw := approvedFullTestJSON(t) var task configread.Task if err := json.Unmarshal(raw, &task); err != nil { t.Fatal(err) @@ -82,33 +74,18 @@ func TestApprovedCallRunnerSynthesizesOpeningBeforeMediaCapture(t *testing.T) { t.Fatal(err) } var sdkCalls atomic.Int32 - wav, _, err := media.EncodeMonoWAV([]byte{1, 0, 2, 0}, 1024) - if err != nil { - t.Fatal(err) - } - var server *httptest.Server - server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/audio" { - _, _ = w.Write(wav) - return - } - sdkCalls.Add(1) - var payload struct { - Input struct { - Text string `json:"text"` - } `json:"input"` - } - if err := json.NewDecoder(r.Body).Decode(&payload); err != nil || payload.Input.Text != "Example greeting" { - http.Error(w, "opening text changed", http.StatusBadRequest) - return - } - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"output": map[string]any{"audio": map[string]any{"url": server.URL + "/audio?sig=fixture"}}}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + serveApprovedTTSMock(t, w, r, func(payload map[string]any) { + sdkCalls.Add(1) + if payload["input"].(map[string]any)["text"] != "Example greeting" { + t.Error("opening text changed") + } + }) })) defer server.Close() - tts := providers["tts-example"] - tts.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation" - providers[tts.ProviderRef] = tts + tts := providers["ali_bailian"] + tts.WSEndpoint = "ws" + strings.TrimPrefix(server.URL, "http") + providers[tts.Code] = tts bound, err := ai.Bind(task, providers) if err != nil { t.Fatal(err) diff --git a/internal/rpc/approved_snapshot_test.go b/internal/rpc/approved_snapshot_test.go index 5659321..ece638c 100644 --- a/internal/rpc/approved_snapshot_test.go +++ b/internal/rpc/approved_snapshot_test.go @@ -24,7 +24,7 @@ func TestApprovedAgentReceivesFrozenASROnlySnapshot(t *testing.T) { if call.DispatcherID != req.DispatcherId || call.TenantID != req.TenantId || call.TaskID != req.TaskId || call.SIPRevision != 8 || call.SelectedTrunkID != req.SelectedTrunkId { t.Fatal("Agent did not receive Dispatcher-authorized call binding") } - if call.AI.Mode != "asr_only" || call.AI.LLM != nil || call.AI.TTS != nil || call.AI.ASR.Provider.Credential != "example-only-not-a-real-secret" { + if call.AI.Mode != "asr_only" || call.AI.LLM != nil || call.AI.TTS != nil || call.AI.ASR.Provider.Credential != "sk-******" { t.Fatal("Agent did not receive the exact approved ASR-only provider snapshot") } return nil @@ -38,12 +38,12 @@ func TestApprovedAgentReceivesFrozenASROnlySnapshot(t *testing.T) { } } -func TestApprovedAgentRejectsUnsupportedAIWithoutPersistingAttempt(t *testing.T) { +func TestApprovedAgentRejectsProtocolOverrideWithoutPersistingAttempt(t *testing.T) { now := time.Date(2026, 9, 20, 10, 0, 0, 0, time.UTC) req := approvedTestRequest(t, now) - req.TaskConfigJson = bytes.Replace(req.TaskConfigJson, []byte(`"language":"zh-CN"`), []byte(`"language":"fr-FR"`), 1) - if !strings.Contains(string(req.TaskConfigJson), `"language":"fr-FR"`) { - t.Fatal("unsupported language was not inserted into fixture") + req.TaskConfigJson = bytes.Replace(req.TaskConfigJson, []byte(`"params":{}`), []byte(`"params":{"reqid":"override"}`), 1) + if !strings.Contains(string(req.TaskConfigJson), `"reqid":"override"`) { + t.Fatal("protocol override was not inserted into fixture") } hash, err := configread.ExecutionBindingSHA256(req.TaskConfigJson, req.ProvidersJson, req.SipRevision) if err != nil { diff --git a/internal/rpc/recording_server_flow_test.go b/internal/rpc/recording_server_flow_test.go index a901987..c0b77c3 100644 --- a/internal/rpc/recording_server_flow_test.go +++ b/internal/rpc/recording_server_flow_test.go @@ -54,7 +54,7 @@ func recordingRPCFixture(t *testing.T) (*RecordingServer, *store.Store, context. readExample("config-read-providers", &providers) snapshot.Providers = make(map[string]configread.Provider) for _, provider := range providers.Providers { - snapshot.Providers[provider.ProviderRef] = provider + snapshot.Providers[provider.Code] = provider } snapshot.SIP.Trunks = []byte(strings.Replace(string(snapshot.SIP.Trunks), `"max_concurrent_calls":null`, `"max_concurrent_calls":2`, 1)) if err := database.ApplyDiscoverySnapshot(recordingDispatcherID, []configread.DiscoveredTask{{TaskID: snapshot.Task.TaskID, TenantID: snapshot.Task.TenantID, TaskRevision: snapshot.Task.TaskRevision, Status: "running"}}); err != nil { @@ -72,7 +72,7 @@ func recordingRPCFixture(t *testing.T) (*RecordingServer, *store.Store, context. } now := time.Date(2026, 9, 21, 1, 30, 0, 0, time.UTC) if err := database.ReserveExecute(recordingDispatcherID, command.EventID, store.CallReservation{ - TrunkID: "trunk-mock", SIPRevision: 8, CallerID: "BD00000000", DialedCallee: command.Callee, Deadline: now.Add(2 * time.Minute), + TrunkID: "trunk-mock", SIPRevision: 8, TaskRevision: 1, CallerID: "BD00000000", DialedCallee: command.Callee, Deadline: now.Add(2 * time.Minute), }, now); err != nil { t.Fatal(err) } diff --git a/internal/rpc/task_trunk_concurrency_test.go b/internal/rpc/task_trunk_concurrency_test.go new file mode 100644 index 0000000..f8fe7aa --- /dev/null +++ b/internal/rpc/task_trunk_concurrency_test.go @@ -0,0 +1,56 @@ +package rpc + +import ( + "context" + "encoding/json" + "path/filepath" + "testing" + "time" + + "git.ipao.vip/rogee/go-sip/internal/configread" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +func TestExecuteApprovedRejectsInvalidOrDisabledTaskLineBeforeAttempt(t *testing.T) { + for _, tc := range []struct { + name, trunks string + code codes.Code + }{ + {"disabled selected line", `[{"trunk_id":"trunk-mock","concurrency":0}]`, codes.FailedPrecondition}, + {"unlisted selected line", `[{"trunk_id":"other-line","concurrency":5}]`, codes.FailedPrecondition}, + {"legacy format", `["trunk-mock"]`, codes.InvalidArgument}, + {"missing concurrency", `[{"trunk_id":"trunk-mock"}]`, codes.InvalidArgument}, + {"negative concurrency", `[{"trunk_id":"trunk-mock","concurrency":-1}]`, codes.InvalidArgument}, + {"duplicate line IDs", `[{"trunk_id":"trunk-mock","concurrency":1},{"trunk_id":"trunk-mock","concurrency":2}]`, codes.InvalidArgument}, + } { + t.Run(tc.name, func(t *testing.T) { + now := time.Date(2026, 9, 20, 10, 0, 0, 0, time.UTC) + req := approvedTestRequest(t, now) + var task map[string]json.RawMessage + if err := json.Unmarshal(req.TaskConfigJson, &task); err != nil { + t.Fatal(err) + } + task["allowed_trunk_ids"] = json.RawMessage(tc.trunks) + var err error + req.TaskConfigJson, err = json.Marshal(task) + if err != nil { + t.Fatal(err) + } + // Rebind the deliberately invalid task so rejection is about its + // structure/authorization, not a mismatched byte digest. + req.BindingSha256, err = configread.ExecutionBindingSHA256(req.TaskConfigJson, req.ProvidersJson, req.SipRevision) + if err != nil { + t.Fatal(err) + } + attempts := 0 + s := activatedApprovedServer(t, now, filepath.Join(t.TempDir(), "agent.json"), req.DispatcherId, 1, func(context.Context, ApprovedExecution) error { + attempts++ + return nil + }) + if _, err := s.ExecuteApproved(context.Background(), req); status.Code(err) != tc.code || attempts != 0 { + t.Fatalf("attempts=%d err=%v, want code=%v and no attempt", attempts, err, tc.code) + } + }) + } +} diff --git a/internal/store/calls.go b/internal/store/calls.go index 33de1ec..6a7bde4 100644 --- a/internal/store/calls.go +++ b/internal/store/calls.go @@ -37,6 +37,7 @@ type Execution struct { type CallReservation struct { TrunkID string SIPRevision int64 + TaskRevision int64 CallerID string DialedCallee string Deadline time.Time @@ -103,7 +104,7 @@ func (s *Store) RecordExecute(cmd ExecuteCommand) (Execution, bool, error) { // sending a single originate. A crash or timeout after this point leaves an // occupied dispatching/unknown execution, never an automatic redial. func (s *Store) ReserveExecute(dispatcherID, eventID string, selected CallReservation, at time.Time) error { - if dispatcherID == "" || eventID == "" || selected.TrunkID == "" || selected.SIPRevision <= 0 || selected.CallerID == "" || selected.DialedCallee == "" || !selected.Deadline.After(at) { + if dispatcherID == "" || eventID == "" || selected.TrunkID == "" || selected.SIPRevision <= 0 || selected.TaskRevision <= 0 || selected.CallerID == "" || selected.DialedCallee == "" || !selected.Deadline.After(at) { return errors.New("invalid originate reservation or expired deadline") } tx, err := s.db.Begin() @@ -154,9 +155,20 @@ func (s *Store) ReserveExecute(dispatcherID, eventID string, selected CallReserv if snapshot.Task.DispatcherID != dispatcherID || snapshot.Task.TenantID != cmd.TenantID || snapshot.Task.TaskID != cmd.TaskID || snapshot.SIP.Revision != revision || snapshot.Quota.TenantID != cmd.TenantID { return errors.New("admission snapshot identity mismatch") } + if snapshot.Task.TaskRevision != selected.TaskRevision { + return fmt.Errorf("%w: task revision changed after trunk selection", ErrNotReady) + } if snapshot.Task.MaxConcurrentCalls <= 0 || snapshot.Quota.MaxConcurrentCalls <= 0 { return fmt.Errorf("%w: missing quota", ErrNotReady) } + taskTrunkLimits, err := snapshot.Task.TrunkLimits() + if err != nil { + return fmt.Errorf("invalid task trunk limits: %w", err) + } + taskTrunkLimit, allowed := taskTrunkLimits[selected.TrunkID] + if !allowed || taskTrunkLimit == 0 { + return fmt.Errorf("%w: selected trunk is not enabled for this task", ErrNotReady) + } var trunks []struct { TrunkID string `json:"trunk_id"` Enabled bool `json:"enabled"` @@ -193,9 +205,16 @@ func (s *Store) ReserveExecute(dispatcherID, eventID string, selected CallReserv if err != nil { return fmt.Errorf("count trunk occupancy: %w", err) } + taskTrunkUsed, err := count(`SELECT COUNT(*) FROM dispatcher_inbox WHERE dispatcher_id=? AND tenant_id=? AND task_id=? AND selected_trunk_id=? AND `+occupied, dispatcherID, cmd.TenantID, cmd.TaskID, selected.TrunkID) + if err != nil { + return fmt.Errorf("count task trunk occupancy: %w", err) + } if tenantUsed >= snapshot.Quota.MaxConcurrentCalls || taskUsed >= snapshot.Task.MaxConcurrentCalls || trunkUsed >= trunkLimit { return ErrCapacity } + if taskTrunkUsed >= taskTrunkLimit { + return fmt.Errorf("%w: task trunk occupancy=%d limit=%d", ErrCapacity, taskTrunkUsed, taskTrunkLimit) + } result, err := tx.Exec(`UPDATE dispatcher_inbox SET status='dispatching',selected_trunk_id=?,caller_id=?,dialed_callee=?,deadline=?,snapshot_json=? WHERE dispatcher_id=? AND event_id=? AND status='pending'`, selected.TrunkID, selected.CallerID, selected.DialedCallee, selected.Deadline.UTC().Format(time.RFC3339Nano), body, dispatcherID, eventID) if err != nil { @@ -403,9 +422,22 @@ func (s *Store) PendingExecuteCount(dispatcherID string, tenantID int64, taskID // TrunkOccupancy includes dispatching and unknown executions. No timeout or // lease expiry releases a call whose real end has not been confirmed. func (s *Store) TrunkOccupancy(dispatcherID string) (map[string]int64, error) { - rows, err := s.db.Query(`SELECT selected_trunk_id, COUNT(*) FROM dispatcher_inbox + return s.trunkOccupancy(`SELECT selected_trunk_id, COUNT(*) FROM dispatcher_inbox WHERE dispatcher_id=? AND status IN ('dispatching','dispatched','unknown') AND selected_trunk_id IS NOT NULL GROUP BY selected_trunk_id`, dispatcherID) +} + +// TaskTrunkOccupancy counts only this task's calls on each line, retaining the +// same unknown-execution fence as the global line occupancy. +func (s *Store) TaskTrunkOccupancy(dispatcherID string, tenantID int64, taskID string) (map[string]int64, error) { + return s.trunkOccupancy(`SELECT selected_trunk_id, COUNT(*) FROM dispatcher_inbox + WHERE dispatcher_id=? AND tenant_id=? AND task_id=? + AND status IN ('dispatching','dispatched','unknown') + AND selected_trunk_id IS NOT NULL GROUP BY selected_trunk_id`, dispatcherID, tenantID, taskID) +} + +func (s *Store) trunkOccupancy(query string, args ...any) (map[string]int64, error) { + rows, err := s.db.Query(query, args...) if err != nil { return nil, fmt.Errorf("count occupied SIP trunks: %w", err) } diff --git a/internal/store/calls_test.go b/internal/store/calls_test.go index adaf31b..db0217c 100644 --- a/internal/store/calls_test.go +++ b/internal/store/calls_test.go @@ -35,7 +35,7 @@ func currentCall(id string) ExecuteCommand { return ExecuteCommand{DispatcherID: currentDispatcherID, EventID: id, TenantID: 1001, TaskID: "task-asr", Callee: "15003164745", IssuedAt: "2026-09-21T01:30:00Z"} } func currentReservation() CallReservation { - return CallReservation{TrunkID: "trunk-mock", SIPRevision: 8, CallerID: "BD00000000", DialedCallee: "15003164745", Deadline: time.Date(2026, 9, 21, 1, 32, 0, 0, time.UTC)} + return CallReservation{TrunkID: "trunk-mock", SIPRevision: 8, TaskRevision: 1, CallerID: "BD00000000", DialedCallee: "15003164745", Deadline: time.Date(2026, 9, 21, 1, 32, 0, 0, time.UTC)} } func TestExecuteInboxPreventsDuplicateOriginationAfterRestart(t *testing.T) { diff --git a/internal/store/snapshot_test.go b/internal/store/snapshot_test.go index 6462d9e..f0ed7a5 100644 --- a/internal/store/snapshot_test.go +++ b/internal/store/snapshot_test.go @@ -28,7 +28,7 @@ func TestStoreSnapshotPreservesAgentAndAllowsNewSIPRevision(t *testing.T) { if err != nil { t.Fatal(err) } - if loaded.SIP.Revision != 8 || loaded.Quota.MaxConcurrentCalls != 3 || loaded.Providers["asr-example"].Credential != "example-only-not-a-real-secret" || !strings.Contains(string(loaded.Task.Agent.Raw), `"interim":false`) { + if loaded.SIP.Revision != 8 || loaded.Quota.MaxConcurrentCalls != 3 || loaded.Providers["volcengine"].Credential != "sk-******" || !strings.Contains(string(loaded.Task.Agent.Raw), `"params":{}`) { t.Fatalf("binding was lost across restart: SIP=%d quota=%d agent=%s", loaded.SIP.Revision, loaded.Quota.MaxConcurrentCalls, loaded.Task.Agent.Raw) } snapshot.SIP.Revision = 9 diff --git a/internal/store/store.go b/internal/store/store.go index 47b2aa5..3dfbb1d 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -403,7 +403,7 @@ func saveSnapshot(tx *sql.Tx, snapshot configread.Snapshot) error { continue } provider, ok := snapshot.Providers[ref] - if !ok || provider.ProviderRef != ref || provider.Credential == "" { + if !ok || provider.Code != ref || provider.Credential == "" { return fmt.Errorf("task provider %q is unavailable", ref) } providers[ref] = provider diff --git a/internal/store/store_test.go b/internal/store/store_test.go index b9245d5..2f13aae 100644 --- a/internal/store/store_test.go +++ b/internal/store/store_test.go @@ -36,7 +36,7 @@ func currentStoreSnapshot(t *testing.T) configread.Snapshot { read("config-read-providers", &p) s.Providers = make(map[string]configread.Provider) for _, provider := range p.Providers { - s.Providers[provider.ProviderRef] = provider + s.Providers[provider.Code] = provider } return s } diff --git a/internal/store/task_trunk_concurrency_test.go b/internal/store/task_trunk_concurrency_test.go new file mode 100644 index 0000000..b247299 --- /dev/null +++ b/internal/store/task_trunk_concurrency_test.go @@ -0,0 +1,195 @@ +package store + +import ( + "encoding/json" + "errors" + "fmt" + "path/filepath" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "git.ipao.vip/rogee/go-sip/internal/configread" +) + +func trunkLimitSnapshot(t *testing.T, taskID string, revision, limit int64) configread.Snapshot { + t.Helper() + snapshot := currentStoreSnapshot(t) + var task map[string]json.RawMessage + if err := json.Unmarshal(snapshot.Task.Raw, &task); err != nil { + t.Fatal(err) + } + set := func(key string, value any) { + raw, err := json.Marshal(value) + if err != nil { + t.Fatal(err) + } + task[key] = raw + } + set("task_id", taskID) + set("task_revision", revision) + set("max_concurrent_calls", 100) + set("allowed_trunk_ids", []configread.AllowedTrunk{{TrunkID: "trunk-mock", Concurrency: limit}}) + raw, err := json.Marshal(task) + if err != nil { + t.Fatal(err) + } + if err := json.Unmarshal(raw, &snapshot.Task); err != nil { + t.Fatal(err) + } + snapshot.SIP.Trunks = []byte(strings.Replace(string(snapshot.SIP.Trunks), `"max_concurrent_calls":null`, `"max_concurrent_calls":100`, 1)) + snapshot.Quota.MaxConcurrentCalls = 100 + return snapshot +} + +func taskTrunkStore(t *testing.T, snapshots ...configread.Snapshot) (*Store, string) { + t.Helper() + path := filepath.Join(t.TempDir(), "state.db") + s, err := Open(path) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = s.Close() }) + tasks := make([]configread.DiscoveredTask, 0, len(snapshots)) + for _, snapshot := range snapshots { + tasks = append(tasks, configread.DiscoveredTask{TaskID: snapshot.Task.TaskID, TenantID: snapshot.Task.TenantID, TaskRevision: snapshot.Task.TaskRevision, Status: "running"}) + } + if err := s.ApplyDiscoverySnapshot(currentDispatcherID, tasks); err != nil { + t.Fatal(err) + } + for _, snapshot := range snapshots { + if err := s.SaveSnapshot(snapshot); err != nil { + t.Fatal(err) + } + } + if err := s.MarkReadyForSIP(currentDispatcherID, 8); err != nil { + t.Fatal(err) + } + return s, path +} + +func recordTrunkCall(t *testing.T, s *Store, taskID, eventID string) ExecuteCommand { + t.Helper() + cmd := currentCall(eventID) + cmd.TaskID = taskID + if _, created, err := s.RecordExecute(cmd); err != nil || !created { + t.Fatalf("record call: created=%t err=%v", created, err) + } + return cmd +} + +func TestTaskTrunkReservationRejectsDisabledAndUnlistedLines(t *testing.T) { + for _, tc := range []struct { + name, selected string + limit int64 + }{ + {"zero", "trunk-mock", 0}, + {"unlisted", "another-trunk", 5}, + } { + t.Run(tc.name, func(t *testing.T) { + s, _ := taskTrunkStore(t, trunkLimitSnapshot(t, "task-asr", 1, tc.limit)) + cmd := recordTrunkCall(t, s, "task-asr", "call-disabled") + selection := currentReservation() + selection.TrunkID = tc.selected + at := time.Date(2026, 9, 21, 1, 30, 0, 0, time.UTC) + if err := s.ReserveExecute(cmd.DispatcherID, cmd.EventID, selection, at); !errors.Is(err, ErrNotReady) { + t.Fatalf("disabled/unlisted task line reserved: %v", err) + } + }) + } +} + +func TestTaskTrunkLimitIsScopedToTaskAndUnknownSurvivesRestart(t *testing.T) { + s, path := taskTrunkStore(t, trunkLimitSnapshot(t, "task-asr", 1, 1), trunkLimitSnapshot(t, "task-other", 1, 1)) + at := time.Date(2026, 9, 21, 1, 30, 0, 0, time.UTC) + first := recordTrunkCall(t, s, "task-asr", "call-first") + if err := s.ReserveExecute(first.DispatcherID, first.EventID, currentReservation(), at); err != nil { + t.Fatal(err) + } + other := recordTrunkCall(t, s, "task-other", "call-other") + if err := s.ReserveExecute(other.DispatcherID, other.EventID, currentReservation(), at); err != nil { + t.Fatalf("another task lost its independent per-line capacity: %v", err) + } + if err := s.MarkExecuteUnknown(first.DispatcherID, first.EventID); err != nil { + t.Fatal(err) + } + if err := s.Close(); err != nil { + t.Fatal(err) + } + s, err := Open(path) + if err != nil { + t.Fatal(err) + } + defer s.Close() + if err := s.MarkReadyForSIP(currentDispatcherID, 8); err != nil { + t.Fatal(err) + } + occupancy, err := s.TaskTrunkOccupancy(currentDispatcherID, 1001, "task-asr") + if err != nil || occupancy["trunk-mock"] != 1 { + t.Fatalf("task trunk occupancy after restart=%v err=%v", occupancy, err) + } + second := recordTrunkCall(t, s, "task-asr", "call-second") + if err := s.ReserveExecute(second.DispatcherID, second.EventID, currentReservation(), at); !errors.Is(err, ErrCapacity) { + t.Fatalf("unknown call released task-line capacity on restart: %v", err) + } + if err := s.FinishExecute(first.DispatcherID, first.EventID); err != nil { + t.Fatal(err) + } + if err := s.ReserveExecute(second.DispatcherID, second.EventID, currentReservation(), at); err != nil { + t.Fatalf("confirmed end did not release task-line capacity: %v", err) + } +} + +func TestTaskTrunkConcurrentReservationsCannotExceedLimit(t *testing.T) { + s, _ := taskTrunkStore(t, trunkLimitSnapshot(t, "task-asr", 1, 5)) + at := time.Date(2026, 9, 21, 1, 30, 0, 0, time.UTC) + commands := make([]ExecuteCommand, 20) + for i := range commands { + commands[i] = recordTrunkCall(t, s, "task-asr", fmt.Sprintf("call-%d", i)) + } + var accepted atomic.Int64 + var wg sync.WaitGroup + for _, cmd := range commands { + wg.Go(func() { + err := s.ReserveExecute(cmd.DispatcherID, cmd.EventID, currentReservation(), at) + if err == nil { + accepted.Add(1) + } else if !errors.Is(err, ErrCapacity) { + t.Errorf("reserve: %v", err) + } + }) + } + wg.Wait() + if accepted.Load() != 5 { + t.Fatalf("accepted=%d, want task-line limit 5", accepted.Load()) + } +} + +func TestTaskTrunkEditRechecksLimitAndSelectionRevision(t *testing.T) { + s, _ := taskTrunkStore(t, trunkLimitSnapshot(t, "task-asr", 1, 2)) + at := time.Date(2026, 9, 21, 1, 30, 0, 0, time.UTC) + first := recordTrunkCall(t, s, "task-asr", "call-before-edit") + if err := s.ReserveExecute(first.DispatcherID, first.EventID, currentReservation(), at); err != nil { + t.Fatal(err) + } + if state, err := s.CompleteEdit(trunkLimitSnapshot(t, "task-asr", 2, 1), "edit-line-limit"); err != nil || state != "applied" { + t.Fatalf("edit=%q err=%v", state, err) + } + second := recordTrunkCall(t, s, "task-asr", "call-after-edit") + if err := s.ReserveExecute(second.DispatcherID, second.EventID, currentReservation(), at); !errors.Is(err, ErrNotReady) { + t.Fatalf("old selection revision was accepted after edit: %v", err) + } + selection := currentReservation() + selection.TaskRevision = 2 + if err := s.ReserveExecute(second.DispatcherID, second.EventID, selection, at); !errors.Is(err, ErrCapacity) { + t.Fatalf("new lower task-line limit was ignored: %v", err) + } + if err := s.FinishExecute(first.DispatcherID, first.EventID); err != nil { + t.Fatal(err) + } + if err := s.ReserveExecute(second.DispatcherID, second.EventID, selection, at); err != nil { + t.Fatalf("capacity did not reopen after confirmed end: %v", err) + } +} diff --git a/scripts/check-current-contracts.py b/scripts/check-current-contracts.py index 2d8c3a2..90af42b 100644 --- a/scripts/check-current-contracts.py +++ b/scripts/check-current-contracts.py @@ -11,7 +11,7 @@ def git(root, *args): return subprocess.check_output(['git', '-C', str(root), *args], text=True, stderr=subprocess.STDOUT).strip() -def check_checkout(root): +def check_checkout(root, updating=False): if (root / 'contracts/local').exists(): raise ValueError('contracts/local is a forbidden parallel contract source') entry = git(root, 'ls-files', '--stage', '--', 'contracts/schema').split() @@ -25,7 +25,9 @@ def check_checkout(root): raise ValueError('contract submodule origin must be ' + URL) commit = git(sub, 'rev-parse', 'HEAD') if commit != entry[1]: - raise ValueError('contract HEAD does not match the project pin; stage the verified contracts/schema pointer') + if not updating: + raise ValueError('contract HEAD does not match the project pin; stage the verified contracts/schema pointer') + git(sub, 'merge-base', '--is-ancestor', entry[1], commit) if git(sub, 'status', '--porcelain', '--untracked-files=all'): raise ValueError('contract submodule has uncommitted changes; verify and commit them in the shared repository') return commit @@ -39,7 +41,7 @@ def main(): checker = importlib.util.module_from_spec(spec) spec.loader.exec_module(checker) count = checker.verify_bundle(checker_path.parent) - print(f'shared contract {commit}: {count} JSON files; source/bundle hashes and offline references valid') + print(f'shared contract {commit}: {count} JSON files; offline references and historical provenance valid') if __name__ == '__main__': diff --git a/scripts/coverage.sh b/scripts/coverage.sh index d1c30e1..a5a7893 100755 --- a/scripts/coverage.sh +++ b/scripts/coverage.sh @@ -4,6 +4,9 @@ set -eu profile=${COVER_PROFILE:-/tmp/go-sip-coverage.out} business_profile=${COVER_BUSINESS_PROFILE:-/tmp/go-sip-coverage-business.out} -GOMAXPROCS=${GOMAXPROCS:-2} go test -p 1 -coverpkg=./... -coverprofile="$profile" ./... +# Coverage patterns also match local module replacements. Instrument only +# this module's explicit package list, not the checked-in third-party SDK. +packages=$(go list ./... | paste -sd, -) +GOMAXPROCS=${GOMAXPROCS:-2} go test -p 1 -coverpkg="$packages" -coverprofile="$profile" ./... awk 'NR == 1 || index($0, "/gen/") == 0' "$profile" > "$business_profile" go tool cover -func="$business_profile" diff --git a/scripts/test_contract_checkout.py b/scripts/test_contract_checkout.py index 3597c78..d1503c4 100644 --- a/scripts/test_contract_checkout.py +++ b/scripts/test_contract_checkout.py @@ -49,6 +49,22 @@ class ContractCheckoutTest(unittest.TestCase): with self.assertRaisesRegex(ValueError, 'does not match'): checker.check_checkout(self.root) + def test_update_rechecks_an_unstaged_fast_forward(self): + (self.sub / 'test.schema.json').write_text('{"changed": true}\n') + self.git(self.sub, 'add', '.') + self.git(self.sub, 'commit', '-qm', 'changed') + self.assertEqual(checker.check_checkout(self.root, updating=True), self.git(self.sub, 'rev-parse', 'HEAD')) + (self.sub / 'untracked.txt').write_text('unexpected\n') + with self.assertRaisesRegex(ValueError, 'uncommitted'): + checker.check_checkout(self.root, updating=True) + + def test_update_rejects_divergent_contract_history(self): + self.git(self.sub, 'checkout', '--orphan', 'unrelated') + self.git(self.sub, 'add', '.') + self.git(self.sub, 'commit', '-qm', 'unrelated root') + with self.assertRaises(subprocess.CalledProcessError): + checker.check_checkout(self.root, updating=True) + def test_dirty_submodule_is_rejected(self): (self.sub / 'untracked.txt').write_text('unexpected\n') with self.assertRaisesRegex(ValueError, 'uncommitted'): diff --git a/scripts/test_coverage.py b/scripts/test_coverage.py new file mode 100644 index 0000000..b81adf0 --- /dev/null +++ b/scripts/test_coverage.py @@ -0,0 +1,41 @@ +#!/usr/bin/env python3 +"""Coverage must measure application packages, not local replacement modules.""" +import os +from pathlib import Path +import subprocess +import tempfile +import unittest + + +class CoverageBoundaryTests(unittest.TestCase): + def test_explicit_project_packages_and_generated_code_exclusion(self): + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + go = root / 'go' + go.write_text('''#!/bin/sh +set -eu +case "$1" in +list) printf 'example/app/internal/ai\\nexample/app/gen/agent\\n' ;; +test) + printf '%s\\n' "$@" > "$COMMANDS" + printf 'mode: set\\nexample/app/internal/ai/a.go:1.1,2.2 1 1\\nexample/app/gen/agent/a.go:1.1,2.2 1 0\\n' > "$COVER_PROFILE" + ;; +tool) exit 0 ;; +*) exit 1 ;; +esac +''') + go.chmod(0o700) + env = dict(os.environ, PATH=str(root) + os.pathsep + os.environ['PATH'], + COMMANDS=str(root / 'commands'), COVER_PROFILE=str(root / 'all.out'), + COVER_BUSINESS_PROFILE=str(root / 'business.out')) + subprocess.run(['sh', str(Path(__file__).with_name('coverage.sh'))], env=env, check=True) + args = (root / 'commands').read_text() + self.assertIn('-coverpkg=example/app/internal/ai,example/app/gen/agent', args) + self.assertNotIn('-coverpkg=./...', args) + report = (root / 'business.out').read_text() + self.assertIn('/internal/ai/', report) + self.assertNotIn('/gen/', report) + + +if __name__ == '__main__': + unittest.main() diff --git a/scripts/update-contracts.sh b/scripts/update-contracts.sh index aef0493..0fc358c 100755 --- a/scripts/update-contracts.sh +++ b/scripts/update-contracts.sh @@ -6,7 +6,9 @@ sub=contracts/schema if [[ ! -e "$sub/.git" ]]; then git submodule update --init -- "$sub" fi -python3 scripts/check-current-contracts.py +# A previously fetched fast-forward may be unstaged after a failed check. +# Still reject dirty submodule files, wrong origins and divergent history. +python3 -c 'from pathlib import Path; from runpy import run_path; run_path("scripts/check-current-contracts.py")["check_checkout"](Path.cwd(), updating=True)' git -C "$sub" fetch origin main git -C "$sub" merge --ff-only FETCH_HEAD if [[ "$(git -C "$sub" rev-parse HEAD)" != "$(git -C "$sub" rev-parse FETCH_HEAD)" ]]; then