diff --git a/docs/docs/index.md b/docs/docs/index.md index b87d3a53..4ea01b9a 100644 --- a/docs/docs/index.md +++ b/docs/docs/index.md @@ -30,8 +30,8 @@ 限流、幂等和服务生命周期契约。 - [Telegram 长轮询 Adapter](telegram.md):Issue #31 的文档先行契约,固定单 Binding、Bot 身份校验、普通文本映射、Dispatch 聚合回复和生命周期边界。 -- [企业微信自建应用 Text Webhook](wecom.md):Issue #60 的文档先行契约,固定 callback - 验签/AES 解密、可信 Binding 路由、文本入站和可靠回复边界。 +- [企业微信自建应用 Channel Adapter](wecom.md):Issue #60/#98 的契约,固定 callback + 验签/AES 解密、可信 Binding 路由、文本/媒体入站和可靠回复边界。 - [Telegram live E2E 示例](https://github.com/XnLemon/trpc-agent-service/tree/main/examples/telegram-e2e): Issue #33 的真实 Bot API 传输冒烟测试和手动 CI 运行说明。 - [PostgreSQL 控制面与启动装配](postgresql-control-plane.md):Issue #37 的实现契约, @@ -76,7 +76,7 @@ cd trpc-agent-service - [数据模型](data-model.md) — 核心表结构、Session/Event/Memory/Summary/Audit 和租户约束 - [Channel Binding](channel-binding.md) — 租户级通道绑定、候选发现与可信入站路由 - [Telegram 长轮询 Adapter](telegram.md) — 单 Binding Telegram long polling、文本映射与安全边界 -- [企业微信自建应用 Text Webhook](wecom.md) — 自建应用 callback、文本入站与回复 Outbox +- [企业微信自建应用 Channel Adapter](wecom.md) — 自建应用 callback、文本/媒体入站与回复 Outbox - [Gateway、Execution Plan 与 HTTP/SSE](gateway.md) — 可信主体、固定执行计划、Runner Registry、 Dispatch、健康检查、优雅停机和普通/流式 API - [PostgreSQL 控制面与启动装配](postgresql-control-plane.md) — 六类控制面表的 migration 顺序、 diff --git a/docs/docs/issue-98-native-media.md b/docs/docs/issue-98-native-media.md new file mode 100644 index 00000000..7d1913ba --- /dev/null +++ b/docs/docs/issue-98-native-media.md @@ -0,0 +1,83 @@ +# Issue #98:原生媒体附件与富回复 + +Issue #98 在 Issue #77 的通道能力基础上补齐协议中立附件契约。目标不是让通道 payload +绕过 Gateway,而是在验证后的 channel 边界内下载、限流、校验并持久化媒体,然后只把 +安全的 `attachment.Reference` 和可加载的 `ContentPart` 交给 Runner。 + +## MVP 范围 + +- 入站附件类型固定为 `image`、`video`、`audio`、`document`,引用包含 ID、MIME、 + 大小、SHA-256、原始文件名和 provider 文件身份。 +- Telegram 入站会保留图片、文档、音频和视频的 provider file id,使用受控 + `MediaDownloader` 下载后立即转存到租户隔离 attachment store。 +- WeCom 入站会在签名、AES 解密、`ReceiveID` 和 `AgentID` 全部通过后处理 + `image`、`file`、`voice`、`video`;未知类型和缺少 `MediaId` 的回调 fail closed。 +- Gateway 将附件绑定到 durable `message_event` 后,Runner 才能通过 tenant/event/reference + 读取内容;Runner 不接触 Telegram 下载链接、WeCom `media_id` 下载 URL、access token 或 + channel secret。 +- 出站 Outbox 支持结构化 `kind + attachment ref + fallback`;Telegram 和 WeCom 对 + `image`、`document` 走原生发送,`audio`、`video` 或不支持能力时使用确定性文本 fallback。 + +## 入站生命周期 + +```text +channel callback/update + -> 验签、解密、绑定候选校验 + -> 识别 provider media/file id + -> 受控 downloader 下载,限制大小和响应形态 + -> attachment store 写入 bytes + metadata + -> Gateway 记录 message_event 并绑定 attachment + -> Runner 按模型能力读取 ContentParts +``` + +附件 ID 使用 `(tenant, binding, external_message_id, ordinal/provider_id)` 语义生成,保证 +重复回调和重试不会产生新的逻辑附件。attachment store 必须返回 defensive copy,并按 +tenant/event/reference 验证读取;未绑定到 durable event 的附件不能被 Runner 加载。 + +## 模型边界 + +`trpc-agent-go` 已支持 `ContentParts`,本仓库的 Responses 适配器按模型 profile 的显式输入 +能力转换: + +| 附件 | Runner 内容 | OpenAI Responses 映射 | +| --- | --- | --- | +| text | `ContentPart{Type:text}` | `input_text` | +| image | image content part | `input_image` | +| document/file | file content part | `input_file` | +| mp3/wav audio | audio content part | `input_audio` | +| video | 保留附件和 fallback 文本 | 不声明视频理解 | + +视频可以安全收发和存储,但当前不承诺模型视频理解。抽帧、OCR、ASR 和视频分析是后续独立 +能力,不属于 #98 MVP。 + +## 出站语义 + +`ReplyOutbox` 的媒体段是不可变结构:`Kind` 表示协议中立类型,`Attachment` 指向已验证内容, +`Payload` 可作为 caption,`Fallback` 是目标通道无法表达原生媒体时必须使用的确定性文本。 + +- Telegram 使用 attachment reader 构造 SDK upload,图片走 `SendPhoto`,文档走 + `SendDocument`;视频和音频先保守发送 fallback。 +- WeCom 图片和文档先上传临时素材,再分别发送 `image` 或 `file` 应用消息;未配置 reader 或 + 不支持的 kind 发送 fallback。 +- Provider 成功 receipt 继续写回 Outbox 状态机;token、provider URL、原始响应和消息正文不进入 + 日志或审计。 + +## 存储与上线 + +当前 in-memory 和 PostgreSQL runtime store 都实现了 attachment store。PostgreSQL 版本复用 +现有 object boundary,但二进制内容仍落在数据库内,适合受限 MVP 和 deterministic E2E,不适合 +长期生产视频流量。生产视频或大文件上线前应补 S3/COS 一类流式对象存储实现,并明确: + +- 每租户数量、单文件大小、总容量和 MIME allow-list; +- 保留期、引用计数、清理任务和 dead-letter 后的处置; +- downloader 超时、取消、重试和限速; +- 跨 tenant/binding 的读取拒绝和 provider secret 脱敏; +- 部署文档中的对象存储 endpoint、凭据轮换和迁移策略。 + +## 验证证据 + +- Telegram 入站、原生图片/文档出站、fallback 和 cancellation 使用 fake SDK/reader 测试。 +- WeCom 入站媒体回调使用加密 XML 和 fake downloader 测试,证明不使用 `PicUrl`。 +- WeCom 出站图片/文档使用 `httptest` 验证临时素材上传和发送 payload。 +- Gateway、in-memory/PostgreSQL runtime store 和 migrations 覆盖 tenant/event/reference 绑定、 + 幂等和清理。 diff --git a/docs/docs/ops.md b/docs/docs/ops.md index f56fc034..b49e6c4e 100644 --- a/docs/docs/ops.md +++ b/docs/docs/ops.md @@ -1,8 +1,29 @@ # 运维、可观测性与生产风险 > 本页把 [生产架构设计](architecture.md) 转成可执行的发布、监控、恢复和风险检查表。 -> 当前仓库只有控制面领域模型、快照和最小 Runner spine;Gateway、队列、真实 IM/Storage -> Adapter、Dashboard 和告警规则仍是后续平台实现,不应把本页当作已经部署的运行手册。 +> 本页同时记录已在仓库实现的运行时能力和仍待生产化的能力。代码、测试或示例可验证 +> 不代表相应的基础设施已经按本页要求部署并通过容量或灾备验收。 + +## 当前实现状态 + +下表区分“有领域契约或设计”与“默认运行链路已经可用”。发布和答辩应以此表为准,避免把 +规划中的架构能力表述为已完成的生产部署。 + +| 能力 | 状态 | 当前边界与验证 | 生产化缺口或后续门禁 | +| --- | --- | --- | --- | +| Gateway、Execution Plan、HTTP Chat/SSE 与 Agent Worker | 已实现 | 可信租户路由、版本化执行快照、限流和幂等均在服务运行链路中 | 仍需按部署规模验证跨节点容量和 SLO | +| 控制面与运行时存储 | 已实现 | PostgreSQL 控制面;InMemory 和 PostgreSQL Runtime Store 覆盖 Session/Event、Memory、Knowledge、Vector、Artifact 和 Reply Outbox | PostgreSQL 向量/知识检索尚未替代为专用索引服务;InMemory 仅适合单进程开发/测试 | +| 可靠回复投递 | 已实现 | Reply Outbox 有租约、fencing、指数退避、dead-letter 与 reconcile 测试 | 尚未形成通用的供应商 `Retry-After` 解析和按 chat/provider 限频队列 | +| Telegram 与企业微信通道 | 已实现,默认装配范围有限 | 两个通道已有入站处理、附件/媒体 MVP、图片和文档出站以及 fallback;企业微信环境 bootstrap 仍是单个静态身份配置 | 需要按目标供应商补齐更丰富交互、平台级多账号装配和限频/429 演练 | +| Trace、metrics、Prometheus 与审计 | 已实现 | OpenTelemetry、Prometheus 导出、脱敏字段和审计事件均有代码、文档或示例验证 | 告警阈值、保留策略和真实 exporter 背压容量尚未完成生产验收 | +| 灰度和回滚 | 部分实现 | Tenant/App 级 canary revision 与版本化回滚指针已可用 | 不含百分比 rollout、稳定分桶、指标自动止损或自动回滚 | +| 故障注入 | 部分实现 | deterministic E2E 覆盖 Gateway、Worker、Outbox lease/retry/restart 和并发边界 | 未完成真实数据库故障转移、OTel 背压、IM 429/Retry-After 与模型超时的演练 | +| Redis、对象存储与独立向量库 | 未实现 | Backend Profile 和 Runtime Capability 已定义可路由边界 | 需要真实 provider、租户装配、迁移、隔离契约和运行验证;当前优先跟踪 Redis Runtime Store | +| Admin 身份与工具治理 | 部分实现 | Admin 使用静态 Bearer Token;工具支持 allow/deny/approval 策略,运行时记录调用上限和审计数据 | 需要 OIDC/JWT、RBAC、token rotation、参数级策略、主体授权和可扣减预算 | +| 容量、备份与灾备 | 未实现 | 已有容量模型、部署示例和故障注入测试 | 需要可重复压测、备份恢复、容量基线和真实基础设施演练 | + +本页的 runbook 描述目标运行约束;当表中能力尚未实现时,它是后续实现的验收条件,而不是 +已经可由默认环境保证的行为。 ## 运行边界与值班目标 @@ -207,10 +228,14 @@ backpressure。高峰保护使用租户级 token bucket、全局队列上限、 | 回复重试风暴 | IM 429/5xx、固定间隔重试、无 per-chat 限速 | 供应商封禁、用户刷屏、队列雪崩 | retry multiplier、429、DLQ、outbox age | 指数退避+jitter、解析 Retry-After、按通道/chat 分桶、最大预算和 DLQ | | goroutine/事件泄漏 | context 未传递、Runner Event channel 未排空、consumer 无关闭边界 | Worker 内存上涨、滚动发布卡住、重复消费 | goroutine、FD、channel backlog、shutdown duration 持续上升 | owner 明确;context deadline;有界 drain;supervisor/health check;超时交给幂等重投递 | -## 当前实现状态与后续门禁 +## 后续实现门禁 -本仓库目前可以验证 Tenant、Agent App/Revision、Model Profile、Backend Profile、无密钥 -Execution Plan、Runner policy 和 Tenant-scoped Session 的模型/边界测试;不能验证真实 IM -验签、跨节点 CAS、队列至少一次投递、SQL/Redis 迁移或生产告警。后续实现每落地一个 Adapter -都必须补充:双租户隔离测试、重复/乱序/验签失败测试、provider 一致性契约测试、故障注入、 -审计字段检查和 `mkdocs build --strict`。 +当前仓库已经覆盖 Tenant、Agent App/Revision、Model Profile、Backend Profile、无密钥 +Execution Plan、Runner policy、Tenant-scoped PostgreSQL/内存 Runtime Store、Telegram/企业微信 +入站验证与可靠回复投递的模型、边界或 E2E 测试。它们不替代跨节点容量、真实基础设施故障 +恢复或外部后端迁移的生产验收。 + +后续每落地一个外部 Adapter 或生产运行能力,都必须补充:双租户隔离测试、重复/乱序/验签 +失败测试、provider 一致性契约测试、目标基础设施故障注入、审计字段检查和 +`mkdocs build --strict`。涉及数据后端的变更还必须证明 migration、shadow read、回滚窗口和 +provider 不可用时的失败语义。 diff --git a/docs/docs/telegram.md b/docs/docs/telegram.md index a72e234b..76b00d39 100644 --- a/docs/docs/telegram.md +++ b/docs/docs/telegram.md @@ -1,6 +1,6 @@ # Telegram 长轮询 Channel Adapter -> Issue #31 的基础实现;Issue #77 在此之上补充 webhook、媒体/rich update 与统一回复渲染。 +> Issue #31 的基础实现;Issue #77 在此之上补充 webhook、媒体/rich update 与统一回复渲染,Issue #98 补齐受控附件入站和图片/文档原生出站。 ## 1. 交付边界 @@ -9,13 +9,13 @@ Telegram 适配器是一个绑定级别的协议入口,不创建第二套租 ```text Telegram getUpdates - -> Update.Message 校验和规范化 + -> Update.Message 校验和规范化为文本或媒体附件 -> 已验证的 channels.RoutingTarget -> gateway.Channel Principal -> gateway.DispatchService -> 完整消费脱敏 DispatchEvent -> 聚合并分段 - -> Telegram sendMessage + -> Telegram sendMessage / SendPhoto / SendDocument ``` long polling 与 webhook 共用同一个 Adapter、幂等和 Gateway 边界;命令、回调仍 fail closed。 @@ -38,6 +38,9 @@ Webhook 由调用方拥有 HTTP listener,`telegram.Webhook` 只负责精确 pa | `Workers` | 零值为 1;大于 1 必须由调用方显式配置,并由 SDK 同步处理 handler 生命周期 | | `ErrorHook` | 只接收稳定的适配器错误类别,不接收 SDK/provider 原始错误、token 或 endpoint 凭据 | | `Factory` | 注入式 Bot factory;生产实现才依赖 `github.com/go-telegram/bot`,测试不创建网络 client | +| `Attachments` | 可选 runtime attachment store;配置后媒体 bytes 在进入 Runner 前先按 tenant/event/reference 持久化 | +| `MediaDownloader` | 可选受控下载器;未配置时生产 SDK client 会在具备 `GetFile` 能力时创建默认下载器 | +| `MaxAttachmentBytes` | 单附件大小上限,零值使用协议中立默认限制 | 构造函数接收 Context,先创建带默认 update handler 的 client,再调用 `getMe`,把返回的 Bot user ID 规范化为十进制字符串并与 `Target.ProviderAccountID` 精确比较。创建失败、`getMe` 失败或身份 @@ -58,7 +61,7 @@ Bot factory 对 SDK 使用以下固定策略: ## 3. 入站规范化与幂等 -第一版只接受 `Update.Message` 中的普通文本: +基础文本字段仍按以下规则规范化: | Telegram 字段 | Gateway 字段 | 规则 | | --- | --- | --- | @@ -69,8 +72,12 @@ Bot factory 对 SDK 使用以下固定策略: | `Message.MessageThreadID` | `ExternalThreadID` | 大于零时保留;发送回复时原样作为 forum thread | | `Update.ID` + trusted `BindingID` | `ExternalMessageID` / `RequestID` | 使用长度前缀编码生成稳定、无碰撞的 binding-aware ID | -编辑消息、channel post、callback/inline、service update、无 sender/chat/text、未知 chat 类型和 -媒体-only update 都以稳定的非敏感原因忽略或拒绝,且不得进入 Dispatch。所有合法消息先用固定 +Issue #98 之后,图片、文档、音频和视频会保留 provider file id,经受控 downloader 下载、限流、 +校验并写入 attachment store;没有 attachment store 时仍保留兼容的 caption 或 `[telegram media]` +文本标记,不把 provider 下载 URL 或 token 传入 Runner。 + +编辑消息、channel post、callback/inline、service update、无 sender/chat、未知 chat 类型和 +不支持的 rich/service update 都以稳定的非敏感原因忽略或拒绝,且不得进入 Dispatch。所有合法消息先用固定 principal 调用 `IdempotencyStore.Begin`: - pending duplicate 不再次调用 Dispatch,也不启动隐藏 retry; @@ -141,8 +148,8 @@ README 和 MkDocs 状态应明确区分已交付与后续能力: - [x] `getMe` 身份校验、tenant/Binding/Runner 隔离、普通文本映射和 binding-aware 幂等通过测试; - [x] Dispatch 完整消费、单逻辑回复、4096 code point 分段、forum thread 路由和失败脱敏通过测试; - [x] cancellation、polling error、send failure、duplicate delivery 和资源生命周期通过测试; -- [x] Telegram long polling 已实现;Webhook、持久化幂等/outbox、媒体、跨节点 ownership - 和其他 rich update 明确保持未勾选。 +- [x] Telegram long polling、webhook、持久化 outbox、媒体附件入站、图片/文档原生出站和 fallback + 已实现;其他 rich update、音频/视频原生出站和视频理解保持非目标或后续能力。 参考:[Telegram Bot API](https://core.telegram.org/bots/api)、 [getUpdates](https://core.telegram.org/bots/api#getting-updates)、 @@ -160,7 +167,7 @@ trace 或错误。CI 使用受保护的 `telegram-e2e` Environment,至少配 `TELEGRAM_SENDER_BOT_TOKEN`。一个 Bot Token 不能模拟普通用户向自己发送入站消息, 所以当前 workflow 必须显式配置第二个受控测试 Bot;本地人工运行可以不配置发送者。 -示例和 CI 都只验证普通文本;命令、媒体、rich update、Webhook、持久化 outbox 和 -生产模型供应商仍不属于该 E2E 范围。详见 +示例和 CI 都只验证普通文本;命令、媒体和 rich update 不属于该 live E2E 范围,媒体行为由 +deterministic fake 测试覆盖。详见 [Telegram live E2E example](https://github.com/XnLemon/trpc-agent-service/tree/main/examples/telegram-e2e) 和 Issue #33。 diff --git a/docs/docs/wecom.md b/docs/docs/wecom.md index 4e580429..073e3f09 100644 --- a/docs/docs/wecom.md +++ b/docs/docs/wecom.md @@ -1,7 +1,7 @@ -# 企业微信自建应用 Text Webhook +# 企业微信自建应用 Channel Adapter > 本页是 Issue #60 的基础契约;Issue #77 扩展了多账号 registry、有界 worker、群聊 -> delivery/receipt,以及独立的 Public WeChat/微信客服 provider boundary。 +> delivery/receipt,以及独立的 Public WeChat/微信客服 provider boundary。Issue #98 增加受控媒体附件入站和图片/文件原生出站。 ## 目标与边界 @@ -10,7 +10,7 @@ Issue #60 将既有 Telegram 长轮询和新的企业微信自建应用接入同 ```text Telegram long polling / WeCom HTTPS callback - -> verified Binding + normalized text + -> verified Binding + normalized text/media -> Gateway Dispatch -> Session/Event + Reply Outbox -> channel Provider @@ -18,12 +18,13 @@ Telegram long polling / WeCom HTTPS callback 基础阶段支持: -- 企业微信自建应用的 URL 验证、签名校验、AES 解密和文本消息; +- 企业微信自建应用的 URL 验证、签名校验、AES 解密、文本消息和受控媒体附件; - 内部成员单聊;Issue #77 增加群聊入站/出站目标; - Binding-aware Session identity、持久化入站幂等和受控文本回复; - 既有 Outbox 的 retry、lease/fencing、dead-letter 与重启恢复语义。 -媒体、卡片、被动 XML 回复和第三方应用仍不在自建应用 provider 范围。公众号和微信客服 +卡片、被动 XML 回复和第三方应用仍不在自建应用 provider 范围。媒体仅通过 Issue #98 的 +attachment store / rich Outbox 边界进入,不允许 provider URL 或 token 穿过 Gateway。公众号和微信客服 使用 `trpcservice/channels/wechat` 中互不兼容的显式 provider,不复用 WeCom credential。 `channels.ChannelWeCom` 的持久化值继续为 `wecom` 以保持已有 Binding 和 Admin API @@ -90,7 +91,7 @@ Secret value 不进入 Binding、digest、Event、Outbox、日志、trace、错 Handler 自己拥有 ACK 后的 bounded execution drain,并由 Runtime 的 `BeginShutdown/Close` 取消和 join。 -共享 Adapter conformance 测试验证:只接受 text、消息身份稳定、取消原样传播、未验证 +共享 Adapter conformance 测试验证:只接受 text 和受控媒体、消息身份稳定、取消原样传播、未验证 payload 不能选择 Tenant/App/Binding、重复入站不重新执行 Runner,以及失败不暴露 供应商原始错误或 Secret。 @@ -122,22 +123,24 @@ Handler 仅从 route key 得到候选;它不得从 XML、header、query 或 ca `nonce`、密文和 receive ID 验证、解密后才产生 `VerifiedBinding`。随后重新读取 Tenant、 Binding、App 的可信快照来构造 `RoutingTarget`。 -GET URL 验证成功只返回解密后的 `echostr`。POST 只接受严格 XML text payload;未知或 -非 text 消息、无效时间戳/nonce/签名、未知/过期候选、inactive Tenant/Binding/App、 +GET URL 验证成功只返回解密后的 `echostr`。POST 接受严格 XML `text`、`image`、`file`、 +`voice` 和 `video` payload;未知类型、缺少 `MediaId`、未配置 attachment store/downloader、 +无效时间戳/nonce/签名、未知/过期候选、inactive Tenant/Binding/App、 receive ID 或 AgentID 不匹配都在 Runner 前失败关闭。对外响应不透露候选、Tenant 或 Secret 细节。 ## 消息、Session 与回复地址 -自建应用文本入站规范化为: +自建应用入站规范化为文本或媒体附件: | Gateway 字段 | 企业微信来源 | 约束 | | --- | --- | --- | | `ExternalMessageID` | `MsgId` | 不能为空,作为 durable idempotency key | | `ExternalUserID` | `FromUserName` | 稳定成员 UserID | -| `Content` | `Content` | text-only,交给 Gateway 再次规范化 | +| `Content` | `Content` 或稳定媒体 marker | 交给 Gateway 再次规范化 | +| `Attachments` | `MediaId` 对应内容 | 仅保存内部 attachment reference,不保存 provider 下载 URL | | direct peer | `FromUserName` | 生成单聊 session 和回复收件人 | -| group chat | 不适用 | 此自建应用 callback 只支持单聊;群机器人、公众号和微信客服需要独立 Adapter | +| group chat | `ChatId` | 生成群聊 session 和回复 chat target;群机器人、公众号和微信客服需要独立 Adapter | Session 和 user identity 继续由 `RoutingTarget.RunnerIdentity` 以 Channel、Binding、 conversation kind、外部稳定 ID 和可选 thread 的长度前缀编码构造。两个 Binding、两个群 @@ -149,9 +152,10 @@ conversation kind、外部稳定 ID 和可选 thread 的长度前缀编码构造 Provider。这样同一个 Bot 或 WeCom App 才能回复多个用户/会话,重启、重试和 dead-letter 不会丢失目的地。 -现有 reply materializer 在进入 Outbox 前按企业微信文本限制生成持久化片段;企业微信 -Provider 再校验每个片段,使用应用 `access_token` 调用发送应用消息接口,并将成功返回的 -provider receipt 写入既有 Outbox 状态机。HTTP、token 或供应商错误仅映射为稳定 +现有 reply materializer 在进入 Outbox 前生成持久化文本或媒体片段;企业微信 Provider +再校验每个片段,文本直接发送,图片和文档先上传临时素材再发送 `image`/`file` 应用消息。 +不支持的音频、视频或未配置 attachment reader 的媒体片段使用 Outbox 内的 deterministic +fallback。成功返回的 provider receipt 写入既有 Outbox 状态机。HTTP、token 或供应商错误仅映射为稳定 retryable/permanent error class。Provider 不把原始 body、URL、token 或消息内容写入日志。 ## 可靠性、关联与审计 @@ -169,7 +173,7 @@ duplicate ingress、delivery 成功/重试/dead-letter,以及 channel/error cl ## 验收矩阵 - URL 验证、AES decrypt/encrypt、签名/receive ID/AgentID failures 和安全错误响应; -- direct text 到 Gateway/Runner/Outbox;WeCom group callback 拒绝; +- direct text/media 到 Gateway/Runner/Outbox;WeCom group callback 按 `ChatId` 规范化为 group; - duplicate、并发、乱序和跨 Tenant/Binding identity; - Context cancellation、retry/dead-letter、stale fence 和 restart recovery; - Telegram Adapter 仍维持现有 long-polling、直接回复和 lifecycle 行为; @@ -195,7 +199,9 @@ WECOM_SECRET_REF callback path 的最后一段是公开 route key;环境变量不保存 route key,也不通过 path 直接构造 Tenant/App/Binding 身份。Bootstrap 会用 candidate index、scoped credential resolver 和当前 Tenant/App/Binding 快照完成可信路由,并为该 tenant 启动一个 -binding-aware Outbox worker。未设置任何 `WECOM_*` 变量时,现有无 WeCom 的环境行为保持 +binding-aware Outbox worker。若 runtime store 实现 attachment store,bootstrap 会同时启用 +WeCom 受控媒体 downloader 和图片/文件原生出站;否则媒体回调 fail closed,出站媒体走文本 +fallback。未设置任何 `WECOM_*` 变量时,现有无 WeCom 的环境行为保持 不变;只设置其中一部分会拒绝启动。 ## 运维前提 diff --git a/docs/mkdocs.yml b/docs/mkdocs.yml index 7e9655b1..36fed15c 100644 --- a/docs/mkdocs.yml +++ b/docs/mkdocs.yml @@ -43,8 +43,9 @@ nav: - "Issue #82 Agent App Registry and Tenant Canary": issue-82-agent-app-registry.md - Gateway、Execution Plan 与 HTTP/SSE: gateway.md - Telegram 长轮询 Adapter: telegram.md - - 企业微信自建应用 Text Webhook: wecom.md + - 企业微信自建应用 Channel Adapter: wecom.md - "Issue #77 Channel Capabilities": issue-77-channel-capabilities.md + - "Issue #98 Native Media": issue-98-native-media.md - Tenant 运行时持久化: runtime-storage.md - Runtime Capabilities (#75): runtime-capabilities.md - "Issue #76 Stateless Worker and Migration": issue-76-stateless-worker.md diff --git a/examples/telegram-e2e/README.md b/examples/telegram-e2e/README.md index e4606562..ec721526 100644 --- a/examples/telegram-e2e/README.md +++ b/examples/telegram-e2e/README.md @@ -37,7 +37,7 @@ go run ./examples/telegram-e2e The command prints a unique ordinary-text marker. Open the receiver Bot in Telegram, send that marker, and confirm the `telegram-e2e-ok:` reply. Commands, -media, and rich updates are intentionally outside this first E2E. Press +media, and rich updates are intentionally outside this live E2E; media behavior is covered by deterministic fake tests. Press `Ctrl+C` to stop the local polling process cleanly. If PowerShell can reach `api.telegram.org` but this command reports diff --git a/examples/wecom-e2e/README.md b/examples/wecom-e2e/README.md index cff0508b..7a294b60 100644 --- a/examples/wecom-e2e/README.md +++ b/examples/wecom-e2e/README.md @@ -28,6 +28,10 @@ The test applies repository migrations, sends a signed encrypted text callback, checks the durable reply and provider delivery, then sends the same callback again to prove the Runner is not executed twice. +Native media ingress and image/file egress are covered by deterministic channel +unit tests; this E2E keeps using text so it does not need live WeCom media +credentials or public provider endpoints. + CI runs this test in a dedicated PostgreSQL service through the **WeCom deterministic E2E** job. It intentionally does not require `WECOM_*`, Telegram, model, or other production credentials. diff --git a/migrations/0014_runtime_attachments.up.sql b/migrations/0014_runtime_attachments.up.sql new file mode 100644 index 00000000..2cfe2405 --- /dev/null +++ b/migrations/0014_runtime_attachments.up.sql @@ -0,0 +1,33 @@ +-- Issue #98: tenant-scoped attachment metadata and event ownership. +SET LOCAL search_path = pg_catalog, public, pg_temp; + +CREATE TABLE public.runtime_attachment ( + tenant_id TEXT NOT NULL, + attachment_id TEXT NOT NULL CHECK (length(btrim(attachment_id)) BETWEEN 1 AND 256), + kind TEXT NOT NULL CHECK (kind IN ('image', 'video', 'audio', 'document')), + mime_type TEXT NOT NULL CHECK (length(btrim(mime_type)) BETWEEN 1 AND 256), + name TEXT NOT NULL DEFAULT '' CHECK (length(name) <= 512), + size BIGINT NOT NULL CHECK (size > 0 AND size <= 67108864), + sha256 TEXT NOT NULL CHECK (sha256 ~ '^[0-9a-f]{64}$'), + provider TEXT NOT NULL DEFAULT '' CHECK (length(provider) <= 64), + provider_id TEXT NOT NULL DEFAULT '' CHECK (length(provider_id) <= 512), + event_id TEXT, + expires_at TIMESTAMPTZ NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + PRIMARY KEY (tenant_id, attachment_id), + FOREIGN KEY (tenant_id) REFERENCES public.tenant(tenant_id), + FOREIGN KEY (tenant_id, attachment_id) + REFERENCES public.runtime_object(tenant_id, object_key) + ON DELETE CASCADE, + FOREIGN KEY (tenant_id, event_id) + REFERENCES public.message_event(tenant_id, event_id) + ON DELETE SET NULL (event_id) +); + +CREATE INDEX runtime_attachment_cleanup_idx + ON public.runtime_attachment (tenant_id, expires_at) + WHERE event_id IS NULL; + +REVOKE ALL ON TABLE public.runtime_attachment FROM PUBLIC; +GRANT SELECT, INSERT, UPDATE, DELETE ON public.runtime_attachment TO tenant_app_writer; +GRANT ALL PRIVILEGES ON public.runtime_attachment TO migration_owner; diff --git a/migrations/0015_runtime_reply_media.up.sql b/migrations/0015_runtime_reply_media.up.sql new file mode 100644 index 00000000..8a2ee9b1 --- /dev/null +++ b/migrations/0015_runtime_reply_media.up.sql @@ -0,0 +1,46 @@ +-- Issue #98: durable media reply descriptors with deterministic text fallback. +SET LOCAL search_path = pg_catalog, public, pg_temp; + +ALTER TABLE public.reply_outbox + ADD COLUMN reply_kind TEXT NOT NULL DEFAULT 'text' + CHECK (reply_kind IN ('text', 'image', 'video', 'audio', 'document')), + ADD COLUMN attachment_id TEXT NOT NULL DEFAULT '' + CHECK (length(attachment_id) <= 256), + ADD COLUMN attachment_kind TEXT NOT NULL DEFAULT '' + CHECK (attachment_kind IN ('', 'image', 'video', 'audio', 'document')), + ADD COLUMN attachment_mime_type TEXT NOT NULL DEFAULT '' + CHECK (length(attachment_mime_type) <= 256), + ADD COLUMN attachment_name TEXT NOT NULL DEFAULT '' + CHECK (length(attachment_name) <= 512), + ADD COLUMN attachment_size BIGINT NOT NULL DEFAULT 0 + CHECK (attachment_size >= 0 AND attachment_size <= 67108864), + ADD COLUMN attachment_sha256 TEXT NOT NULL DEFAULT '' + CHECK (attachment_sha256 = '' OR attachment_sha256 ~ '^[0-9a-f]{64}$'), + ADD COLUMN attachment_provider TEXT NOT NULL DEFAULT '' + CHECK (length(attachment_provider) <= 64), + ADD COLUMN attachment_provider_id TEXT NOT NULL DEFAULT '' + CHECK (length(attachment_provider_id) <= 512), + ADD COLUMN fallback TEXT NOT NULL DEFAULT '', + ADD CONSTRAINT reply_outbox_media_contract CHECK ( + ( + reply_kind = 'text' + AND attachment_id = '' + AND attachment_kind = '' + AND attachment_mime_type = '' + AND attachment_name = '' + AND attachment_size = 0 + AND attachment_sha256 = '' + AND attachment_provider = '' + AND attachment_provider_id = '' + AND fallback = '' + ) + OR ( + reply_kind <> 'text' + AND reply_kind = attachment_kind + AND btrim(attachment_id) <> '' + AND btrim(attachment_mime_type) <> '' + AND attachment_size > 0 + AND attachment_sha256 ~ '^[0-9a-f]{64}$' + AND btrim(fallback) <> '' + ) + ); diff --git a/migrations/history_test.go b/migrations/history_test.go index 870a8961..493f48c8 100644 --- a/migrations/history_test.go +++ b/migrations/history_test.go @@ -15,10 +15,13 @@ func TestOrderedFilesAreContiguousAndDigestable(t *testing.T) { if err != nil { t.Fatal(err) } - if len(files) != 13 || files[0].version != 1 || files[1].version != 2 || files[2].version != 3 || files[3].version != 4 || files[4].version != 5 || files[5].version != 6 || files[6].version != 7 || files[7].version != 8 || files[8].version != 9 || files[9].version != 10 || files[10].version != 11 || files[11].version != 12 || files[12].version != 13 { + if len(files) != 15 { t.Fatalf("migration order = %+v", files) } - for _, migration := range files { + for index, migration := range files { + if migration.version != index+1 { + t.Fatalf("migration order = %+v", files) + } if migration.name == "" || len(migration.digest) != 64 || migration.sql == "" { t.Fatalf("invalid migration metadata = %+v", migration) } diff --git a/migrations/migration_test.go b/migrations/migration_test.go index 505675a0..e5ff0e84 100644 --- a/migrations/migration_test.go +++ b/migrations/migration_test.go @@ -57,6 +57,8 @@ func TestPostgreSQLControlPlaneMigration(t *testing.T) { "0011_reply_trace_parent.up.sql", "0012_runtime_capabilities.up.sql", "0013_execution_queue.up.sql", + "0014_runtime_attachments.up.sql", + "0015_runtime_reply_media.up.sql", } { path := filepath.Join(migrationDir, name) contents, err := os.ReadFile(path) // #nosec G304 -- names are fixed migration files under the repository root. diff --git a/trpcservice/attachment/attachment.go b/trpcservice/attachment/attachment.go new file mode 100644 index 00000000..7b7e3421 --- /dev/null +++ b/trpcservice/attachment/attachment.go @@ -0,0 +1,219 @@ +// Package attachment defines protocol-neutral, tenant-scoped media references. +package attachment + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "mime" + "strings" + "time" + "unicode/utf8" +) + +var ( + // ErrInvalid reports malformed attachment metadata or content. + ErrInvalid = errors.New("invalid attachment") +) + +const ( + maxReferenceIDRunes = 256 + maxProviderRunes = 64 + maxProviderIDRunes = 512 + maxNameRunes = 512 + maxMIMETypeRunes = 256 + maxSizeBytes = 64 << 20 +) + +// DefaultRetention is the retention applied to an uploaded attachment when +// its caller does not provide an earlier expiry time. +const DefaultRetention = 24 * time.Hour + +// Kind identifies the protocol-neutral media category of an attachment. +type Kind string + +const ( + // KindImage identifies an image attachment. + KindImage Kind = "image" + // KindVideo identifies a video attachment. + KindVideo Kind = "video" + // KindAudio identifies an audio attachment. + KindAudio Kind = "audio" + // KindDocument identifies a document or other file attachment. + KindDocument Kind = "document" +) + +// Reference identifies one attachment owned by the authenticated tenant. It +// intentionally contains no provider URL, credential, or provider fetch token. +type Reference struct { + ID string + Kind Kind + MIMEType string + Name string + Size int64 + SHA256 string + Provider string + ProviderID string +} + +// Normalize validates a Reference and returns its canonical representation. +func (reference Reference) Normalize() (Reference, error) { + value := reference + value.ID = strings.TrimSpace(value.ID) + value.Name = strings.TrimSpace(value.Name) + value.MIMEType = strings.ToLower(strings.TrimSpace(value.MIMEType)) + value.SHA256 = strings.ToLower(strings.TrimSpace(value.SHA256)) + value.Provider = strings.ToLower(strings.TrimSpace(value.Provider)) + value.ProviderID = strings.TrimSpace(value.ProviderID) + if !validText(value.ID, maxReferenceIDRunes, true) { + return Reference{}, fmt.Errorf("%w: reference ID is invalid", ErrInvalid) + } + switch value.Kind { + case KindImage, KindVideo, KindAudio, KindDocument: + default: + return Reference{}, fmt.Errorf("%w: kind is invalid", ErrInvalid) + } + if !validText(value.Name, maxNameRunes, false) || !validText(value.MIMEType, maxMIMETypeRunes, true) || !validText(value.Provider, maxProviderRunes, false) || !validText(value.ProviderID, maxProviderIDRunes, false) { + return Reference{}, fmt.Errorf("%w: name or MIME type is invalid", ErrInvalid) + } + mediaType, params, err := mime.ParseMediaType(value.MIMEType) + if err != nil || mediaType != value.MIMEType || len(params) != 0 || !strings.Contains(value.MIMEType, "/") || !matchesKind(value.Kind, value.MIMEType) { + return Reference{}, fmt.Errorf("%w: MIME type is invalid", ErrInvalid) + } + if value.Size < 1 || value.Size > maxSizeBytes { + return Reference{}, fmt.Errorf("%w: size is invalid", ErrInvalid) + } + if len(value.SHA256) != sha256.Size*2 || !isLowerHex(value.SHA256) { + return Reference{}, fmt.Errorf("%w: digest is invalid", ErrInvalid) + } + return value, nil +} + +// Upload describes verified provider metadata for one attachment before its +// bytes have been persisted. The byte count is enforced by Store implementations. +type Upload struct { + ID string + Kind Kind + MIMEType string + Name string + Size int64 + Provider string + ProviderID string + ExpiresAt time.Time +} + +// Normalize validates Upload metadata and applies the default retention. +func (upload Upload) Normalize(now time.Time) (Upload, error) { + value := upload + value.ID = strings.TrimSpace(value.ID) + value.Name = strings.TrimSpace(value.Name) + value.MIMEType = strings.ToLower(strings.TrimSpace(value.MIMEType)) + value.Provider = strings.ToLower(strings.TrimSpace(value.Provider)) + value.ProviderID = strings.TrimSpace(value.ProviderID) + if !validText(value.ID, maxReferenceIDRunes, true) { + return Upload{}, fmt.Errorf("%w: upload ID is invalid", ErrInvalid) + } + if !validText(value.Name, maxNameRunes, false) || !validText(value.MIMEType, maxMIMETypeRunes, true) || !validText(value.Provider, maxProviderRunes, false) || !validText(value.ProviderID, maxProviderIDRunes, false) { + return Upload{}, fmt.Errorf("%w: upload metadata is invalid", ErrInvalid) + } + switch value.Kind { + case KindImage, KindVideo, KindAudio, KindDocument: + default: + return Upload{}, fmt.Errorf("%w: kind is invalid", ErrInvalid) + } + mediaType, params, err := mime.ParseMediaType(value.MIMEType) + if err != nil || mediaType != value.MIMEType || len(params) != 0 || !matchesKind(value.Kind, value.MIMEType) { + return Upload{}, fmt.Errorf("%w: MIME type is invalid", ErrInvalid) + } + if value.Size < 1 || value.Size > maxSizeBytes { + return Upload{}, fmt.Errorf("%w: size is invalid", ErrInvalid) + } + if now.IsZero() { + now = time.Now().UTC() + } + if value.ExpiresAt.IsZero() { + value.ExpiresAt = now.Add(DefaultRetention) + } + value.ExpiresAt = value.ExpiresAt.UTC() + if !value.ExpiresAt.After(now) || value.ExpiresAt.After(now.Add(DefaultRetention)) { + return Upload{}, fmt.Errorf("%w: expiry is invalid", ErrInvalid) + } + return value, nil +} + +// Content is the verified data loaded for one Reference. Callers own Data and +// must not retain it beyond the model request that consumes it. +type Content struct { + Data []byte +} + +// Validate checks that Content is the exact content described by reference. +func (content Content) Validate(reference Reference) error { + if _, err := reference.Normalize(); err != nil { + return err + } + if int64(len(content.Data)) != reference.Size { + return fmt.Errorf("%w: content size does not match reference", ErrInvalid) + } + digest := sha256.Sum256(content.Data) + if hex.EncodeToString(digest[:]) != reference.SHA256 { + return fmt.Errorf("%w: content digest does not match reference", ErrInvalid) + } + return nil +} + +// Clone returns an independent copy of Content. +func (content Content) Clone() Content { + return Content{Data: append([]byte(nil), content.Data...)} +} + +// Reader loads tenant-owned attachment data for a durable event and validated +// Reference. The reader must enforce tenant scope, event ownership, and return +// a defensive copy of the bytes. +type Reader interface { + Load(context.Context, string, string, Reference) (Content, error) +} + +// Binder associates verified attachment references with a durable message +// event before a Runner may load their bytes. +type Binder interface { + BindAttachments(context.Context, string, string, []Reference) error +} + +func validText(value string, maximum int, required bool) bool { + if (required && value == "") || !utf8.ValidString(value) || len([]rune(value)) > maximum { + return false + } + for _, character := range value { + if character < 0x20 || character == 0x7f { + return false + } + } + return true +} + +func isLowerHex(value string) bool { + for _, character := range value { + if !('0' <= character && character <= '9' || 'a' <= character && character <= 'f') { + return false + } + } + return true +} + +func matchesKind(kind Kind, contentType string) bool { + switch kind { + case KindImage: + return strings.HasPrefix(contentType, "image/") + case KindVideo: + return strings.HasPrefix(contentType, "video/") + case KindAudio: + return strings.HasPrefix(contentType, "audio/") + case KindDocument: + return !strings.HasPrefix(contentType, "image/") && !strings.HasPrefix(contentType, "video/") && !strings.HasPrefix(contentType, "audio/") + default: + return false + } +} diff --git a/trpcservice/attachment/attachment_test.go b/trpcservice/attachment/attachment_test.go new file mode 100644 index 00000000..14d945db --- /dev/null +++ b/trpcservice/attachment/attachment_test.go @@ -0,0 +1,129 @@ +package attachment + +import ( + "crypto/sha256" + "encoding/hex" + "errors" + "strings" + "testing" + "time" +) + +func TestReferenceNormalizeAndContentValidate(t *testing.T) { + data := []byte("image") + digest := sha256.Sum256(data) + reference := Reference{ID: " attachment-1 ", Kind: KindImage, MIMEType: " IMAGE/PNG ", Name: " chart.png ", Size: int64(len(data)), SHA256: " " + hex.EncodeToString(digest[:]) + " ", Provider: " Telegram ", ProviderID: " file-1 "} + normalized, err := reference.Normalize() + if err != nil || normalized.ID != "attachment-1" || normalized.MIMEType != "image/png" || normalized.Name != "chart.png" || normalized.Provider != "telegram" || normalized.ProviderID != "file-1" { + t.Fatalf("Normalize = %+v, %v", normalized, err) + } + if err := (Content{Data: data}).Validate(normalized); err != nil { + t.Fatalf("Validate = %v", err) + } + wrongSize := normalized + wrongSize.Size++ + if err := (Content{Data: data}).Validate(wrongSize); !errors.Is(err, ErrInvalid) { + t.Fatalf("size mismatch content error = %v", err) + } + if err := (Content{Data: []byte("other")}).Validate(normalized); !errors.Is(err, ErrInvalid) { + t.Fatalf("mismatched content error = %v", err) + } + clone := (Content{Data: data}).Clone() + clone.Data[0] = 'I' + if string(data) != "image" { + t.Fatalf("Clone aliased source data: %q", data) + } +} + +func TestUploadNormalizeAppliesBoundedRetentionAndMediaFamily(t *testing.T) { + now := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + normalized, err := (Upload{ID: " attachment-1 ", Kind: KindVideo, MIMEType: " VIDEO/MP4 ", Name: " clip.mp4 ", Size: 1, Provider: " Telegram ", ProviderID: " file-1 "}).Normalize(now) + if err != nil || normalized.ID != "attachment-1" || normalized.MIMEType != "video/mp4" || normalized.Name != "clip.mp4" || normalized.Provider != "telegram" || normalized.ProviderID != "file-1" || !normalized.ExpiresAt.Equal(now.Add(DefaultRetention)) { + t.Fatalf("Normalize = %+v, %v", normalized, err) + } + zeroNow, err := (Upload{ID: "attachment-1", Kind: KindDocument, MIMEType: "application/pdf", Size: 1}).Normalize(time.Time{}) + if err != nil || zeroNow.ExpiresAt.IsZero() { + t.Fatalf("zero-now Normalize = %+v, %v", zeroNow, err) + } + if _, err := (Upload{ID: "attachment-1", Kind: KindImage, MIMEType: "application/pdf", Size: 1}).Normalize(now); !errors.Is(err, ErrInvalid) { + t.Fatalf("mismatched media family error = %v", err) + } + if _, err := (Upload{ID: "attachment-1", Kind: KindDocument, MIMEType: "application/pdf", Size: 1, ExpiresAt: now.Add(DefaultRetention + time.Second)}).Normalize(now); !errors.Is(err, ErrInvalid) { + t.Fatalf("excess retention error = %v", err) + } +} + +func TestReferenceNormalizeRejectsMalformedMetadata(t *testing.T) { + data := []byte("x") + digest := sha256.Sum256(data) + valid := Reference{ID: "attachment-1", Kind: KindImage, MIMEType: "image/png", Name: "chart.png", Size: int64(len(data)), SHA256: hex.EncodeToString(digest[:]), Provider: "telegram", ProviderID: "file-1"} + for _, test := range []struct { + name string + mutate func(*Reference) + wantError string + }{ + {name: "empty id", mutate: func(reference *Reference) { reference.ID = "" }, wantError: "reference ID"}, + {name: "control id", mutate: func(reference *Reference) { reference.ID = "bad\nid" }, wantError: "reference ID"}, + {name: "unknown kind", mutate: func(reference *Reference) { reference.Kind = "sticker" }, wantError: "kind"}, + {name: "bad optional text", mutate: func(reference *Reference) { reference.Name = "bad\nname" }, wantError: "name or MIME type"}, + {name: "mime parameters", mutate: func(reference *Reference) { reference.MIMEType = "image/png; charset=utf-8" }, wantError: "MIME type"}, + {name: "mime family mismatch", mutate: func(reference *Reference) { reference.Kind, reference.MIMEType = KindAudio, "image/png" }, wantError: "MIME type"}, + {name: "zero size", mutate: func(reference *Reference) { reference.Size = 0 }, wantError: "size"}, + {name: "oversized", mutate: func(reference *Reference) { reference.Size = maxSizeBytes + 1 }, wantError: "size"}, + {name: "short digest", mutate: func(reference *Reference) { reference.SHA256 = "abc" }, wantError: "digest"}, + {name: "non-hex digest", mutate: func(reference *Reference) { reference.SHA256 = strings.Repeat("g", sha256.Size*2) }, wantError: "digest"}, + } { + t.Run(test.name, func(t *testing.T) { + reference := valid + test.mutate(&reference) + _, err := reference.Normalize() + if !errors.Is(err, ErrInvalid) || !strings.Contains(err.Error(), test.wantError) { + t.Fatalf("Normalize error = %v, want %q", err, test.wantError) + } + }) + } +} + +func TestUploadNormalizeRejectsMalformedMetadata(t *testing.T) { + now := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + valid := Upload{ID: "attachment-1", Kind: KindDocument, MIMEType: "application/pdf", Name: "brief.pdf", Size: 1, Provider: "wecom", ProviderID: "media-1", ExpiresAt: now.Add(time.Hour)} + for _, test := range []struct { + name string + mutate func(*Upload) + }{ + {name: "empty id", mutate: func(upload *Upload) { upload.ID = "" }}, + {name: "control provider id", mutate: func(upload *Upload) { upload.ProviderID = "bad\nmedia" }}, + {name: "unknown kind", mutate: func(upload *Upload) { upload.Kind = "sticker" }}, + {name: "mime parameters", mutate: func(upload *Upload) { upload.MIMEType = "application/pdf; charset=utf-8" }}, + {name: "media family mismatch", mutate: func(upload *Upload) { upload.Kind, upload.MIMEType = KindDocument, "audio/mpeg" }}, + {name: "zero size", mutate: func(upload *Upload) { upload.Size = 0 }}, + {name: "past expiry", mutate: func(upload *Upload) { upload.ExpiresAt = now }}, + } { + t.Run(test.name, func(t *testing.T) { + upload := valid + test.mutate(&upload) + if _, err := upload.Normalize(now); !errors.Is(err, ErrInvalid) { + t.Fatalf("Normalize accepted invalid upload: %+v err=%v", upload, err) + } + }) + } +} + +func TestMatchesKindCoversMediaFamilies(t *testing.T) { + for _, test := range []struct { + kind Kind + mime string + want bool + }{ + {kind: KindImage, mime: "image/png", want: true}, + {kind: KindVideo, mime: "video/mp4", want: true}, + {kind: KindAudio, mime: "audio/mpeg", want: true}, + {kind: KindDocument, mime: "application/pdf", want: true}, + {kind: KindDocument, mime: "image/png"}, + {kind: Kind("unknown"), mime: "application/octet-stream"}, + } { + if got := matchesKind(test.kind, test.mime); got != test.want { + t.Fatalf("matchesKind(%q, %q) = %t, want %t", test.kind, test.mime, got, test.want) + } + } +} diff --git a/trpcservice/bootstrap/environment.go b/trpcservice/bootstrap/environment.go index 6c7a423f..8b797b66 100644 --- a/trpcservice/bootstrap/environment.go +++ b/trpcservice/bootstrap/environment.go @@ -333,14 +333,22 @@ func environmentWeComComponents(config environmentConfig, channelsRepo channels. return nil, nil, nil } credentials := environmentWeComCredentialResolver{tenantID: config.tenantID, config: *config.wecom} + var attachments runtimestorage.AttachmentStore + if store, ok := runtimeStore.(runtimestorage.AttachmentStore); ok { + attachments = store + } + var mediaDownloader wecom.MediaDownloader + if attachments != nil { + mediaDownloader = &wecom.HTTPMediaDownloader{} + } factory := func(dispatcher gateway.DispatchService) (http.Handler, error) { - return wecom.New(wecom.Config{Candidates: channelsRepo, Tenants: tenantsRepo, Apps: appsRepo, Credentials: credentials, Dispatcher: dispatcher, AuditWriter: auditWriter, Observability: config.telemetry}) + return wecom.New(wecom.Config{Candidates: channelsRepo, Tenants: tenantsRepo, Apps: appsRepo, Credentials: credentials, Dispatcher: dispatcher, Attachments: attachments, MediaDownloader: mediaDownloader, AuditWriter: auditWriter, Observability: config.telemetry}) } owner, err := environmentWeComOwnerFunc() if err != nil { return nil, nil, err } - worker, err := newEnvironmentWeComWorker(outbox.Config{Store: runtimeStore, Provider: &wecom.BindingProvider{Bindings: channelsRepo, Credentials: credentials}, Channel: "wecom", ProviderName: "wecom", TenantID: config.tenantID, Owner: owner, LeaseDuration: 30 * time.Second, AuditWriter: auditWriter, Observability: config.telemetry}) + worker, err := newEnvironmentWeComWorker(outbox.Config{Store: runtimeStore, Provider: &wecom.BindingProvider{Bindings: channelsRepo, Credentials: credentials, Attachments: attachments}, Channel: "wecom", ProviderName: "wecom", TenantID: config.tenantID, Owner: owner, LeaseDuration: 30 * time.Second, AuditWriter: auditWriter, Observability: config.telemetry}) return factory, worker, err } diff --git a/trpcservice/bootstrap/responses_model.go b/trpcservice/bootstrap/responses_model.go index c7b2400a..a8213ad6 100644 --- a/trpcservice/bootstrap/responses_model.go +++ b/trpcservice/bootstrap/responses_model.go @@ -4,6 +4,7 @@ import ( "bufio" "bytes" "context" + "encoding/base64" "encoding/json" "fmt" "net/http" @@ -37,8 +38,19 @@ type responsesInputItem struct { Content []responsesContentPart `json:"content"` } type responsesContentPart struct { - Type string `json:"type"` - Text string `json:"text"` + Type string `json:"type"` + Text string `json:"text,omitempty"` + ImageURL string `json:"image_url,omitempty"` + Detail string `json:"detail,omitempty"` + FileID string `json:"file_id,omitempty"` + FileURL string `json:"file_url,omitempty"` + FileData string `json:"file_data,omitempty"` + Filename string `json:"filename,omitempty"` + InputAudio *responsesInputAudio `json:"input_audio,omitempty"` +} +type responsesInputAudio struct { + Data string `json:"data"` + Format string `json:"format"` } func (m *responsesModel) stream(ctx context.Context, request *trpcmodel.Request, out chan<- *trpcmodel.Response) { @@ -87,17 +99,125 @@ func responsesInput(request *trpcmodel.Request) []responsesInputItem { if role == "" { role = string(trpcmodel.RoleUser) } - text := message.Content - if text == "" { - for _, part := range message.ContentParts { - if part.Text != nil { - text += *part.Text - } + input = append(input, responsesInputItem{Role: role, Content: responsesContent(message)}) + } + return input +} + +func responsesContent(message trpcmodel.Message) []responsesContentPart { + parts := make([]responsesContentPart, 0, 1+len(message.ContentParts)) + text := message.Content + if text == "" { + for _, part := range message.ContentParts { + if part.Text != nil { + text += *part.Text } } - input = append(input, responsesInputItem{Role: role, Content: []responsesContentPart{{Type: "input_text", Text: text}}}) } - return input + if text != "" || len(message.ContentParts) == 0 { + parts = append(parts, responsesTextPart(text)) + } + for _, part := range message.ContentParts { + if part.Type == trpcmodel.ContentTypeText || part.Text != nil { + continue + } + if converted, ok := responsesContentFromPart(part); ok { + parts = append(parts, converted) + } + } + if len(parts) == 0 { + parts = append(parts, responsesTextPart("")) + } + return parts +} + +func responsesContentFromPart(part trpcmodel.ContentPart) (responsesContentPart, bool) { + switch part.Type { + case trpcmodel.ContentTypeImage: + if part.Image == nil { + return responsesTextPart("[image attachment omitted: image data is missing]"), true + } + if strings.TrimSpace(part.Image.URL) != "" { + return responsesContentPart{Type: "input_image", ImageURL: strings.TrimSpace(part.Image.URL), Detail: part.Image.Detail}, true + } + if len(part.Image.Data) == 0 { + return responsesTextPart("[image attachment omitted: image data is missing]"), true + } + return responsesContentPart{Type: "input_image", ImageURL: dataURL("image", part.Image.Format, part.Image.Data), Detail: part.Image.Detail}, true + case trpcmodel.ContentTypeFile: + return responsesFilePart(part.File) + case trpcmodel.ContentTypeAudio: + return responsesAudioPart(part.Audio) + case trpcmodel.ContentTypeVideo: + return responsesTextPart("[video attachment omitted: model video input is not enabled]"), true + default: + return responsesContentPart{}, false + } +} + +func responsesFilePart(file *trpcmodel.File) (responsesContentPart, bool) { + if file == nil { + return responsesTextPart("[file attachment omitted: file data is missing]"), true + } + if id := strings.TrimSpace(file.FileID); id != "" { + return responsesContentPart{Type: "input_file", FileID: id}, true + } + if fileURL := strings.TrimSpace(file.URL); fileURL != "" { + return responsesContentPart{Type: "input_file", FileURL: fileURL}, true + } + if len(file.Data) == 0 { + return responsesTextPart("[file attachment omitted: file data is missing]"), true + } + return responsesContentPart{Type: "input_file", FileData: dataURL("application", file.MimeType, file.Data), Filename: responsesFilename(file.Name)}, true +} + +func responsesAudioPart(audio *trpcmodel.Audio) (responsesContentPart, bool) { + if audio == nil || len(audio.Data) == 0 { + return responsesTextPart("[audio attachment omitted: audio data is missing]"), true + } + format := responsesAudioFormat(audio.Format) + if format == "" { + return responsesTextPart("[audio attachment omitted: unsupported audio format]"), true + } + return responsesContentPart{Type: "input_audio", InputAudio: &responsesInputAudio{Data: base64.StdEncoding.EncodeToString(audio.Data), Format: format}}, true +} + +func responsesTextPart(text string) responsesContentPart { + return responsesContentPart{Type: "input_text", Text: text} +} + +func dataURL(defaultFamily, format string, data []byte) string { + mediaType := strings.ToLower(strings.TrimSpace(format)) + if mediaType == "" { + mediaType = defaultFamily + "/octet-stream" + } else if !strings.Contains(mediaType, "/") { + mediaType = defaultFamily + "/" + strings.TrimPrefix(mediaType, ".") + } + return "data:" + mediaType + ";base64," + base64.StdEncoding.EncodeToString(data) +} + +func responsesFilename(name string) string { + if trimmed := strings.TrimSpace(name); trimmed != "" { + return trimmed + } + return "attachment" +} + +func responsesAudioFormat(format string) string { + value := strings.ToLower(strings.TrimSpace(format)) + if _, subtype, ok := strings.Cut(value, "/"); ok { + value = subtype + } + switch value { + case "mpeg", "mpga": + return "mp3" + case "x-wav", "wave": + return "wav" + case "mp3", "wav": + return value + default: + return "" + } } func (m *responsesModel) doRequest(ctx context.Context, body []byte) (*http.Response, error) { diff --git a/trpcservice/bootstrap/responses_model_test.go b/trpcservice/bootstrap/responses_model_test.go index 00260c66..ee48ddb0 100644 --- a/trpcservice/bootstrap/responses_model_test.go +++ b/trpcservice/bootstrap/responses_model_test.go @@ -28,6 +28,103 @@ func TestResponsesModelInfoAndInputDefaults(t *testing.T) { } } +func TestResponsesModelMapsContentParts(t *testing.T) { + input := responsesInput(&trpcmodel.Request{Messages: []trpcmodel.Message{{ + Role: trpcmodel.RoleUser, + Content: "describe", + ContentParts: []trpcmodel.ContentPart{ + {Type: trpcmodel.ContentTypeImage, Image: &trpcmodel.Image{Data: []byte{1, 2}, Detail: "auto", Format: "png"}}, + {Type: trpcmodel.ContentTypeFile, File: &trpcmodel.File{Name: "brief.pdf", Data: []byte("pdf"), MimeType: "application/pdf"}}, + {Type: trpcmodel.ContentTypeAudio, Audio: &trpcmodel.Audio{Data: []byte("mp3"), Format: "audio/mpeg"}}, + {Type: trpcmodel.ContentTypeVideo, Video: &trpcmodel.Video{Data: []byte("mp4"), Format: "mp4"}}, + }, + }}}) + if len(input) != 1 || len(input[0].Content) != 5 { + t.Fatalf("input = %#v", input) + } + parts := input[0].Content + if parts[0].Type != "input_text" || parts[0].Text != "describe" { + t.Fatalf("text part = %#v", parts[0]) + } + if parts[1].Type != "input_image" || parts[1].ImageURL != "data:image/png;base64,AQI=" || parts[1].Detail != "auto" { + t.Fatalf("image part = %#v", parts[1]) + } + if parts[2].Type != "input_file" || parts[2].Filename != "brief.pdf" || parts[2].FileData != "data:application/pdf;base64,cGRm" { + t.Fatalf("file part = %#v", parts[2]) + } + if parts[3].Type != "input_audio" || parts[3].InputAudio == nil || parts[3].InputAudio.Data != "bXAz" || parts[3].InputAudio.Format != "mp3" { + t.Fatalf("audio part = %#v", parts[3]) + } + if parts[4].Type != "input_text" || parts[4].Text != "[video attachment omitted: model video input is not enabled]" { + t.Fatalf("video fallback part = %#v", parts[4]) + } +} + +func TestResponsesModelMapsContentPartFallbacksAndReferences(t *testing.T) { + parts := responsesContent(trpcmodel.Message{ContentParts: []trpcmodel.ContentPart{ + {Type: trpcmodel.ContentTypeImage, Image: &trpcmodel.Image{URL: " https://files.example/image.png ", Detail: "low"}}, + {Type: trpcmodel.ContentTypeImage}, + {Type: trpcmodel.ContentTypeFile, File: &trpcmodel.File{FileID: " file-123 "}}, + {Type: trpcmodel.ContentTypeFile, File: &trpcmodel.File{URL: " https://files.example/brief.pdf "}}, + {Type: trpcmodel.ContentTypeFile, File: &trpcmodel.File{Data: []byte("raw")}}, + {Type: trpcmodel.ContentTypeFile}, + {Type: trpcmodel.ContentTypeAudio, Audio: &trpcmodel.Audio{Data: []byte("wav"), Format: "audio/x-wav"}}, + {Type: trpcmodel.ContentTypeAudio, Audio: &trpcmodel.Audio{Data: []byte("ogg"), Format: "audio/ogg"}}, + {Type: "unsupported"}, + }}) + if len(parts) != 8 { + t.Fatalf("parts = %#v", parts) + } + assertResponsesPart(t, parts[0], responsesPartWant{typ: "input_image", imageURL: "https://files.example/image.png", detail: "low"}) + assertResponsesPart(t, parts[1], responsesPartWant{typ: "input_text", text: "[image attachment omitted: image data is missing]"}) + assertResponsesPart(t, parts[2], responsesPartWant{typ: "input_file", fileID: "file-123"}) + assertResponsesPart(t, parts[3], responsesPartWant{typ: "input_file", fileURL: "https://files.example/brief.pdf"}) + assertResponsesPart(t, parts[4], responsesPartWant{typ: "input_file", fileData: "data:application/octet-stream;base64,cmF3", filename: "attachment"}) + assertResponsesPart(t, parts[5], responsesPartWant{typ: "input_text", text: "[file attachment omitted: file data is missing]"}) + assertResponsesPart(t, parts[6], responsesPartWant{typ: "input_audio", audioFormat: "wav"}) + assertResponsesPart(t, parts[7], responsesPartWant{typ: "input_text", text: "[audio attachment omitted: unsupported audio format]"}) + + empty := responsesContent(trpcmodel.Message{ContentParts: []trpcmodel.ContentPart{{Type: "unsupported"}}}) + if len(empty) != 1 || empty[0].Type != "input_text" || empty[0].Text != "" { + t.Fatalf("empty content fallback = %#v", empty) + } +} + +type responsesPartWant struct { + typ, text, imageURL, detail, fileID, fileURL, fileData, filename, audioFormat string +} + +func assertResponsesPart(t *testing.T, got responsesContentPart, want responsesPartWant) { + t.Helper() + if got.Type != want.typ || got.Text != want.text || got.ImageURL != want.imageURL || got.Detail != want.detail || got.FileID != want.fileID || got.FileURL != want.fileURL || got.FileData != want.fileData || got.Filename != want.filename { + t.Fatalf("content part = %#v, want %#v", got, want) + } + if want.audioFormat == "" && got.InputAudio != nil { + t.Fatalf("unexpected audio payload = %#v", got.InputAudio) + } + if want.audioFormat != "" && (got.InputAudio == nil || got.InputAudio.Format != want.audioFormat) { + t.Fatalf("audio part = %#v, want format %q", got, want.audioFormat) + } +} + +func TestResponsesAudioFormatVariants(t *testing.T) { + for _, test := range []struct { + input string + want string + }{ + {input: "mpeg", want: "mp3"}, + {input: "audio/mpga", want: "mp3"}, + {input: "wave", want: "wav"}, + {input: "wav", want: "wav"}, + {input: "flac", want: ""}, + {input: "", want: ""}, + } { + if got := responsesAudioFormat(test.input); got != test.want { + t.Fatalf("responsesAudioFormat(%q) = %q, want %q", test.input, got, test.want) + } + } +} + func TestResponsesModelGenerateContentRejectsInvalidArguments(t *testing.T) { model := &responsesModel{} if responses, err := model.GenerateContent(nil, &trpcmodel.Request{}); responses != nil || err == nil { diff --git a/trpcservice/channels/telegram/provider.go b/trpcservice/channels/telegram/provider.go index 1b0c6d47..62c92257 100644 --- a/trpcservice/channels/telegram/provider.go +++ b/trpcservice/channels/telegram/provider.go @@ -1,14 +1,17 @@ package telegram import ( + "bytes" "context" "errors" "strconv" "sync" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/outbox" runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" "github.com/go-telegram/bot" + "github.com/go-telegram/bot/models" ) // Provider delivers durable outbox segments through one trusted Telegram @@ -16,19 +19,47 @@ import ( // idempotency key; Telegram itself remains at-least-once unless it supports a // matching external idempotency facility. type Provider struct { - client BotClient - chatID int64 - threadID int - mu sync.Mutex - receipts map[string]string + client BotClient + chatID int64 + threadID int + attachments attachment.Reader + mu sync.Mutex + receipts map[string]string +} + +type telegramPhotoSender interface { + SendPhoto(context.Context, *bot.SendPhotoParams) (*models.Message, error) +} + +type telegramDocumentSender interface { + SendDocument(context.Context, *bot.SendDocumentParams) (*models.Message, error) +} + +// ProviderOption configures optional Telegram reply-provider dependencies. +type ProviderOption func(*Provider) + +// WithAttachmentReader enables native media delivery from durable attachment +// references. Nil readers are ignored and keep the text-only fallback path. +func WithAttachmentReader(reader attachment.Reader) ProviderOption { + return func(provider *Provider) { + if reader != nil { + provider.attachments = reader + } + } } // NewProvider creates a Telegram reply provider for a chat and optional thread. -func NewProvider(client BotClient, chatID int64, threadID int) (*Provider, error) { +func NewProvider(client BotClient, chatID int64, threadID int, options ...ProviderOption) (*Provider, error) { if client == nil || chatID == 0 || threadID < 0 { return nil, outbox.ErrInvalid } - return &Provider{client: client, chatID: chatID, threadID: threadID, receipts: map[string]string{}}, nil + provider := &Provider{client: client, chatID: chatID, threadID: threadID, receipts: map[string]string{}} + for _, option := range options { + if option != nil { + option(provider) + } + } + return provider, nil } // Deliver sends one durable reply segment and returns the provider message ID. @@ -43,12 +74,9 @@ func (p *Provider) Deliver(ctx context.Context, value runtimestorage.ReplyOutbox return receipt, nil } p.mu.Unlock() - message, err := p.client.SendMessage(ctx, &bot.SendMessageParams{ChatID: p.chatID, MessageThreadID: p.threadID, Text: value.Payload}) + message, err := p.deliverMessage(ctx, value) if err != nil { - if errors.Is(err, context.Canceled) { - return "", &outbox.DeliveryError{Class: "canceled", Retryable: true} - } - return "", &outbox.DeliveryError{Class: "provider_error", Retryable: true} + return "", err } if message == nil || message.ID <= 0 { return "", &outbox.DeliveryError{Class: "provider_invalid_receipt", Retryable: false} @@ -60,6 +88,101 @@ func (p *Provider) Deliver(ctx context.Context, value runtimestorage.ReplyOutbox return receipt, nil } +func (p *Provider) deliverMessage(ctx context.Context, value runtimestorage.ReplyOutbox) (*models.Message, error) { + if value.Kind == "" || value.Kind == runtimestorage.ReplyKindText { + return p.sendText(ctx, value.Payload) + } + normalized, err := runtimestorage.NormalizeReplyOutbox(value) + if err != nil { + return nil, &outbox.DeliveryError{Class: "invalid", Retryable: false} + } + switch normalized.Kind { + case runtimestorage.ReplyKindImage: + return p.sendPhoto(ctx, normalized) + case runtimestorage.ReplyKindDocument: + return p.sendDocument(ctx, normalized) + default: + return p.sendText(ctx, normalized.Fallback) + } +} + +func (p *Provider) sendPhoto(ctx context.Context, value runtimestorage.ReplyOutbox) (*models.Message, error) { + sender, ok := p.client.(telegramPhotoSender) + if !ok || p.attachments == nil { + return p.sendText(ctx, value.Fallback) + } + file, ok, err := p.attachmentUpload(ctx, value) + if err != nil { + return nil, err + } + if !ok { + return p.sendText(ctx, value.Fallback) + } + message, err := sender.SendPhoto(ctx, &bot.SendPhotoParams{ChatID: p.chatID, MessageThreadID: p.threadID, Photo: file, Caption: value.Payload}) + if err != nil { + return nil, telegramDeliveryError(ctx, err) + } + return message, nil +} + +func (p *Provider) sendDocument(ctx context.Context, value runtimestorage.ReplyOutbox) (*models.Message, error) { + sender, ok := p.client.(telegramDocumentSender) + if !ok || p.attachments == nil { + return p.sendText(ctx, value.Fallback) + } + file, ok, err := p.attachmentUpload(ctx, value) + if err != nil { + return nil, err + } + if !ok { + return p.sendText(ctx, value.Fallback) + } + message, err := sender.SendDocument(ctx, &bot.SendDocumentParams{ChatID: p.chatID, MessageThreadID: p.threadID, Document: file, Caption: value.Payload}) + if err != nil { + return nil, telegramDeliveryError(ctx, err) + } + return message, nil +} + +func (p *Provider) attachmentUpload(ctx context.Context, value runtimestorage.ReplyOutbox) (*models.InputFileUpload, bool, error) { + content, err := p.attachments.Load(ctx, value.TenantID, value.EventID, value.Attachment) + if err != nil { + if errors.Is(err, context.Canceled) || errors.Is(ctx.Err(), context.Canceled) { + return nil, false, &outbox.DeliveryError{Class: "canceled", Retryable: true} + } + if errors.Is(err, context.DeadlineExceeded) || errors.Is(ctx.Err(), context.DeadlineExceeded) { + return nil, false, &outbox.DeliveryError{Class: "timeout", Retryable: true} + } + return nil, false, nil + } + if err := content.Validate(value.Attachment); err != nil { + return nil, false, nil + } + name := value.Attachment.Name + if name == "" { + name = value.Attachment.ID + } + return &models.InputFileUpload{Filename: name, Data: bytes.NewReader(content.Data)}, true, nil +} + +func (p *Provider) sendText(ctx context.Context, text string) (*models.Message, error) { + message, err := p.client.SendMessage(ctx, &bot.SendMessageParams{ChatID: p.chatID, MessageThreadID: p.threadID, Text: text}) + if err != nil { + return nil, telegramDeliveryError(ctx, err) + } + return message, nil +} + +func telegramDeliveryError(ctx context.Context, err error) error { + if errors.Is(err, context.Canceled) || errors.Is(ctx.Err(), context.Canceled) { + return &outbox.DeliveryError{Class: "canceled", Retryable: true} + } + if errors.Is(err, context.DeadlineExceeded) || errors.Is(ctx.Err(), context.DeadlineExceeded) { + return &outbox.DeliveryError{Class: "timeout", Retryable: true} + } + return &outbox.DeliveryError{Class: "provider_error", Retryable: true} +} + // Reconcile checks whether a previously attempted segment can be confirmed. func (p *Provider) Reconcile(_ context.Context, value runtimestorage.ReplyOutbox) (outbox.DeliveryStatus, string, error) { if p == nil { diff --git a/trpcservice/channels/telegram/provider_test.go b/trpcservice/channels/telegram/provider_test.go index 167ab979..bb4dac3d 100644 --- a/trpcservice/channels/telegram/provider_test.go +++ b/trpcservice/channels/telegram/provider_test.go @@ -2,9 +2,13 @@ package telegram import ( "context" + "crypto/sha256" + "encoding/hex" "errors" + "io" "testing" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/outbox" runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" "github.com/go-telegram/bot" @@ -12,10 +16,14 @@ import ( ) type providerBot struct { - message *models.Message - err error - params *bot.SendMessageParams - calls int + message *models.Message + err error + params *bot.SendMessageParams + photoParams *bot.SendPhotoParams + documentParams *bot.SendDocumentParams + calls int + photoCalls int + documentCalls int } func (b *providerBot) Start(context.Context) {} @@ -27,6 +35,36 @@ func (b *providerBot) SendMessage(_ context.Context, params *bot.SendMessagePara b.params = params return b.message, b.err } +func (b *providerBot) SendPhoto(_ context.Context, params *bot.SendPhotoParams) (*models.Message, error) { + b.photoCalls++ + b.photoParams = params + return b.message, b.err +} +func (b *providerBot) SendDocument(_ context.Context, params *bot.SendDocumentParams) (*models.Message, error) { + b.documentCalls++ + b.documentParams = params + return b.message, b.err +} + +type providerAttachmentReader struct { + content attachment.Content + err error + tenantID string + eventID string + reference attachment.Reference + calls int +} + +func (reader *providerAttachmentReader) Load(_ context.Context, tenantID, eventID string, reference attachment.Reference) (attachment.Content, error) { + reader.calls++ + reader.tenantID = tenantID + reader.eventID = eventID + reader.reference = reference + if reader.err != nil { + return attachment.Content{}, reader.err + } + return reader.content.Clone(), nil +} func TestProviderUsesStableReceiptAndReconcile(t *testing.T) { botClient := &providerBot{message: &models.Message{ID: 42}} @@ -51,6 +89,218 @@ func TestProviderUsesStableReceiptAndReconcile(t *testing.T) { } } +func TestProviderDeliversImageAttachmentNatively(t *testing.T) { + data := []byte("png") + reference := providerAttachmentReference(t, attachment.KindImage, "image/png", "chart.png", data) + botClient := &providerBot{message: &models.Message{ID: 77}} + reader := &providerAttachmentReader{content: attachment.Content{Data: data}} + provider, err := NewProvider(botClient, 99, 7, WithAttachmentReader(reader)) + if err != nil { + t.Fatal(err) + } + reply := runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-image", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Attachment: reference, Fallback: "[image attachment: chart.png]", + } + receipt, err := provider.Deliver(context.Background(), reply) + if err != nil || receipt != "77" { + t.Fatalf("image deliver = %q, %v", receipt, err) + } + if botClient.photoCalls != 1 || botClient.calls != 0 || reader.calls != 1 { + t.Fatalf("calls text=%d photo=%d reader=%d", botClient.calls, botClient.photoCalls, reader.calls) + } + if reader.tenantID != reply.TenantID || reader.eventID != reply.EventID || reader.reference != reference { + t.Fatalf("reader args = %q %q %+v", reader.tenantID, reader.eventID, reader.reference) + } + if botClient.photoParams.ChatID != int64(99) || botClient.photoParams.MessageThreadID != 7 || botClient.photoParams.Caption != "caption" { + t.Fatalf("photo params = %+v", botClient.photoParams) + } + upload, ok := botClient.photoParams.Photo.(*models.InputFileUpload) + if !ok || upload.Filename != "chart.png" { + t.Fatalf("photo upload = %#v", botClient.photoParams.Photo) + } + if got, err := io.ReadAll(upload.Data); err != nil || string(got) != string(data) { + t.Fatalf("photo data = %q, %v", got, err) + } + if _, err := provider.Deliver(context.Background(), reply); err != nil || botClient.photoCalls != 1 { + t.Fatalf("idempotent image delivery calls=%d err=%v", botClient.photoCalls, err) + } +} + +func TestProviderDeliversDocumentAttachmentNatively(t *testing.T) { + data := []byte("document") + reference := providerAttachmentReference(t, attachment.KindDocument, "application/pdf", "brief.pdf", data) + botClient := &providerBot{message: &models.Message{ID: 88}} + provider, err := NewProvider(botClient, 99, 0, WithAttachmentReader(&providerAttachmentReader{content: attachment.Content{Data: data}})) + if err != nil { + t.Fatal(err) + } + receipt, err := provider.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-document", ReplyID: "reply-document", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindDocument, Payload: "caption", Attachment: reference, Fallback: "[document attachment: brief.pdf]", + }) + if err != nil || receipt != "88" || botClient.documentCalls != 1 || botClient.calls != 0 { + t.Fatalf("document deliver = %q text=%d document=%d err=%v", receipt, botClient.calls, botClient.documentCalls, err) + } + if botClient.documentParams.Caption != "caption" { + t.Fatalf("document caption = %q", botClient.documentParams.Caption) + } + upload, ok := botClient.documentParams.Document.(*models.InputFileUpload) + if !ok || upload.Filename != "brief.pdf" { + t.Fatalf("document upload = %#v", botClient.documentParams.Document) + } +} + +func TestProviderFallsBackForUnsupportedOrUnavailableMedia(t *testing.T) { + data := []byte("mp4") + reference := providerAttachmentReference(t, attachment.KindVideo, "video/mp4", "clip.mp4", data) + botClient := &providerBot{message: &models.Message{ID: 99}} + provider, err := NewProvider(botClient, 9, 0) + if err != nil { + t.Fatal(err) + } + receipt, err := provider.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-video", ReplyID: "reply-video", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindVideo, Payload: "caption", Attachment: reference, Fallback: "[video attachment: clip.mp4]", + }) + if err != nil || receipt != "99" || botClient.calls != 1 || botClient.photoCalls != 0 || botClient.params.Text != "[video attachment: clip.mp4]" { + t.Fatalf("fallback deliver = %q text=%d photo=%d params=%+v err=%v", receipt, botClient.calls, botClient.photoCalls, botClient.params, err) + } +} + +func TestProviderImageWithoutReaderFallsBackToText(t *testing.T) { + imageData := []byte("png") + imageReference := providerAttachmentReference(t, attachment.KindImage, "image/png", "chart.png", imageData) + botClient := &providerBot{message: &models.Message{ID: 11}} + provider, err := NewProvider(botClient, 9, 0) + if err != nil { + t.Fatal(err) + } + receipt, err := provider.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-image", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Attachment: imageReference, Fallback: "[image attachment: chart.png]", + }) + if err != nil || receipt != "11" || botClient.photoCalls != 0 || botClient.calls != 1 || botClient.params.Text != "[image attachment: chart.png]" { + t.Fatalf("fallback receipt=%q text=%d photo=%d params=%+v err=%v", receipt, botClient.calls, botClient.photoCalls, botClient.params, err) + } +} + +func TestProviderMissingAttachmentFallsBackToText(t *testing.T) { + imageData := []byte("png") + imageReference := providerAttachmentReference(t, attachment.KindImage, "image/png", "chart.png", imageData) + botClient := &providerBot{message: &models.Message{ID: 12}} + reader := &providerAttachmentReader{err: runtimestorage.ErrNotFound} + provider, err := NewProvider(botClient, 9, 0, WithAttachmentReader(reader)) + if err != nil { + t.Fatal(err) + } + receipt, err := provider.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-missing", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Attachment: imageReference, Fallback: "[image attachment: chart.png]", + }) + if err != nil || receipt != "12" || botClient.calls != 1 || botClient.photoCalls != 0 { + t.Fatalf("missing fallback receipt=%q text=%d photo=%d err=%v", receipt, botClient.calls, botClient.photoCalls, err) + } +} + +func TestProviderTamperedAttachmentFallsBackToText(t *testing.T) { + imageData := []byte("png") + imageReference := providerAttachmentReference(t, attachment.KindImage, "image/png", "chart.png", imageData) + botClient := &providerBot{message: &models.Message{ID: 13}} + reader := &providerAttachmentReader{content: attachment.Content{Data: []byte("tampered")}} + provider, err := NewProvider(botClient, 9, 0, WithAttachmentReader(reader)) + if err != nil { + t.Fatal(err) + } + receipt, err := provider.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-tampered", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Attachment: imageReference, Fallback: "[image attachment: chart.png]", + }) + if err != nil || receipt != "13" || botClient.calls != 1 || botClient.photoCalls != 0 { + t.Fatalf("tampered fallback receipt=%q text=%d photo=%d err=%v", receipt, botClient.calls, botClient.photoCalls, err) + } +} + +func TestProviderDocumentEmptyNameUsesReferenceID(t *testing.T) { + data := []byte("pdf") + reference := providerAttachmentReference(t, attachment.KindDocument, "application/pdf", "", data) + botClient := &providerBot{message: &models.Message{ID: 14}} + provider, err := NewProvider(botClient, 9, 0, WithAttachmentReader(&providerAttachmentReader{content: attachment.Content{Data: data}})) + if err != nil { + t.Fatal(err) + } + _, err = provider.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-doc", ReplyID: "reply-doc-empty-name", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindDocument, Payload: "caption", Attachment: reference, Fallback: "[document attachment]", + }) + if err != nil { + t.Fatal(err) + } + upload, ok := botClient.documentParams.Document.(*models.InputFileUpload) + if !ok || upload.Filename != reference.ID { + t.Fatalf("document upload = %#v", botClient.documentParams.Document) + } +} + +func TestProviderPreservesAttachmentCancellation(t *testing.T) { + data := []byte("png") + reference := providerAttachmentReference(t, attachment.KindImage, "image/png", "chart.png", data) + reader := &providerAttachmentReader{err: context.Canceled} + provider, err := NewProvider(&providerBot{message: &models.Message{ID: 1}}, 9, 0, WithAttachmentReader(reader)) + if err != nil { + t.Fatal(err) + } + _, err = provider.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-image", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Attachment: reference, Fallback: "[image attachment: chart.png]", + }) + var deliveryErr *outbox.DeliveryError + if !errors.As(err, &deliveryErr) || deliveryErr.Class != "canceled" || !deliveryErr.Retryable { + t.Fatalf("canceled attachment = %#v", err) + } +} + +func TestProviderClassifiesMediaTimeoutAndInvalidContract(t *testing.T) { + data := []byte("png") + reference := providerAttachmentReference(t, attachment.KindImage, "image/png", "chart.png", data) + timeoutProvider, err := NewProvider(&providerBot{message: &models.Message{ID: 1}}, 9, 0, WithAttachmentReader(&providerAttachmentReader{err: context.DeadlineExceeded})) + if err != nil { + t.Fatal(err) + } + _, err = timeoutProvider.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-timeout", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Attachment: reference, Fallback: "[image attachment: chart.png]", + }) + var deliveryErr *outbox.DeliveryError + if !errors.As(err, &deliveryErr) || deliveryErr.Class != "timeout" || !deliveryErr.Retryable { + t.Fatalf("timeout attachment = %#v", err) + } + + sendTimeout, err := NewProvider(&providerBot{message: &models.Message{ID: 1}, err: context.DeadlineExceeded}, 9, 0, WithAttachmentReader(&providerAttachmentReader{content: attachment.Content{Data: data}})) + if err != nil { + t.Fatal(err) + } + _, err = sendTimeout.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-send-timeout", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Attachment: reference, Fallback: "[image attachment: chart.png]", + }) + if !errors.As(err, &deliveryErr) || deliveryErr.Class != "timeout" || !deliveryErr.Retryable { + t.Fatalf("send timeout = %#v", err) + } + + invalidProvider, err := NewProvider(&providerBot{message: &models.Message{ID: 1}}, 9, 0) + if err != nil { + t.Fatal(err) + } + _, err = invalidProvider.Deliver(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-invalid", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Fallback: "[image attachment]", + }) + if !errors.As(err, &deliveryErr) || deliveryErr.Class != "invalid" || deliveryErr.Retryable { + t.Fatalf("invalid media contract = %#v", err) + } +} + func TestProviderRedactsAndClassifiesSendFailures(t *testing.T) { botClient := &providerBot{err: errors.New("secret provider response")} provider, err := NewProvider(botClient, 99, 0) @@ -117,3 +367,13 @@ func TestProviderValidationAndReceiptFailureBranches(t *testing.T) { t.Fatal("nil context deliver unexpectedly succeeded") } } + +func providerAttachmentReference(t *testing.T, kind attachment.Kind, contentType, name string, data []byte) attachment.Reference { + t.Helper() + digest := sha256.Sum256(data) + reference := attachment.Reference{ID: "attachment-" + string(kind), Kind: kind, MIMEType: contentType, Name: name, Size: int64(len(data)), SHA256: hex.EncodeToString(digest[:])} + if _, err := reference.Normalize(); err != nil { + t.Fatalf("attachment reference = %v", err) + } + return reference +} diff --git a/trpcservice/channels/telegram/telegram.go b/trpcservice/channels/telegram/telegram.go index 582ac9bd..5089290d 100644 --- a/trpcservice/channels/telegram/telegram.go +++ b/trpcservice/channels/telegram/telegram.go @@ -3,9 +3,14 @@ package telegram import ( + "bytes" "context" + "crypto/sha256" + "encoding/hex" "errors" "fmt" + "io" + "mime" "net/http" "net/url" "strconv" @@ -13,23 +18,27 @@ import ( "sync" "time" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/audit" "github.com/XnLemon/trpc-agent-service/trpcservice/channels" "github.com/XnLemon/trpc-agent-service/trpcservice/channels/replies" "github.com/XnLemon/trpc-agent-service/trpcservice/gateway" "github.com/XnLemon/trpc-agent-service/trpcservice/metrics" "github.com/XnLemon/trpc-agent-service/trpcservice/observability" + runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" "github.com/go-telegram/bot" "github.com/go-telegram/bot/models" ) const ( - defaultPollTimeout = time.Minute - minimumPollTimeout = 2 * time.Second - maximumPollTimeout = 10 * time.Minute - maximumTokenRunes = 1024 - maximumReplyRunes = 4096 - failureReply = "Sorry, I couldn't process that message." + defaultPollTimeout = time.Minute + minimumPollTimeout = 2 * time.Second + maximumPollTimeout = 10 * time.Minute + maximumTokenRunes = 1024 + maximumReplyRunes = 4096 + failureReply = "Sorry, I couldn't process that message." + defaultAttachmentBytes = 64 << 20 + maximumAttachmentBytes = 64 << 20 ) var ( @@ -58,6 +67,8 @@ var ( ErrSendMessage = errors.New("telegram send message failed") // ErrPolling reports a redacted SDK polling error delivered to ErrorHook. ErrPolling = errors.New("telegram polling failed") + // ErrAttachment reports a redacted media download or storage failure. + ErrAttachment = errors.New("telegram attachment processing failed") ) // ErrorOperation identifies the safe operation category supplied to ErrorHook. @@ -95,6 +106,92 @@ type BotClient interface { SendMessage(context.Context, *bot.SendMessageParams) (*models.Message, error) } +// MediaDownloader downloads one authenticated Telegram file into the adapter's +// attachment store. Implementations must not expose provider URLs or tokens to +// callers. +type MediaDownloader interface { + Download(context.Context, string) (io.ReadCloser, error) +} + +type telegramFileClient interface { + GetFile(context.Context, *bot.GetFileParams) (*models.File, error) + FileDownloadLink(*models.File) string +} + +type telegramMediaDownloader struct { + client telegramFileClient + httpClient bot.HttpClient + maximum int64 +} + +func (downloader telegramMediaDownloader) Download(ctx context.Context, fileID string) (io.ReadCloser, error) { + if ctx == nil || downloader.client == nil || downloader.httpClient == nil || fileID == "" || downloader.maximum < 1 { + return nil, ErrAttachment + } + if err := ctx.Err(); err != nil { + return nil, err + } + file, err := downloader.client.GetFile(ctx, &bot.GetFileParams{FileID: fileID}) + if err != nil { + return nil, telegramAttachmentError(ctx) + } + fileURL, err := downloader.fileDownloadURL(file) + if err != nil { + return nil, ErrAttachment + } + request, err := http.NewRequestWithContext(ctx, http.MethodGet, fileURL.String(), nil) + if err != nil { + return nil, ErrAttachment + } + response, err := downloader.httpClient.Do(request) + if err != nil { + return nil, telegramAttachmentError(ctx) + } + data, err := readTelegramMediaResponse(ctx, response, downloader.maximum) + if err != nil { + return nil, err + } + return io.NopCloser(bytes.NewReader(data)), nil +} + +func (downloader telegramMediaDownloader) fileDownloadURL(file *models.File) (*url.URL, error) { + if file == nil || file.FilePath == "" { + return nil, ErrAttachment + } + fileURL, err := url.Parse(downloader.client.FileDownloadLink(file)) + if err != nil || fileURL.Scheme != "https" || fileURL.Host == "" || fileURL.User != nil || fileURL.RawQuery != "" || fileURL.Fragment != "" { + return nil, ErrAttachment + } + return fileURL, nil +} + +func readTelegramMediaResponse(ctx context.Context, response *http.Response, maximum int64) ([]byte, error) { + if response == nil || response.Body == nil { + return nil, ErrAttachment + } + defer response.Body.Close() + if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices || response.ContentLength > maximum { + return nil, ErrAttachment + } + data, err := io.ReadAll(io.LimitReader(response.Body, maximum+1)) + if err != nil { + return nil, telegramAttachmentError(ctx) + } + if int64(len(data)) > maximum { + return nil, ErrAttachment + } + return data, nil +} + +func telegramAttachmentError(ctx context.Context) error { + if ctx != nil { + if err := ctx.Err(); err != nil { + return err + } + } + return ErrAttachment +} + // BotFactoryConfig contains non-secret options for constructing one BotClient. // The token is passed separately to BotFactory.New and is never stored in this // configuration value. @@ -160,21 +257,33 @@ type Config struct { Factory BotFactory // Observability supplies provider-neutral trace and metric hooks. Observability observability.Provider + // Attachments is the explicit durable boundary for native inbound media. + // When nil, media keeps the legacy text-marker behavior. + Attachments runtimestorage.AttachmentStore + // MediaDownloader optionally replaces the authenticated Telegram file + // downloader. It is primarily useful for deterministic tests. + MediaDownloader MediaDownloader + // MaxAttachmentBytes bounds each downloaded media object. Zero defaults to + // the protocol-neutral attachment limit. + MaxAttachmentBytes int64 } // Adapter owns one trusted Telegram Binding and routes its updates through the // existing Gateway contracts. It does not create or cache a Runner directly. type Adapter struct { - client BotClient - dispatcher gateway.DispatchService - principal gateway.Principal - target channels.RoutingTarget - idempotency *gateway.IdempotencyStore - ownIdempotency bool - errorHook ErrorHook - audit audit.Recorder - telemetry observability.Provider - metrics metrics.Catalog + client BotClient + dispatcher gateway.DispatchService + principal gateway.Principal + target channels.RoutingTarget + idempotency *gateway.IdempotencyStore + ownIdempotency bool + errorHook ErrorHook + audit audit.Recorder + telemetry observability.Provider + metrics metrics.Catalog + attachments runtimestorage.AttachmentStore + mediaDownloader MediaDownloader + maxAttachmentBytes int64 mu sync.RWMutex closed bool @@ -184,13 +293,14 @@ type Adapter struct { var _ channels.PollingAdapter = (*Adapter)(nil) type normalizedConfig struct { - token string - target channels.RoutingTarget - principal gateway.Principal - apiBaseURL string - pollTimeout time.Duration - workers int - providerAcctID string + token string + target channels.RoutingTarget + principal gateway.Principal + apiBaseURL string + pollTimeout time.Duration + workers int + maxAttachmentBytes int64 + providerAcctID string } // New validates the trusted route, constructs the Bot client, and verifies its @@ -216,7 +326,8 @@ func New(ctx context.Context, config Config) (*Adapter, error) { adapter := &Adapter{ dispatcher: config.Dispatcher, principal: normalized.principal, target: normalized.target, idempotency: idempotency, ownIdempotency: ownIdempotency, errorHook: config.ErrorHook, - audit: audit.Recorder{Writer: config.AuditWriter, TenantID: normalized.target.TenantID}, + audit: audit.Recorder{Writer: config.AuditWriter, TenantID: normalized.target.TenantID}, + attachments: config.Attachments, maxAttachmentBytes: normalized.maxAttachmentBytes, } if config.Observability == nil { config.Observability = observability.NewNoopProvider() @@ -241,6 +352,12 @@ func New(ctx context.Context, config Config) (*Adapter, error) { _ = adapter.closeOwnedIdempotency() return nil, err } + adapter.mediaDownloader = config.MediaDownloader + if adapter.mediaDownloader == nil { + if fileClient, ok := client.(telegramFileClient); ok && config.Attachments != nil { + adapter.mediaDownloader = telegramMediaDownloader{client: fileClient, httpClient: configuredHTTPClient(config.HTTPClient, normalized.pollTimeout), maximum: normalized.maxAttachmentBytes} + } + } return adapter, nil } @@ -280,11 +397,15 @@ func normalizeConfig(ctx context.Context, config Config) (normalizedConfig, erro if err != nil { return normalizedConfig{}, err } + maxAttachmentBytes, err := normalizeAttachmentBytes(config.MaxAttachmentBytes) + if err != nil { + return normalizedConfig{}, err + } principal, err := gateway.NewChannelPrincipal(config.Target) if err != nil { return normalizedConfig{}, fmt.Errorf("%w: trusted principal is invalid", ErrInvalid) } - return normalizedConfig{token: token, target: config.Target, principal: principal, apiBaseURL: apiBaseURL, pollTimeout: pollTimeout, workers: workers, providerAcctID: strconv.FormatInt(providerAccountID, 10)}, nil + return normalizedConfig{token: token, target: config.Target, principal: principal, apiBaseURL: apiBaseURL, pollTimeout: pollTimeout, workers: workers, maxAttachmentBytes: maxAttachmentBytes, providerAcctID: strconv.FormatInt(providerAccountID, 10)}, nil } func (adapter *Adapter) verifyIdentity(ctx context.Context, providerAccountID string) error { @@ -397,7 +518,7 @@ func (adapter *Adapter) HandleUpdate(ctx context.Context, update *models.Update) if client == nil || adapter.idempotency == nil { return ErrNotReady } - message, err := normalizeUpdate(adapter.target, update) + message, err := adapter.normalizeUpdate(ctx, update) if err != nil { adapter.report(ErrorOperationUpdate, err) return err @@ -609,6 +730,186 @@ func normalizeUpdate(target channels.RoutingTarget, update *models.Update) (gate return normalized, nil } +func (adapter *Adapter) normalizeUpdate(ctx context.Context, update *models.Update) (gateway.InboundMessage, error) { + inbound, err := normalizeUpdate(adapter.target, update) + if err != nil || adapter.attachments == nil || adapter.mediaDownloader == nil || update == nil || update.Message == nil { + return inbound, err + } + references, err := adapter.ingestAttachments(ctx, inbound.ExternalMessageID, update.Message) + if err != nil { + return gateway.InboundMessage{}, err + } + if len(references) == 0 { + return inbound, nil + } + inbound.Attachments = references + inbound.ContentType = gateway.ContentTypeMedia + return inbound.Normalize() +} + +type telegramAttachment struct { + fileID string + kind attachment.Kind + mimeType string + name string +} + +func (adapter *Adapter) ingestAttachments(ctx context.Context, externalMessageID string, message *models.Message) ([]attachment.Reference, error) { + descriptors := nativeAttachments(message) + if len(descriptors) == 0 { + return nil, nil + } + references := make([]attachment.Reference, 0, len(descriptors)) + for index, descriptor := range descriptors { + if err := ctx.Err(); err != nil { + return nil, err + } + reader, err := adapter.mediaDownloader.Download(ctx, descriptor.fileID) + if err != nil { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return nil, err + } + return nil, ErrAttachment + } + if reader == nil { + return nil, ErrAttachment + } + data, readErr := io.ReadAll(io.LimitReader(reader, adapter.maxAttachmentBytes+1)) + closeErr := reader.Close() + if readErr != nil || closeErr != nil { + if contextErr := ctx.Err(); contextErr != nil { + return nil, contextErr + } + return nil, ErrAttachment + } + if int64(len(data)) == 0 || int64(len(data)) > adapter.maxAttachmentBytes { + return nil, ErrAttachment + } + upload := attachment.Upload{ + ID: attachmentID(externalMessageID, index, descriptor.fileID), Kind: descriptor.kind, + MIMEType: descriptor.mimeType, Name: descriptor.name, Size: int64(len(data)), + Provider: "telegram", ProviderID: descriptor.fileID, + } + reference, err := adapter.attachments.PutAttachment(ctx, adapter.target.TenantID, upload, bytes.NewReader(data)) + if err != nil { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return nil, err + } + return nil, ErrAttachment + } + references = append(references, reference) + } + return references, nil +} + +func nativeAttachments(message *models.Message) []telegramAttachment { + if message == nil { + return nil + } + if photo := largestPhoto(message.Photo); photo != nil { + return []telegramAttachment{{fileID: photo.FileID, kind: attachment.KindImage, mimeType: "image/jpeg", name: photo.FileID + ".jpg"}} + } + if message.Video != nil && strings.TrimSpace(message.Video.FileID) != "" { + return []telegramAttachment{{fileID: message.Video.FileID, kind: attachment.KindVideo, mimeType: mediaMIME(attachment.KindVideo, message.Video.MimeType), name: mediaName(message.Video.FileName, message.Video.FileID, ".mp4")}} + } + if message.Animation != nil && strings.TrimSpace(message.Animation.FileID) != "" { + return []telegramAttachment{{fileID: message.Animation.FileID, kind: attachment.KindVideo, mimeType: mediaMIME(attachment.KindVideo, message.Animation.MimeType), name: mediaName(message.Animation.FileName, message.Animation.FileID, ".mp4")}} + } + if message.Audio != nil && strings.TrimSpace(message.Audio.FileID) != "" { + return []telegramAttachment{{fileID: message.Audio.FileID, kind: attachment.KindAudio, mimeType: mediaMIME(attachment.KindAudio, message.Audio.MimeType), name: mediaName(message.Audio.FileName, message.Audio.FileID, ".mp3")}} + } + if message.Voice != nil && strings.TrimSpace(message.Voice.FileID) != "" { + return []telegramAttachment{{fileID: message.Voice.FileID, kind: attachment.KindAudio, mimeType: mediaMIME(attachment.KindAudio, message.Voice.MimeType), name: message.Voice.FileID + ".ogg"}} + } + if message.Document != nil && strings.TrimSpace(message.Document.FileID) != "" { + mimeType := strings.TrimSpace(strings.ToLower(message.Document.MimeType)) + kind := attachment.KindDocument + if strings.HasPrefix(mimeType, "image/") { + kind = attachment.KindImage + } else if strings.HasPrefix(mimeType, "video/") { + kind = attachment.KindVideo + } else if strings.HasPrefix(mimeType, "audio/") { + kind = attachment.KindAudio + } + if !validMIME(mimeType) || !kindSupportsMIME(kind, mimeType) { + mimeType = "application/octet-stream" + kind = attachment.KindDocument + } + return []telegramAttachment{{fileID: message.Document.FileID, kind: kind, mimeType: mimeType, name: mediaName(message.Document.FileName, message.Document.FileID, "")}} + } + if message.VideoNote != nil && strings.TrimSpace(message.VideoNote.FileID) != "" { + return []telegramAttachment{{fileID: message.VideoNote.FileID, kind: attachment.KindVideo, mimeType: "video/mp4", name: message.VideoNote.FileID + ".mp4"}} + } + return nil +} + +func largestPhoto(photos []models.PhotoSize) *models.PhotoSize { + var largest *models.PhotoSize + for index := range photos { + photo := &photos[index] + if strings.TrimSpace(photo.FileID) == "" { + continue + } + if largest == nil || photo.FileSize > largest.FileSize || photo.Width*photo.Height > largest.Width*largest.Height { + largest = photo + } + } + return largest +} + +func mediaMIME(kind attachment.Kind, value string) string { + value = strings.ToLower(strings.TrimSpace(value)) + if validMIME(value) && kindSupportsMIME(kind, value) { + return value + } + switch kind { + case attachment.KindVideo: + return "video/mp4" + case attachment.KindAudio: + return "audio/mpeg" + default: + return "application/octet-stream" + } +} + +func validMIME(value string) bool { + parsed, params, err := mime.ParseMediaType(value) + return err == nil && parsed == value && len(params) == 0 && strings.Contains(value, "/") +} + +func kindSupportsMIME(kind attachment.Kind, value string) bool { + switch kind { + case attachment.KindImage: + return strings.HasPrefix(value, "image/") + case attachment.KindVideo: + return strings.HasPrefix(value, "video/") + case attachment.KindAudio: + return strings.HasPrefix(value, "audio/") + case attachment.KindDocument: + return !strings.HasPrefix(value, "image/") && !strings.HasPrefix(value, "video/") && !strings.HasPrefix(value, "audio/") + default: + return false + } +} + +func mediaName(name, fileID, suffix string) string { + name = strings.TrimSpace(name) + if name == "" { + return fileID + suffix + } + for _, character := range name { + if character < 0x20 || character == 0x7f || character == '/' || character == '\\' { + return fileID + suffix + } + } + return name +} + +func attachmentID(externalMessageID string, ordinal int, providerID string) string { + digest := sha256.Sum256([]byte(encodeParts(externalMessageID, strconv.Itoa(ordinal), providerID))) + return "att_" + hex.EncodeToString(digest[:]) +} + func messageContent(message *models.Message) (string, string, bool) { if message == nil { return "", "", false @@ -799,6 +1100,16 @@ func normalizeWorkers(value int) (int, error) { return value, nil } +func normalizeAttachmentBytes(value int64) (int64, error) { + if value == 0 { + return defaultAttachmentBytes, nil + } + if value < 1 || value > maximumAttachmentBytes { + return 0, fmt.Errorf("%w: attachment size limit is outside supported bounds", ErrInvalid) + } + return value, nil +} + func hasControl(value string) bool { for _, character := range value { if character < 0x20 || character == 0x7f { diff --git a/trpcservice/channels/telegram/telegram_test.go b/trpcservice/channels/telegram/telegram_test.go index 36a940e8..3cb37ddb 100644 --- a/trpcservice/channels/telegram/telegram_test.go +++ b/trpcservice/channels/telegram/telegram_test.go @@ -1,22 +1,27 @@ package telegram import ( + "bytes" "context" "crypto/sha256" "encoding/hex" "errors" "fmt" + "io" "net/http" + "net/http/httptest" "strings" "sync" "testing" "time" "github.com/XnLemon/trpc-agent-service/trpcservice/agent" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/audit" "github.com/XnLemon/trpc-agent-service/trpcservice/channels" "github.com/XnLemon/trpc-agent-service/trpcservice/channels/inmemory" "github.com/XnLemon/trpc-agent-service/trpcservice/gateway" + attachmentmemory "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/inmemory" "github.com/XnLemon/trpc-agent-service/trpcservice/tenant" "github.com/go-telegram/bot" "github.com/go-telegram/bot/models" @@ -194,9 +199,13 @@ func TestNewRejectsNonTelegramTargetAndInvalidRuntimeOptions(t *testing.T) { target := newTrustedTarget(t, channels.ChannelTelegram, "options", "12345") for name, config := range map[string]Config{ - "negative worker": {BotToken: "token", Target: target, Dispatcher: &dispatchStub{}, Workers: -1}, - "short poll": {BotToken: "token", Target: target, Dispatcher: &dispatchStub{}, PollTimeout: time.Second}, - "http api": {BotToken: "token", Target: target, Dispatcher: &dispatchStub{}, APIBaseURL: "http://insecure.example"}, + "negative worker": {BotToken: "token", Target: target, Dispatcher: &dispatchStub{}, Workers: -1}, + "short poll": {BotToken: "token", Target: target, Dispatcher: &dispatchStub{}, PollTimeout: time.Second}, + "http api": {BotToken: "token", Target: target, Dispatcher: &dispatchStub{}, APIBaseURL: "http://insecure.example"}, + "negative attachment bytes": {BotToken: "token", Target: target, Dispatcher: &dispatchStub{}, MaxAttachmentBytes: -1}, + "oversized attachment bytes": { + BotToken: "token", Target: target, Dispatcher: &dispatchStub{}, MaxAttachmentBytes: maximumAttachmentBytes + 1, + }, } { t.Run(name, func(t *testing.T) { if _, err := New(context.Background(), config); !errors.Is(err, ErrInvalid) { @@ -569,6 +578,262 @@ func TestMessageContentAndHasMediaBranches(t *testing.T) { } } +func TestNativeMediaIsPersistedAndDispatchedAsReference(t *testing.T) { + target := newTrustedTarget(t, channels.ChannelTelegram, "native-media", "12345") + dispatcher := &dispatchStub{events: []gateway.DispatchEvent{{Type: gateway.DispatchEventDone, Done: true}}} + data := []byte("telegram-image") + adapter, err := New(context.Background(), Config{ + BotToken: "12345:runtime-secret", Target: target, Dispatcher: dispatcher, + Factory: &fakeFactory{client: &fakeBot{me: &models.User{ID: 12345, IsBot: true}}}, + Attachments: attachmentmemory.New(), MediaDownloader: fakeMediaDownloader{data: data}, + }) + if err != nil { + t.Fatal(err) + } + defer func() { _ = adapter.Close() }() + update := &models.Update{ID: 201, Message: &models.Message{ + ID: 201, From: &models.User{ID: 42}, Chat: models.Chat{ID: 100, Type: models.ChatTypePrivate}, + Caption: "please inspect", Photo: []models.PhotoSize{{FileID: "photo-small", FileSize: 3}, {FileID: "photo-large", FileSize: len(data)}}, + }} + if err := adapter.HandleUpdate(context.Background(), update); err != nil { + t.Fatalf("HandleUpdate = %v", err) + } + requests := dispatcher.requests() + if len(requests) != 1 || requests[0].Message.Content != "please inspect" || len(requests[0].Message.Attachments) != 1 { + t.Fatalf("dispatch request = %+v", requests) + } + reference := requests[0].Message.Attachments[0] + if reference.Kind != attachment.KindImage || reference.MIMEType != "image/jpeg" || reference.Provider != "telegram" || reference.ProviderID != "photo-large" || reference.Size != int64(len(data)) { + t.Fatalf("attachment reference = %+v", reference) + } + if !strings.HasPrefix(reference.ID, "att_") || len(reference.SHA256) != sha256.Size*2 { + t.Fatalf("attachment identity = %+v", reference) + } +} + +func TestNativeMediaFailuresAreRedactedAndCancellationPreserved(t *testing.T) { + target := newTrustedTarget(t, channels.ChannelTelegram, "native-media-errors", "12345") + base := Config{BotToken: "12345:runtime-secret", Target: target, Factory: &fakeFactory{client: &fakeBot{me: &models.User{ID: 12345, IsBot: true}}}, Attachments: attachmentmemory.New()} + tooLarge := base + tooLarge.Dispatcher = &dispatchStub{} + tooLarge.MediaDownloader = fakeMediaDownloader{data: []byte("123456"), err: nil} + tooLarge.MaxAttachmentBytes = 5 + adapter, err := New(context.Background(), tooLarge) + if err != nil { + t.Fatal(err) + } + if err := adapter.HandleUpdate(context.Background(), mediaUpdate(202, "oversized")); !errors.Is(err, ErrAttachment) || strings.Contains(err.Error(), "secret") { + t.Fatalf("oversized media err = %v", err) + } + if got := len(tooLarge.Dispatcher.(*dispatchStub).requests()); got != 0 { + t.Fatalf("oversized media reached dispatch: %d", got) + } + _ = adapter.Close() + + canceled := base + canceled.Dispatcher = &dispatchStub{} + ctx, cancel := context.WithCancel(context.Background()) + canceled.MediaDownloader = fakeMediaDownloader{wait: func(context.Context) { cancel() }} + adapter, err = New(context.Background(), canceled) + if err != nil { + t.Fatal(err) + } + defer func() { _ = adapter.Close() }() + if err := adapter.HandleUpdate(ctx, mediaUpdate(203, "cancel")); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled media err = %v", err) + } + if got := len(canceled.Dispatcher.(*dispatchStub).requests()); got != 0 { + t.Fatalf("canceled media reached dispatch: %d", got) + } +} + +func TestTelegramMediaDownloaderFetchesBoundedHTTPSFile(t *testing.T) { + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/file/photos/chart.jpg" { + t.Fatalf("download path = %s", r.URL.Path) + } + _, _ = io.WriteString(w, "telegram-bytes") + })) + defer server.Close() + client := &fakeTelegramFileClient{file: &models.File{FileID: "file-1", FilePath: "photos/chart.jpg"}, link: server.URL + "/file/photos/chart.jpg"} + downloader := telegramMediaDownloader{client: client, httpClient: server.Client(), maximum: 64} + body, err := downloader.Download(context.Background(), "file-1") + if err != nil { + t.Fatalf("Download = %v", err) + } + data, err := io.ReadAll(body) + closeErr := body.Close() + if err != nil || closeErr != nil || string(data) != "telegram-bytes" { + t.Fatalf("download body = %q read=%v close=%v", data, err, closeErr) + } + if client.fileID != "file-1" { + t.Fatalf("GetFile id = %q", client.fileID) + } +} + +func TestTelegramMediaDownloaderRejectsUnsafeOrUnavailableFiles(t *testing.T) { + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/status": + http.Error(w, "nope", http.StatusBadGateway) + case "/declared": + w.Header().Set("Content-Length", "99") + _, _ = io.WriteString(w, "x") + case "/large": + _, _ = io.WriteString(w, "123456") + default: + _, _ = io.WriteString(w, "ok") + } + })) + defer server.Close() + + for _, test := range []struct { + name string + ctx context.Context + downloader telegramMediaDownloader + fileID string + want error + }{ + {name: "nil context", downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{}, httpClient: server.Client(), maximum: 8}, fileID: "file", want: ErrAttachment}, + {name: "missing client", ctx: context.Background(), downloader: telegramMediaDownloader{httpClient: server.Client(), maximum: 8}, fileID: "file", want: ErrAttachment}, + {name: "empty file id", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{}, httpClient: server.Client(), maximum: 8}, want: ErrAttachment}, + {name: "zero maximum", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{}, httpClient: server.Client()}, fileID: "file", want: ErrAttachment}, + {name: "get file error", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{err: errors.New("secret token")}, httpClient: server.Client(), maximum: 8}, fileID: "file", want: ErrAttachment}, + {name: "missing file path", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{file: &models.File{FileID: "file"}}, httpClient: server.Client(), maximum: 8}, fileID: "file", want: ErrAttachment}, + {name: "plain http URL", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{file: &models.File{FilePath: "file"}, link: "http://example.com/file"}, httpClient: server.Client(), maximum: 8}, fileID: "file", want: ErrAttachment}, + {name: "URL with query", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{file: &models.File{FilePath: "file"}, link: server.URL + "/file?token=secret"}, httpClient: server.Client(), maximum: 8}, fileID: "file", want: ErrAttachment}, + {name: "transport error", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{file: &models.File{FilePath: "file"}, link: server.URL + "/file"}, httpClient: &http.Client{Transport: telegramRoundTripFunc(func(*http.Request) (*http.Response, error) { return nil, errors.New("transport secret") })}, maximum: 8}, fileID: "file", want: ErrAttachment}, + {name: "provider status", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{file: &models.File{FilePath: "file"}, link: server.URL + "/status"}, httpClient: server.Client(), maximum: 8}, fileID: "file", want: ErrAttachment}, + {name: "declared size", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{file: &models.File{FilePath: "file"}, link: server.URL + "/declared"}, httpClient: server.Client(), maximum: 8}, fileID: "file", want: ErrAttachment}, + {name: "body too large", ctx: context.Background(), downloader: telegramMediaDownloader{client: &fakeTelegramFileClient{file: &models.File{FilePath: "file"}, link: server.URL + "/large"}, httpClient: server.Client(), maximum: 5}, fileID: "file", want: ErrAttachment}, + } { + t.Run(test.name, func(t *testing.T) { + body, err := test.downloader.Download(test.ctx, test.fileID) + if body != nil { + _ = body.Close() + } + if !errors.Is(err, test.want) || strings.Contains(fmt.Sprint(err), "secret") { + t.Fatalf("Download error = %v, want %v", err, test.want) + } + }) + } + + canceled, cancel := context.WithCancel(context.Background()) + cancel() + downloader := telegramMediaDownloader{client: &fakeTelegramFileClient{file: &models.File{FilePath: "file"}, link: server.URL + "/file"}, httpClient: server.Client(), maximum: 8} + if _, err := downloader.Download(canceled, "file"); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled Download = %v", err) + } + if _, err := readTelegramMediaResponse(context.Background(), nil, 1); !errors.Is(err, ErrAttachment) { + t.Fatalf("nil response = %v", err) + } + if _, err := readTelegramMediaResponse(context.Background(), &http.Response{}, 1); !errors.Is(err, ErrAttachment) { + t.Fatalf("nil body response = %v", err) + } +} + +func TestNativeAttachmentsClassifyTelegramMediaFamilies(t *testing.T) { + for _, test := range []struct { + name string + message *models.Message + want telegramAttachment + }{ + {name: "nil message"}, + {name: "largest photo", message: &models.Message{Photo: []models.PhotoSize{{FileID: ""}, {FileID: "small", FileSize: 1, Width: 10, Height: 10}, {FileID: "large", FileSize: 2, Width: 2, Height: 2}}}, want: telegramAttachment{fileID: "large", kind: attachment.KindImage, mimeType: "image/jpeg", name: "large.jpg"}}, + {name: "video default mime and safe name", message: &models.Message{Video: &models.Video{FileID: "video-1", FileName: "bad/name", MimeType: "application/octet-stream"}}, want: telegramAttachment{fileID: "video-1", kind: attachment.KindVideo, mimeType: "video/mp4", name: "video-1.mp4"}}, + {name: "animation keeps mime", message: &models.Message{Animation: &models.Animation{FileID: "anim-1", FileName: "clip.mp4", MimeType: "video/mp4"}}, want: telegramAttachment{fileID: "anim-1", kind: attachment.KindVideo, mimeType: "video/mp4", name: "clip.mp4"}}, + {name: "audio default suffix", message: &models.Message{Audio: &models.Audio{FileID: "audio-1", MimeType: "audio/mpeg"}}, want: telegramAttachment{fileID: "audio-1", kind: attachment.KindAudio, mimeType: "audio/mpeg", name: "audio-1.mp3"}}, + {name: "voice default", message: &models.Message{Voice: &models.Voice{FileID: "voice-1", MimeType: "audio/ogg"}}, want: telegramAttachment{fileID: "voice-1", kind: attachment.KindAudio, mimeType: "audio/ogg", name: "voice-1.ogg"}}, + {name: "document image", message: &models.Message{Document: &models.Document{FileID: "doc-image", FileName: "scan.png", MimeType: "image/png"}}, want: telegramAttachment{fileID: "doc-image", kind: attachment.KindImage, mimeType: "image/png", name: "scan.png"}}, + {name: "document video", message: &models.Message{Document: &models.Document{FileID: "doc-video", FileName: "clip.mov", MimeType: "video/quicktime"}}, want: telegramAttachment{fileID: "doc-video", kind: attachment.KindVideo, mimeType: "video/quicktime", name: "clip.mov"}}, + {name: "document audio", message: &models.Message{Document: &models.Document{FileID: "doc-audio", FileName: "voice.mp3", MimeType: "audio/mpeg"}}, want: telegramAttachment{fileID: "doc-audio", kind: attachment.KindAudio, mimeType: "audio/mpeg", name: "voice.mp3"}}, + {name: "document invalid MIME fallback", message: &models.Message{Document: &models.Document{FileID: "doc-1", FileName: "bad/name", MimeType: "image/png; charset=utf-8"}}, want: telegramAttachment{fileID: "doc-1", kind: attachment.KindDocument, mimeType: "application/octet-stream", name: "doc-1"}}, + {name: "video note", message: &models.Message{VideoNote: &models.VideoNote{FileID: "round-1"}}, want: telegramAttachment{fileID: "round-1", kind: attachment.KindVideo, mimeType: "video/mp4", name: "round-1.mp4"}}, + } { + t.Run(test.name, func(t *testing.T) { + got := nativeAttachments(test.message) + if test.want == (telegramAttachment{}) { + if got != nil { + t.Fatalf("nativeAttachments = %+v, want nil", got) + } + return + } + if len(got) != 1 || got[0] != test.want { + t.Fatalf("nativeAttachments = %+v, want %+v", got, test.want) + } + }) + } + if got := nativeAttachments(&models.Message{Document: &models.Document{FileID: ""}}); got != nil { + t.Fatalf("empty document file ID = %+v", got) + } + if kindSupportsMIME(attachment.Kind("unknown"), "application/octet-stream") { + t.Fatal("unknown attachment kind matched MIME") + } + if kindSupportsMIME(attachment.KindDocument, "image/png") { + t.Fatal("document attachment accepted image MIME") + } + if got := mediaMIME(attachment.KindAudio, "application/octet-stream"); got != "audio/mpeg" { + t.Fatalf("audio fallback MIME = %q", got) + } + if got := mediaMIME(attachment.KindImage, "application/octet-stream"); got != "application/octet-stream" { + t.Fatalf("image fallback MIME = %q", got) + } +} + +type fakeMediaDownloader struct { + data []byte + err error + wait func(context.Context) +} + +func (downloader fakeMediaDownloader) Download(ctx context.Context, _ string) (io.ReadCloser, error) { + if downloader.wait != nil { + downloader.wait(ctx) + } + if err := ctx.Err(); err != nil { + return nil, err + } + if downloader.err != nil { + return nil, downloader.err + } + return io.NopCloser(bytes.NewReader(downloader.data)), nil +} + +func mediaUpdate(updateID int64, fileID string) *models.Update { + return &models.Update{ID: updateID, Message: &models.Message{ + ID: int(updateID), From: &models.User{ID: 42}, Chat: models.Chat{ID: 100, Type: models.ChatTypePrivate}, + Video: &models.Video{FileID: fileID}, + }} +} + +type fakeTelegramFileClient struct { + file *models.File + err error + link string + fileID string +} + +func (client *fakeTelegramFileClient) GetFile(ctx context.Context, params *bot.GetFileParams) (*models.File, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + client.fileID = params.FileID + if client.err != nil { + return nil, client.err + } + return client.file, nil +} + +func (client *fakeTelegramFileClient) FileDownloadLink(*models.File) string { + return client.link +} + +type telegramRoundTripFunc func(*http.Request) (*http.Response, error) + +func (roundTrip telegramRoundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { + return roundTrip(request) +} + func TestDispatchAndSendFailuresAreRedacted(t *testing.T) { target := newTrustedTarget(t, channels.ChannelTelegram, "failure", "12345") recorder := &errorRecorder{} diff --git a/trpcservice/channels/wecom/binding_provider.go b/trpcservice/channels/wecom/binding_provider.go index 9b668574..91ef81c4 100644 --- a/trpcservice/channels/wecom/binding_provider.go +++ b/trpcservice/channels/wecom/binding_provider.go @@ -8,6 +8,7 @@ import ( "sync" "time" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/channels" "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/outbox" storage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" @@ -25,6 +26,7 @@ type BindingProvider struct { HTTPClient *http.Client BaseURL string Now func() time.Time + Attachments attachment.Reader mu sync.Mutex providers map[string]*Provider @@ -82,7 +84,7 @@ func (p *BindingProvider) provider(ctx context.Context, value storage.ReplyOutbo return provider, nil } } - provider := &Provider{CorpID: binding.Protocol.WeCom.CorpID, AgentID: binding.Protocol.WeCom.AgentID, AppSecret: credentials.AppSecret, HTTPClient: p.HTTPClient, BaseURL: p.BaseURL, Now: p.Now} + provider := &Provider{CorpID: binding.Protocol.WeCom.CorpID, AgentID: binding.Protocol.WeCom.AgentID, AppSecret: credentials.AppSecret, HTTPClient: p.HTTPClient, BaseURL: p.BaseURL, Now: p.Now, Attachments: p.Attachments} p.providers[key] = provider return provider, nil } diff --git a/trpcservice/channels/wecom/provider.go b/trpcservice/channels/wecom/provider.go index 059ae4a9..f289d4fd 100644 --- a/trpcservice/channels/wecom/provider.go +++ b/trpcservice/channels/wecom/provider.go @@ -6,6 +6,8 @@ import ( "encoding/json" "errors" "io" + "mime" + "mime/multipart" "net/http" "net/url" "strconv" @@ -13,11 +15,12 @@ import ( "sync" "time" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/outbox" storage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" ) -// Provider delivers text replies through a WeCom self-built application. +// Provider delivers replies through a WeCom self-built application. // Access tokens are cached in memory and are never persisted or logged. type Provider struct { CorpID string @@ -26,6 +29,7 @@ type Provider struct { HTTPClient *http.Client BaseURL string Now func() time.Time + Attachments attachment.Reader mu sync.Mutex token string tokenExpiry time.Time @@ -36,23 +40,106 @@ var _ outbox.Provider = (*Provider)(nil) const maximumTextBytes = 2048 -// Deliver sends one durable text segment through the WeCom application API. +// HTTPMediaDownloader fetches verified WeCom media through the official API +// and keeps per-binding access-token caches in memory. +type HTTPMediaDownloader struct { + HTTPClient *http.Client + BaseURL string + Now func() time.Time + + mu sync.Mutex + providers map[string]*Provider +} + +var _ MediaDownloader = (*HTTPMediaDownloader)(nil) + +// Download resolves a short-lived access token from the verified Binding +// context and returns the provider response body for the handler to validate +// and persist. +func (downloader *HTTPMediaDownloader) Download(ctx context.Context, request MediaDownloadRequest) (io.ReadCloser, error) { + request, err := normalizeMediaDownloadRequest(request) + if err != nil || downloader == nil || ctx == nil { + return nil, ErrAttachment + } + if err := ctx.Err(); err != nil { + return nil, err + } + return downloader.provider(request).downloadMedia(ctx, request) +} + +func (downloader *HTTPMediaDownloader) provider(request MediaDownloadRequest) *Provider { + key := request.TenantID + "\x00" + request.BindingID + "\x00" + request.CorpID + "\x00" + request.AgentID + downloader.mu.Lock() + defer downloader.mu.Unlock() + if downloader.providers == nil { + downloader.providers = make(map[string]*Provider) + } + if provider := downloader.providers[key]; provider != nil && provider.AppSecret == request.AppSecret { + return provider + } + provider := &Provider{ + CorpID: request.CorpID, AgentID: request.AgentID, AppSecret: request.AppSecret, + HTTPClient: downloader.HTTPClient, BaseURL: downloader.BaseURL, Now: downloader.Now, + } + downloader.providers[key] = provider + return provider +} + +func normalizeMediaDownloadRequest(request MediaDownloadRequest) (MediaDownloadRequest, error) { + request.TenantID = strings.TrimSpace(request.TenantID) + request.BindingID = strings.TrimSpace(request.BindingID) + request.CorpID = strings.TrimSpace(request.CorpID) + request.AgentID = strings.TrimSpace(request.AgentID) + request.AppSecret = strings.TrimSpace(request.AppSecret) + request.MediaID = strings.TrimSpace(request.MediaID) + request.MIMEType = strings.ToLower(strings.TrimSpace(request.MIMEType)) + if request.MaximumBytes == 0 { + request.MaximumBytes = defaultAttachmentBytes + } + if request.TenantID == "" || request.BindingID == "" || request.CorpID == "" || request.AgentID == "" || request.AppSecret == "" || request.MediaID == "" || request.MaximumBytes < 1 || request.MaximumBytes > maximumAttachmentBytes { + return MediaDownloadRequest{}, ErrAttachment + } + if request.Kind != attachment.KindImage && request.Kind != attachment.KindDocument && request.Kind != attachment.KindAudio && request.Kind != attachment.KindVideo { + return MediaDownloadRequest{}, ErrAttachment + } + for _, value := range []string{request.TenantID, request.BindingID, request.CorpID, request.AgentID, request.AppSecret, request.MediaID, request.MIMEType} { + if hasControl(value) { + return MediaDownloadRequest{}, ErrAttachment + } + } + return request, nil +} + +// Deliver sends one durable reply segment through the WeCom application API. // //nolint:gocyclo func (p *Provider) Deliver(ctx context.Context, value storage.ReplyOutbox) (string, error) { if p == nil || strings.TrimSpace(p.CorpID) == "" || strings.TrimSpace(p.AgentID) == "" || strings.TrimSpace(p.AppSecret) == "" || ctx == nil { return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} } - if (value.ReplyTarget.ConversationKind != "direct" && value.ReplyTarget.ConversationKind != "group") || value.ReplyTarget.ReceiverID == "" { - return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} + normalized, err := normalizeDeliveryReply(value) + if err != nil { + return "", err } - if len([]byte(value.Payload)) == 0 || len([]byte(value.Payload)) > maximumTextBytes { + value = normalized + if (value.ReplyTarget.ConversationKind != "direct" && value.ReplyTarget.ConversationKind != "group") || value.ReplyTarget.ReceiverID == "" { return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} } agentID, parseErr := strconv.Atoi(strings.TrimSpace(p.AgentID)) if parseErr != nil || agentID <= 0 || strconv.Itoa(agentID) != strings.TrimSpace(p.AgentID) { return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} } + nativeMedia := nativeWeComMedia(value.Kind) && p.Attachments != nil + text := "" + if !nativeMedia { + text, err = deliveryText(value) + if err != nil { + return "", err + } + if len([]byte(text)) == 0 || len([]byte(text)) > maximumTextBytes { + return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} + } + } key := deliveryKey(value) p.mu.Lock() if value.ReplyID != "" && p.receipts != nil { @@ -66,24 +153,239 @@ func (p *Provider) Deliver(ctx context.Context, value storage.ReplyOutbox) (stri if err != nil { return "", err } - target := struct { - ToUser string `json:"touser,omitempty"` - ChatID string `json:"chatid,omitempty"` - MsgType string `json:"msgtype"` - AgentID int `json:"agentid"` - Text struct { - Content string `json:"content"` - } `json:"text"` - Safe int `json:"safe"` - }{MsgType: "text", AgentID: agentID, Text: struct { - Content string `json:"content"` - }{value.Payload}, Safe: 0} + var receipt string + if nativeMedia { + receipt, err = p.deliverMedia(ctx, token, agentID, value) + } else { + receipt, err = p.deliverText(ctx, token, agentID, value, text) + } + if err != nil { + return "", err + } + p.mu.Lock() + if value.ReplyID != "" { + if p.receipts == nil { + p.receipts = make(map[string]string) + } + p.receipts[key] = receipt + } + p.mu.Unlock() + return receipt, nil +} + +type wecomSendPayload struct { + ToUser string `json:"touser,omitempty"` + ChatID string `json:"chatid,omitempty"` + MsgType string `json:"msgtype"` + AgentID int `json:"agentid"` + Text *wecomTextPayload `json:"text,omitempty"` + Image *wecomMediaPayload `json:"image,omitempty"` + File *wecomMediaPayload `json:"file,omitempty"` + Safe int `json:"safe"` +} + +type wecomTextPayload struct { + Content string `json:"content"` +} + +type wecomMediaPayload struct { + MediaID string `json:"media_id"` +} + +func normalizeDeliveryReply(value storage.ReplyOutbox) (storage.ReplyOutbox, error) { + normalized, err := storage.NormalizeReplyOutbox(value) + if err != nil { + return storage.ReplyOutbox{}, &outbox.DeliveryError{Class: "invalid", Retryable: false} + } + return normalized, nil +} + +func (p *Provider) deliverText(ctx context.Context, token string, agentID int, value storage.ReplyOutbox, text string) (string, error) { + payload := newSendPayload(value, agentID, "text") + payload.Text = &wecomTextPayload{Content: text} + return p.sendMessage(ctx, token, payload) +} + +func (p *Provider) deliverMedia(ctx context.Context, token string, agentID int, value storage.ReplyOutbox) (string, error) { + content, err := p.Attachments.Load(ctx, value.TenantID, value.EventID, value.Attachment) + if err != nil { + return "", attachmentLoadError(ctx, err) + } + if err := content.Validate(value.Attachment); err != nil { + return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} + } + uploadType, msgType := wecomMediaTypes(value.Kind) + mediaID, err := p.uploadTempMedia(ctx, token, uploadType, value.Attachment.Name, content.Data) + if err != nil { + return "", err + } + payload := newSendPayload(value, agentID, msgType) + switch msgType { + case "image": + payload.Image = &wecomMediaPayload{MediaID: mediaID} + case "file": + payload.File = &wecomMediaPayload{MediaID: mediaID} + default: + return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} + } + return p.sendMessage(ctx, token, payload) +} + +func nativeWeComMedia(kind storage.ReplyKind) bool { + return kind == storage.ReplyKindImage || kind == storage.ReplyKindDocument +} + +func wecomMediaTypes(kind storage.ReplyKind) (string, string) { + switch kind { + case storage.ReplyKindImage: + return "image", "image" + case storage.ReplyKindDocument: + return "file", "file" + default: + return "", "" + } +} + +func newSendPayload(value storage.ReplyOutbox, agentID int, msgType string) wecomSendPayload { + payload := wecomSendPayload{MsgType: msgType, AgentID: agentID, Safe: 0} if value.ReplyTarget.ConversationKind == "group" { - target.ChatID = value.ReplyTarget.ReceiverID + payload.ChatID = value.ReplyTarget.ReceiverID } else { - target.ToUser = value.ReplyTarget.ReceiverID + payload.ToUser = value.ReplyTarget.ReceiverID + } + return payload +} + +func (p *Provider) uploadTempMedia(ctx context.Context, token, mediaType, name string, data []byte) (string, error) { + var body bytes.Buffer + writer := multipart.NewWriter(&body) + part, err := writer.CreateFormFile("media", attachmentFileName(name)) + if err != nil { + return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} + } + if _, err := part.Write(data); err != nil { + return "", &outbox.DeliveryError{Class: "provider_error", Retryable: true} + } + if err := writer.Close(); err != nil { + return "", &outbox.DeliveryError{Class: "provider_error", Retryable: true} + } + endpoint := p.baseURL() + "/cgi-bin/media/upload?access_token=" + url.QueryEscape(token) + "&type=" + url.QueryEscape(mediaType) + request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, &body) + if err != nil { + return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} + } + request.Header.Set("Content-Type", writer.FormDataContentType()) + response, err := p.client().Do(request) + if err != nil { + return "", transportDeliveryError(err) + } + defer response.Body.Close() + if response.StatusCode < 200 || response.StatusCode >= 300 { + return "", &outbox.DeliveryError{Class: "unavailable", Retryable: true} + } + var result struct { + ErrCode int `json:"errcode"` + MediaID string `json:"media_id"` + } + if err := json.NewDecoder(io.LimitReader(response.Body, 64<<10)).Decode(&result); err != nil { + return "", &outbox.DeliveryError{Class: "provider_error", Retryable: true} + } + if result.ErrCode != 0 { + return "", p.providerResultError(result.ErrCode, response.StatusCode) + } + if result.MediaID == "" { + return "", &outbox.DeliveryError{Class: "provider_error", Retryable: true} + } + return result.MediaID, nil +} + +func (p *Provider) downloadMedia(ctx context.Context, download MediaDownloadRequest) (io.ReadCloser, error) { + if p == nil || ctx == nil || strings.TrimSpace(download.MediaID) == "" || hasControl(download.MediaID) { + return nil, ErrAttachment + } + if err := ctx.Err(); err != nil { + return nil, err + } + token, err := p.accessToken(ctx) + if err != nil { + return nil, wecomAttachmentError(ctx) + } + endpoint := p.baseURL() + "/cgi-bin/media/get?access_token=" + url.QueryEscape(token) + "&media_id=" + url.QueryEscape(download.MediaID) + request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) + if err != nil { + return nil, ErrAttachment + } + response, err := p.client().Do(request) + if err != nil { + return nil, wecomAttachmentError(ctx) + } + data, err := readWeComMediaResponse(ctx, response, download) + if err != nil { + return nil, err + } + return io.NopCloser(bytes.NewReader(data)), nil +} + +func readWeComMediaResponse(ctx context.Context, response *http.Response, download MediaDownloadRequest) ([]byte, error) { + if response == nil || response.Body == nil { + return nil, ErrAttachment + } + defer response.Body.Close() + if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices { + return nil, ErrAttachment + } + contentType := strings.ToLower(strings.TrimSpace(response.Header.Get("Content-Type"))) + data, err := io.ReadAll(io.LimitReader(response.Body, download.MaximumBytes+1)) + if err != nil { + return nil, wecomAttachmentError(ctx) + } + if int64(len(data)) == 0 || int64(len(data)) > download.MaximumBytes || providerJSONError(contentType, data) || !downloadContentTypeMatches(download.Kind, contentType) { + return nil, ErrAttachment + } + return data, nil +} + +func wecomAttachmentError(ctx context.Context) error { + if ctx != nil { + if err := ctx.Err(); err != nil { + return err + } + } + return ErrAttachment +} + +func providerJSONError(contentType string, _ []byte) bool { + mediaType, _, err := mime.ParseMediaType(contentType) + if err != nil { + mediaType = strings.ToLower(strings.TrimSpace(contentType)) + } + return mediaType == "application/json" || strings.HasSuffix(mediaType, "+json") +} + +func downloadContentTypeMatches(kind attachment.Kind, contentType string) bool { + mediaType, _, err := mime.ParseMediaType(contentType) + if strings.TrimSpace(contentType) == "" || err != nil || mediaType == "application/octet-stream" { + return true + } + switch kind { + case attachment.KindImage: + return strings.HasPrefix(mediaType, "image/") + case attachment.KindVideo: + return strings.HasPrefix(mediaType, "video/") + case attachment.KindAudio: + return strings.HasPrefix(mediaType, "audio/") + case attachment.KindDocument: + return !strings.HasPrefix(mediaType, "image/") && !strings.HasPrefix(mediaType, "video/") && !strings.HasPrefix(mediaType, "audio/") + default: + return false + } +} + +func (p *Provider) sendMessage(ctx context.Context, token string, payload wecomSendPayload) (string, error) { + body, err := json.Marshal(payload) + if err != nil { + return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} } - body, _ := json.Marshal(target) endpoint := p.baseURL() + "/cgi-bin/message/send?access_token=" + url.QueryEscape(token) request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) if err != nil { @@ -92,13 +394,7 @@ func (p *Provider) Deliver(ctx context.Context, value storage.ReplyOutbox) (stri request.Header.Set("Content-Type", "application/json") response, err := p.client().Do(request) if err != nil { - if errors.Is(err, context.DeadlineExceeded) { - return "", &outbox.DeliveryError{Class: "timeout", Retryable: true} - } - if errors.Is(err, context.Canceled) { - return "", &outbox.DeliveryError{Class: "canceled", Retryable: true} - } - return "", &outbox.DeliveryError{Class: "unavailable", Retryable: true} + return "", transportDeliveryError(err) } defer response.Body.Close() if response.StatusCode < 200 || response.StatusCode >= 300 { @@ -106,34 +402,74 @@ func (p *Provider) Deliver(ctx context.Context, value storage.ReplyOutbox) (stri } var result struct { ErrCode int `json:"errcode"` - ErrMsg string `json:"errmsg"` MsgID string `json:"msgid"` } if err := json.NewDecoder(io.LimitReader(response.Body, 64<<10)).Decode(&result); err != nil { return "", &outbox.DeliveryError{Class: "provider_error", Retryable: true} } if result.ErrCode != 0 { - class, retryable := classifyWeCom(result.ErrCode, response.StatusCode) - if class == "unauthenticated" { - p.mu.Lock() - p.token = "" - p.tokenExpiry = time.Time{} - p.mu.Unlock() - } - return "", &outbox.DeliveryError{Class: class, Retryable: retryable} + return "", p.providerResultError(result.ErrCode, response.StatusCode) } if result.MsgID == "" { return "", &outbox.DeliveryError{Class: "provider_error", Retryable: true} } - p.mu.Lock() - if value.ReplyID != "" { - if p.receipts == nil { - p.receipts = make(map[string]string) - } - p.receipts[key] = result.MsgID + return result.MsgID, nil +} + +func attachmentFileName(name string) string { + if strings.TrimSpace(name) == "" { + return "attachment" } + return name +} + +func attachmentLoadError(ctx context.Context, err error) error { + if contextErr := ctx.Err(); contextErr != nil { + return transportDeliveryError(contextErr) + } + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return transportDeliveryError(err) + } + if errors.Is(err, storage.ErrNotFound) || errors.Is(err, attachment.ErrInvalid) { + return &outbox.DeliveryError{Class: "invalid", Retryable: false} + } + return &outbox.DeliveryError{Class: "unavailable", Retryable: true} +} + +func transportDeliveryError(err error) error { + if errors.Is(err, context.DeadlineExceeded) { + return &outbox.DeliveryError{Class: "timeout", Retryable: true} + } + if errors.Is(err, context.Canceled) { + return &outbox.DeliveryError{Class: "canceled", Retryable: true} + } + return &outbox.DeliveryError{Class: "unavailable", Retryable: true} +} + +func (p *Provider) providerResultError(code, status int) error { + class, retryable := classifyWeCom(code, status) + if class == "unauthenticated" { + p.clearToken() + } + return &outbox.DeliveryError{Class: class, Retryable: retryable} +} + +func (p *Provider) clearToken() { + p.mu.Lock() + p.token = "" + p.tokenExpiry = time.Time{} p.mu.Unlock() - return result.MsgID, nil +} + +func deliveryText(value storage.ReplyOutbox) (string, error) { + if value.Kind == "" || value.Kind == storage.ReplyKindText { + return value.Payload, nil + } + normalized, err := storage.NormalizeReplyOutbox(value) + if err != nil { + return "", &outbox.DeliveryError{Class: "invalid", Retryable: false} + } + return normalized.Fallback, nil } // Reconcile reports unknown because WeCom does not expose a stable receipt query for app text sends. diff --git a/trpcservice/channels/wecom/wecom.go b/trpcservice/channels/wecom/wecom.go index 4a1a56cc..064fb6e0 100644 --- a/trpcservice/channels/wecom/wecom.go +++ b/trpcservice/channels/wecom/wecom.go @@ -1,5 +1,5 @@ -// Package wecom implements the text-only HTTPS callback for WeCom self-built -// application Bindings. +// Package wecom implements HTTPS callbacks for WeCom self-built application +// Bindings. package wecom import ( @@ -8,6 +8,7 @@ import ( "crypto/aes" "crypto/cipher" "crypto/sha1" // #nosec G505 -- WeCom requires SHA-1 callback signatures. + "crypto/sha256" "crypto/subtle" "encoding/base64" "encoding/binary" @@ -17,16 +18,19 @@ import ( "io" "net/http" "sort" + "strconv" "strings" "sync" "time" "github.com/XnLemon/trpc-agent-service/trpcservice/agent" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/audit" "github.com/XnLemon/trpc-agent-service/trpcservice/channels" "github.com/XnLemon/trpc-agent-service/trpcservice/gateway" "github.com/XnLemon/trpc-agent-service/trpcservice/metrics" "github.com/XnLemon/trpc-agent-service/trpcservice/observability" + runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" "github.com/XnLemon/trpc-agent-service/trpcservice/tenant" "github.com/google/uuid" ) @@ -36,9 +40,21 @@ var ( ErrInvalid = errors.New("invalid wecom callback") // ErrVerification reports a failed WeCom callback signature or decryption check. ErrVerification = errors.New("wecom callback verification failed") + // ErrAttachment reports a redacted WeCom media ingestion failure. + ErrAttachment = errors.New("wecom attachment processing failed") ) -const wecomBlockSize = 32 +const ( + wecomBlockSize = 32 + defaultAttachmentBytes = 64 << 20 + maximumAttachmentBytes = 64 << 20 + wecomMediaContentRunes = 128 + wecomDefaultFileMIME = "application/octet-stream" + wecomDefaultImageMIME = "image/jpeg" + wecomDefaultVideoMIME = "video/mp4" + wecomDefaultVoiceMIME = "audio/amr" + wecomDefaultVoiceSuffix = ".amr" +) // Credentials is the private credential bundle for one Binding SecretRef. type Credentials struct { @@ -53,6 +69,27 @@ type CredentialResolver interface { Resolve(context.Context, channels.SecretScope) (Credentials, error) } +// MediaDownloadRequest carries the verified Binding context required to fetch +// one WeCom media object. It stays inside the channel boundary and is never +// passed to Runner. +type MediaDownloadRequest struct { + TenantID string + BindingID string + CorpID string + AgentID string + AppSecret string + MediaID string + Kind attachment.Kind + MIMEType string + MaximumBytes int64 +} + +// MediaDownloader downloads one authenticated WeCom media object into the +// handler-owned attachment store. It must not expose provider URLs or tokens. +type MediaDownloader interface { + Download(context.Context, MediaDownloadRequest) (io.ReadCloser, error) +} + // Config contains either a static callback target or the dependencies required // to resolve a current trusted Binding for each callback. type Config struct { @@ -60,11 +97,21 @@ type Config struct { EncodingAESKey string ReceiveID string AgentID string + AppSecret string RouteKey string Target channels.RoutingTarget Dispatcher gateway.DispatchService MaxBodyBytes int64 ExecutionTimeout time.Duration + // Attachments is the explicit durable boundary for native inbound media. + // When nil, media callback types are rejected fail-closed. + Attachments runtimestorage.AttachmentStore + // MediaDownloader performs authenticated provider media downloads before + // bytes enter the protocol-neutral attachment store. + MediaDownloader MediaDownloader + // MaxAttachmentBytes bounds each downloaded media object. Zero defaults to + // the protocol-neutral attachment limit. + MaxAttachmentBytes int64 Candidates channels.CandidateConsumer Tenants tenant.Repository @@ -80,6 +127,8 @@ type callbackState struct { token string receiveID string agentID string + corpID string + appSecret string key []byte principal gateway.Principal } @@ -91,20 +140,23 @@ type Handler struct { // These retained fields preserve the package's focused cryptographic tests // and are populated together with static. Dynamic callbacks use a local // verified state instead. - token, receiveID string - key []byte - routeKey string - dynamic bool - candidates channels.CandidateConsumer - tenants tenant.Repository - apps agent.Repository - credentials CredentialResolver - dispatcher gateway.DispatchService - maxBodyBytes int64 - executionTimeout time.Duration - auditWriter audit.Writer - telemetry observability.Provider - metrics metrics.Catalog + token, receiveID string + key []byte + routeKey string + dynamic bool + candidates channels.CandidateConsumer + tenants tenant.Repository + apps agent.Repository + credentials CredentialResolver + dispatcher gateway.DispatchService + maxBodyBytes int64 + executionTimeout time.Duration + attachments runtimestorage.AttachmentStore + mediaDownloader MediaDownloader + maxAttachmentBytes int64 + auditWriter audit.Writer + telemetry observability.Provider + metrics metrics.Catalog mu sync.Mutex closing bool @@ -115,7 +167,7 @@ type Handler struct { var _ channels.WebhookAdapter = (*Handler)(nil) -// New validates a text callback Handler. Dynamic mode receives the complete +// New validates a callback Handler. Dynamic mode receives the complete // trusted target only after protocol verification. // //nolint:gocyclo @@ -135,11 +187,20 @@ func New(config Config) (*Handler, error) { if config.ExecutionTimeout < 1 { return nil, ErrInvalid } + maxAttachmentBytes, err := normalizeAttachmentBytes(config.MaxAttachmentBytes) + if err != nil || (config.Attachments == nil) != (config.MediaDownloader == nil) { + return nil, ErrInvalid + } baseCtx, cancel := context.WithCancel(context.Background()) if config.Observability == nil { config.Observability = observability.NewNoopProvider() } - handler := &Handler{routeKey: strings.Trim(config.RouteKey, "/"), dispatcher: config.Dispatcher, maxBodyBytes: config.MaxBodyBytes, executionTimeout: config.ExecutionTimeout, auditWriter: config.AuditWriter, baseCtx: baseCtx, cancel: cancel} + handler := &Handler{ + routeKey: strings.Trim(config.RouteKey, "/"), dispatcher: config.Dispatcher, + maxBodyBytes: config.MaxBodyBytes, executionTimeout: config.ExecutionTimeout, + attachments: config.Attachments, mediaDownloader: config.MediaDownloader, maxAttachmentBytes: maxAttachmentBytes, + auditWriter: config.AuditWriter, baseCtx: baseCtx, cancel: cancel, + } handler.telemetry, handler.metrics = config.Observability, metrics.New(config.Observability) if config.Candidates != nil || config.Tenants != nil || config.Apps != nil || config.Credentials != nil { if config.Candidates == nil || config.Tenants == nil || config.Apps == nil || config.Credentials == nil || handler.routeKey != "" { @@ -154,6 +215,10 @@ func New(config Config) (*Handler, error) { cancel() return nil, ErrInvalid } + if config.Attachments != nil && strings.TrimSpace(config.AppSecret) == "" { + cancel() + return nil, ErrInvalid + } if err := config.Target.Validate(); err != nil || config.Target.Channel != channels.ChannelWeCom { cancel() return nil, ErrInvalid @@ -168,7 +233,11 @@ func New(config Config) (*Handler, error) { cancel() return nil, ErrInvalid } - handler.static = &callbackState{token: config.Token, receiveID: config.ReceiveID, agentID: strings.TrimSpace(config.AgentID), key: key, principal: principal} + handler.static = &callbackState{ + token: config.Token, receiveID: config.ReceiveID, agentID: strings.TrimSpace(config.AgentID), + corpID: config.Target.ProviderAccountID, appSecret: strings.TrimSpace(config.AppSecret), + key: key, principal: principal, + } handler.token, handler.receiveID, handler.key = config.Token, config.ReceiveID, key return handler, nil } @@ -265,7 +334,7 @@ func (h *Handler) handleMessage(w http.ResponseWriter, r *http.Request) { return } var message inboundXML - if err := xml.Unmarshal(plain, &message); err != nil || message.MsgType != "text" || strings.TrimSpace(message.Content) == "" || strings.TrimSpace(message.MsgID) == "" || strings.TrimSpace(message.FromUserName) == "" || strings.TrimSpace(message.AgentID) != state.agentID { + if err := xml.Unmarshal(plain, &message); err != nil || !validInboundEnvelope(message, state.agentID) || h.validateInboundMessage(message) != nil { http.Error(w, "bad request", http.StatusBadRequest) return } @@ -280,13 +349,10 @@ func (h *Handler) handleMessage(w http.ResponseWriter, r *http.Request) { go func() { defer h.drains.Done() defer cancel() - inbound := gateway.InboundMessage{Content: message.Content, ContentType: gateway.ContentTypeText, ExternalMessageID: message.MsgID, ExternalUserID: message.FromUserName} - if strings.TrimSpace(message.ChatID) != "" { - inbound.ConversationKind = channels.ConversationGroup - inbound.ExternalChatID = strings.TrimSpace(message.ChatID) - } else { - inbound.ConversationKind = channels.ConversationDirect - inbound.ExternalPeerID = message.FromUserName + inbound, buildErr := h.buildInboundMessage(executionCtx, state, message) + if buildErr != nil { + result <- buildErr + return } stream, dispatchErr := h.dispatcher.Dispatch(executionCtx, gateway.DispatchRequest{Accepted: accepted, Principal: state.principal, RequestID: requestID, TraceID: traceID, Message: inbound}) if dispatchErr == nil && stream != nil { @@ -319,6 +385,207 @@ func (h *Handler) handleMessage(w http.ResponseWriter, r *http.Request) { } } +func validInboundEnvelope(message inboundXML, agentID string) bool { + return strings.TrimSpace(message.MsgID) != "" && strings.TrimSpace(message.FromUserName) != "" && strings.TrimSpace(message.AgentID) == strings.TrimSpace(agentID) +} + +func (h *Handler) validateInboundMessage(message inboundXML) error { + switch normalizedMessageType(message.MsgType) { + case "text": + if strings.TrimSpace(message.Content) == "" { + return ErrInvalid + } + return nil + case "image", "file", "voice", "video": + if h == nil || h.attachments == nil || h.mediaDownloader == nil { + return ErrInvalid + } + _, err := wecomAttachmentDescriptor(message) + return err + default: + return ErrInvalid + } +} + +func (h *Handler) buildInboundMessage(ctx context.Context, state callbackState, message inboundXML) (gateway.InboundMessage, error) { + inbound := gateway.InboundMessage{ + ExternalMessageID: strings.TrimSpace(message.MsgID), + ExternalUserID: strings.TrimSpace(message.FromUserName), + } + if chatID := strings.TrimSpace(message.ChatID); chatID != "" { + inbound.ConversationKind = channels.ConversationGroup + inbound.ExternalChatID = chatID + } else { + inbound.ConversationKind = channels.ConversationDirect + inbound.ExternalPeerID = strings.TrimSpace(message.FromUserName) + } + if normalizedMessageType(message.MsgType) == "text" { + inbound.Content = message.Content + inbound.ContentType = gateway.ContentTypeText + return inbound.Normalize() + } + reference, err := h.ingestAttachment(ctx, state, message) + if err != nil { + return gateway.InboundMessage{}, err + } + inbound.Content = wecomMediaContent(reference) + inbound.ContentType = gateway.ContentTypeMedia + inbound.Attachments = []attachment.Reference{reference} + return inbound.Normalize() +} + +type wecomAttachment struct { + mediaID string + kind attachment.Kind + mimeType string + name string +} + +func (h *Handler) ingestAttachment(ctx context.Context, state callbackState, message inboundXML) (attachment.Reference, error) { + descriptor, err := wecomAttachmentDescriptor(message) + if err != nil { + return attachment.Reference{}, ErrAttachment + } + if err := ctx.Err(); err != nil { + return attachment.Reference{}, err + } + download := MediaDownloadRequest{ + TenantID: state.principal.TenantID(), + CorpID: state.corpID, + AgentID: state.agentID, + AppSecret: state.appSecret, + MediaID: descriptor.mediaID, + Kind: descriptor.kind, + MIMEType: descriptor.mimeType, + MaximumBytes: h.maxAttachmentBytes, + } + if target, ok := state.principal.RoutingTarget(); ok { + download.BindingID = target.BindingID + if download.CorpID == "" { + download.CorpID = target.ProviderAccountID + } + } + reader, err := h.mediaDownloader.Download(ctx, download) + if err != nil { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return attachment.Reference{}, err + } + return attachment.Reference{}, ErrAttachment + } + if reader == nil { + return attachment.Reference{}, ErrAttachment + } + data, readErr := io.ReadAll(io.LimitReader(reader, h.maxAttachmentBytes+1)) + closeErr := reader.Close() + if readErr != nil || closeErr != nil { + if contextErr := ctx.Err(); contextErr != nil { + return attachment.Reference{}, contextErr + } + return attachment.Reference{}, ErrAttachment + } + if int64(len(data)) == 0 || int64(len(data)) > h.maxAttachmentBytes { + return attachment.Reference{}, ErrAttachment + } + bindingID := download.BindingID + upload := attachment.Upload{ + ID: attachmentID(bindingID, strings.TrimSpace(message.MsgID), 0, descriptor.mediaID), + Kind: descriptor.kind, + MIMEType: descriptor.mimeType, + Name: descriptor.name, + Size: int64(len(data)), + Provider: "wecom", + ProviderID: descriptor.mediaID, + } + reference, err := h.attachments.PutAttachment(ctx, state.principal.TenantID(), upload, bytes.NewReader(data)) + if err != nil { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return attachment.Reference{}, err + } + return attachment.Reference{}, ErrAttachment + } + return reference, nil +} + +func wecomAttachmentDescriptor(message inboundXML) (wecomAttachment, error) { + mediaID := strings.TrimSpace(message.MediaID) + if mediaID == "" { + return wecomAttachment{}, ErrInvalid + } + switch normalizedMessageType(message.MsgType) { + case "image": + return wecomAttachment{mediaID: mediaID, kind: attachment.KindImage, mimeType: wecomDefaultImageMIME, name: mediaName("", mediaID, ".jpg")}, nil + case "file": + return wecomAttachment{mediaID: mediaID, kind: attachment.KindDocument, mimeType: wecomDefaultFileMIME, name: mediaName(message.FileName, mediaID, "")}, nil + case "voice": + mimeType, suffix := voiceMIME(message.Format) + return wecomAttachment{mediaID: mediaID, kind: attachment.KindAudio, mimeType: mimeType, name: mediaName("", mediaID, suffix)}, nil + case "video": + return wecomAttachment{mediaID: mediaID, kind: attachment.KindVideo, mimeType: wecomDefaultVideoMIME, name: mediaName("", mediaID, ".mp4")}, nil + default: + return wecomAttachment{}, ErrInvalid + } +} + +func normalizedMessageType(value string) string { + return strings.ToLower(strings.TrimSpace(value)) +} + +func voiceMIME(format string) (string, string) { + switch strings.ToLower(strings.TrimSpace(format)) { + case "mp3": + return "audio/mpeg", ".mp3" + case "wav": + return "audio/wav", ".wav" + case "m4a": + return "audio/mp4", ".m4a" + case "ogg": + return "audio/ogg", ".ogg" + case "speex": + return "audio/speex", ".speex" + default: + return wecomDefaultVoiceMIME, wecomDefaultVoiceSuffix + } +} + +func mediaName(name, providerID, suffix string) string { + name = strings.TrimSpace(name) + if name == "" { + return providerID + suffix + } + for _, character := range name { + if character < 0x20 || character == 0x7f || character == '/' || character == '\\' { + return providerID + suffix + } + } + return name +} + +func wecomMediaContent(reference attachment.Reference) string { + base := "[wecom " + string(reference.Kind) + " attachment" + if reference.Name != "" { + withName := base + ": " + reference.Name + "]" + if len([]rune(withName)) <= wecomMediaContentRunes { + return withName + } + } + return base + "]" +} + +func attachmentID(bindingID, externalMessageID string, ordinal int, providerID string) string { + digest := sha256.Sum256([]byte(encodeParts(bindingID, externalMessageID, strconv.Itoa(ordinal), providerID))) + return "att_" + hex.EncodeToString(digest[:]) +} + +func encodeParts(parts ...string) string { + var builder strings.Builder + for _, part := range parts { + builder.WriteString(strconv.Itoa(len([]byte(part)))) + builder.WriteByte(':') + builder.WriteString(part) + } + return builder.String() +} + func (h *Handler) tryAcceptedIngress(accepted <-chan struct{}, w http.ResponseWriter, ctx context.Context, principal gateway.Principal, message inboundXML, requestID, traceID string) bool { select { case <-accepted: @@ -461,7 +728,11 @@ func (h *Handler) verify(r *http.Request, ciphertext string) ([]byte, callbackSt if decryptErr != nil { return decryptErr } - verifiedState = callbackState{token: credentials.CallbackToken, receiveID: binding.Protocol.WeCom.ReceiveID, agentID: binding.Protocol.WeCom.AgentID, key: key} + verifiedState = callbackState{ + token: credentials.CallbackToken, receiveID: binding.Protocol.WeCom.ReceiveID, + agentID: binding.Protocol.WeCom.AgentID, corpID: binding.Protocol.WeCom.CorpID, + appSecret: credentials.AppSecret, key: key, + } verifiedPlain = plain return nil }) @@ -504,6 +775,26 @@ func validSignature(token, signature, timestamp, nonce, ciphertext string) bool want := hex.EncodeToString(sum[:]) return subtle.ConstantTimeCompare([]byte(signature), []byte(want)) == 1 } + +func normalizeAttachmentBytes(value int64) (int64, error) { + if value == 0 { + return defaultAttachmentBytes, nil + } + if value < 1 || value > maximumAttachmentBytes { + return 0, ErrInvalid + } + return value, nil +} + +func hasControl(value string) bool { + for _, character := range value { + if character < 0x20 || character == 0x7f { + return true + } + } + return false +} + func decodeAESKey(value string) ([]byte, error) { key, err := base64.StdEncoding.DecodeString(value + "=") if err != nil || len(key) != 32 { @@ -564,6 +855,10 @@ type inboundXML struct { MsgType string `xml:"MsgType"` AgentID string `xml:"AgentID"` Content string `xml:"Content"` + MediaID string `xml:"MediaId"` + PicURL string `xml:"PicUrl"` + Format string `xml:"Format"` + FileName string `xml:"FileName"` } var _ http.Handler = (*Handler)(nil) diff --git a/trpcservice/channels/wecom/wecom_test.go b/trpcservice/channels/wecom/wecom_test.go index f9edd991..919ae93d 100644 --- a/trpcservice/channels/wecom/wecom_test.go +++ b/trpcservice/channels/wecom/wecom_test.go @@ -7,6 +7,7 @@ import ( "crypto/cipher" "crypto/rand" "crypto/sha1" // #nosec G505 -- WeCom requires SHA-1 callback signatures. + "crypto/sha256" "encoding/base64" "encoding/binary" "encoding/hex" @@ -22,11 +23,13 @@ import ( "time" "github.com/XnLemon/trpc-agent-service/trpcservice/agent" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/audit" "github.com/XnLemon/trpc-agent-service/trpcservice/channels" "github.com/XnLemon/trpc-agent-service/trpcservice/gateway" "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/outbox" storage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" + attachmentmemory "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/inmemory" "github.com/XnLemon/trpc-agent-service/trpcservice/tenant" ) @@ -56,6 +59,7 @@ func TestDecodeAESKeyAndDecryptRoundTrip(t *testing.T) { func TestNewRejectsInvalidConfiguration(t *testing.T) { stub := &callbackDispatchStub{requests: make(chan gateway.DispatchRequest, 1)} + staticTarget := staticTestTarget(t) for _, config := range []Config{ {}, {Dispatcher: stub, MaxBodyBytes: -1}, @@ -63,6 +67,7 @@ func TestNewRejectsInvalidConfiguration(t *testing.T) { {Dispatcher: stub, Token: "token", ReceiveID: "receive", AgentID: "1"}, {Dispatcher: stub, Token: "token", ReceiveID: "receive", AgentID: "1", EncodingAESKey: "bad", Target: channels.RoutingTarget{}}, {Dispatcher: stub, Token: "token", ReceiveID: "receive", AgentID: "1", EncodingAESKey: base64.RawStdEncoding.EncodeToString(bytes.Repeat([]byte{1}, 32)), Target: channels.RoutingTarget{}, Candidates: &dynamicCandidateConsumer{}}, + {Dispatcher: stub, Token: "token", ReceiveID: "receive", AgentID: "1", EncodingAESKey: base64.RawStdEncoding.EncodeToString(bytes.Repeat([]byte{1}, 32)), Target: staticTarget, Attachments: attachmentmemory.New(), MediaDownloader: &fakeWeComMediaDownloader{}}, } { if handler, err := New(config); handler != nil || !errors.Is(err, ErrInvalid) { t.Fatalf("invalid config = handler %v, err %v", handler, err) @@ -233,6 +238,549 @@ func TestProviderDeliversGroupChat(t *testing.T) { } } +func TestProviderSendsMediaReplyFallbackText(t *testing.T) { + var payload map[string]any + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + if r.URL.Path == "/cgi-bin/gettoken" { + _, _ = io.WriteString(w, `{"errcode":0,"access_token":"token","expires_in":3600}`) + return + } + if r.URL.Path == "/cgi-bin/message/send" { + _ = json.NewDecoder(r.Body).Decode(&payload) + _, _ = io.WriteString(w, `{"errcode":0,"msgid":"media-fallback-1"}`) + return + } + http.NotFound(w, r) + })) + defer server.Close() + p := &Provider{CorpID: "corp", AgentID: "1", AppSecret: "secret", BaseURL: server.URL, HTTPClient: server.Client()} + value := storage.ReplyOutbox{ + Payload: "caption", Kind: storage.ReplyKindImage, + Attachment: wecomReplyReference(t, attachment.KindImage, "image/png", "chart.png", []byte("png")), + Fallback: "[image attachment: chart.png]", + ReplyTarget: storage.ReplyTarget{ + ConversationKind: "direct", + ReceiverID: "user-1", + }, + } + if id, err := p.Deliver(context.Background(), value); err != nil || id != "media-fallback-1" { + t.Fatalf("media fallback deliver = %q, %v", id, err) + } + text, ok := payload["text"].(map[string]any) + if payload["msgtype"] != "text" || !ok || text["content"] != "[image attachment: chart.png]" { + t.Fatalf("media fallback payload = %#v", payload) + } +} + +//nolint:gocyclo // Covers the complete native upload/send contract for both supported media kinds. +func TestProviderSendsNativeMediaReply(t *testing.T) { + for _, test := range []struct { + name string + kind storage.ReplyKind + attachmentKind attachment.Kind + mimeType string + fileName string + data []byte + uploadType string + messageType string + fallback string + }{ + {name: "image", kind: storage.ReplyKindImage, attachmentKind: attachment.KindImage, mimeType: "image/png", fileName: "chart.png", data: []byte("png"), uploadType: "image", messageType: "image", fallback: "[image attachment: chart.png]"}, + {name: "document", kind: storage.ReplyKindDocument, attachmentKind: attachment.KindDocument, mimeType: "application/pdf", fileName: "brief.pdf", data: []byte("pdf"), uploadType: "file", messageType: "file", fallback: "[document attachment: brief.pdf]"}, + } { + t.Run(test.name, func(t *testing.T) { + var payload map[string]any + var uploadCalls int + var sendCalls int + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/cgi-bin/gettoken": + _, _ = io.WriteString(w, `{"errcode":0,"access_token":"token","expires_in":3600}`) + case "/cgi-bin/media/upload": + uploadCalls++ + if r.URL.Query().Get("access_token") != "token" || r.URL.Query().Get("type") != test.uploadType { + t.Errorf("upload query = %s", r.URL.RawQuery) + } + if err := r.ParseMultipartForm(1 << 20); err != nil { + t.Errorf("parse upload = %v", err) + return + } + files := r.MultipartForm.File["media"] + if len(files) != 1 || files[0].Filename != test.fileName { + t.Errorf("upload files = %#v", files) + return + } + file, err := files[0].Open() + if err != nil { + t.Errorf("open upload = %v", err) + return + } + defer func() { _ = file.Close() }() + data, err := io.ReadAll(file) + if err != nil || string(data) != string(test.data) { + t.Errorf("upload data = %q, %v", data, err) + return + } + _, _ = io.WriteString(w, `{"errcode":0,"media_id":"uploaded-`+test.messageType+`"}`) + case "/cgi-bin/message/send": + sendCalls++ + if r.URL.Query().Get("access_token") != "token" { + t.Errorf("send query = %s", r.URL.RawQuery) + } + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + t.Errorf("decode send payload = %v", err) + } + _, _ = io.WriteString(w, `{"errcode":0,"msgid":"native-`+test.messageType+`-1"}`) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + reference := wecomReplyReference(t, test.attachmentKind, test.mimeType, test.fileName, test.data) + reader := &providerAttachmentReader{content: attachment.Content{Data: test.data}} + provider := &Provider{CorpID: "corp", AgentID: "1", AppSecret: "secret", BaseURL: server.URL, HTTPClient: server.Client(), Attachments: reader} + reply := storage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-1", ReplyID: "reply-native-" + test.messageType, + SegmentIndex: 0, SegmentCount: 1, Kind: test.kind, Payload: "caption", + Attachment: reference, Fallback: test.fallback, + ReplyTarget: storage.ReplyTarget{ConversationKind: "direct", ReceiverID: "user-1"}, + } + + receipt, err := provider.Deliver(context.Background(), reply) + if err != nil || receipt != "native-"+test.messageType+"-1" { + t.Fatalf("native deliver = %q, %v", receipt, err) + } + if uploadCalls != 1 || sendCalls != 1 || reader.calls != 1 { + t.Fatalf("calls upload=%d send=%d reader=%d", uploadCalls, sendCalls, reader.calls) + } + if reader.tenantID != reply.TenantID || reader.eventID != reply.EventID || reader.reference != reference { + t.Fatalf("reader args = %q %q %+v", reader.tenantID, reader.eventID, reader.reference) + } + media, ok := payload[test.messageType].(map[string]any) + if payload["msgtype"] != test.messageType || payload["touser"] != "user-1" || payload["agentid"] != float64(1) || !ok || media["media_id"] != "uploaded-"+test.messageType { + t.Fatalf("native payload = %#v", payload) + } + }) + } +} + +func TestProviderPreservesNativeAttachmentCancellation(t *testing.T) { + data := []byte("png") + reference := wecomReplyReference(t, attachment.KindImage, "image/png", "chart.png", data) + provider := &Provider{ + CorpID: "corp", AgentID: "1", AppSecret: "secret", + Attachments: &providerAttachmentReader{err: context.Canceled}, + token: "cached", tokenExpiry: time.Now().Add(time.Hour), + } + _, err := provider.Deliver(context.Background(), storage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-image", SegmentIndex: 0, SegmentCount: 1, + Kind: storage.ReplyKindImage, Payload: "caption", Attachment: reference, Fallback: "[image attachment: chart.png]", + ReplyTarget: storage.ReplyTarget{ConversationKind: "direct", ReceiverID: "user-1"}, + }) + assertDeliveryErrorClass(t, err, "canceled", true) +} + +func TestHTTPMediaDownloaderFetchesMediaWithVerifiedBindingContext(t *testing.T) { + var tokenCalls int + var mediaCalls int + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/cgi-bin/gettoken": + tokenCalls++ + if r.URL.Query().Get("corpid") != "corp" || r.URL.Query().Get("corpsecret") != "secret" { + t.Errorf("token query = %s", r.URL.RawQuery) + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"errcode":0,"access_token":"token","expires_in":3600}`) + case "/cgi-bin/media/get": + mediaCalls++ + if r.URL.Query().Get("access_token") != "token" || r.URL.Query().Get("media_id") != "media-1" { + t.Errorf("media query = %s", r.URL.RawQuery) + } + w.Header().Set("Content-Type", "application/octet-stream") + _, _ = io.WriteString(w, "media-bytes") + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + downloader := &HTTPMediaDownloader{BaseURL: server.URL, HTTPClient: server.Client(), Now: func() time.Time { return time.Unix(100, 0).UTC() }} + request := MediaDownloadRequest{TenantID: "tenant-a", BindingID: "binding-1", CorpID: "corp", AgentID: "1", AppSecret: "secret", MediaID: "media-1", Kind: attachment.KindImage, MIMEType: "image/jpeg", MaximumBytes: 1 << 20} + for i := 0; i < 2; i++ { + body, err := downloader.Download(context.Background(), request) + if err != nil { + t.Fatalf("download = %v", err) + } + data, err := io.ReadAll(body) + closeErr := body.Close() + if err != nil || closeErr != nil || string(data) != "media-bytes" { + t.Fatalf("download body = %q read=%v close=%v", data, err, closeErr) + } + } + if tokenCalls != 1 || mediaCalls != 2 { + t.Fatalf("calls token=%d media=%d", tokenCalls, mediaCalls) + } +} + +func TestHTTPMediaDownloaderRejectsInvalidOrProviderErrorResponses(t *testing.T) { + if _, err := (*HTTPMediaDownloader)(nil).Download(context.Background(), MediaDownloadRequest{}); !errors.Is(err, ErrAttachment) { + t.Fatalf("nil downloader error = %v", err) + } + + for _, test := range []struct { + name string + kind attachment.Kind + mime string + body string + }{ + {name: "provider error JSON", kind: attachment.KindImage, mime: "image/jpeg", body: `{"errcode":40003}`}, + {name: "successful JSON is not document media", kind: attachment.KindDocument, mime: "application/octet-stream", body: `{"errcode":0}`}, + } { + t.Run(test.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/cgi-bin/gettoken" { + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"errcode":0,"access_token":"token","expires_in":3600}`) + return + } + w.Header().Set("Content-Type", "application/json; charset=utf-8") + _, _ = io.WriteString(w, test.body) + })) + defer server.Close() + + downloader := &HTTPMediaDownloader{BaseURL: server.URL, HTTPClient: server.Client()} + request := MediaDownloadRequest{TenantID: "tenant-a", BindingID: "binding-1", CorpID: "corp", AgentID: "1", AppSecret: "secret", MediaID: "media-1", Kind: test.kind, MIMEType: test.mime, MaximumBytes: 1 << 20} + body, err := downloader.Download(context.Background(), request) + if body != nil || !errors.Is(err, ErrAttachment) { + t.Fatalf("provider JSON download = body %v err %v", body, err) + } + }) + } +} + +func TestWeComMediaHelpersCoverFamiliesAndFallbacks(t *testing.T) { + for _, test := range []struct { + name string + message inboundXML + want wecomAttachment + }{ + {name: "file", message: inboundXML{MsgType: " file ", MediaID: " media-file ", FileName: "bad/name"}, want: wecomAttachment{mediaID: "media-file", kind: attachment.KindDocument, mimeType: wecomDefaultFileMIME, name: "media-file"}}, + {name: "voice", message: inboundXML{MsgType: "voice", MediaID: "media-voice", Format: "speex"}, want: wecomAttachment{mediaID: "media-voice", kind: attachment.KindAudio, mimeType: "audio/speex", name: "media-voice.speex"}}, + {name: "video", message: inboundXML{MsgType: "video", MediaID: "media-video"}, want: wecomAttachment{mediaID: "media-video", kind: attachment.KindVideo, mimeType: wecomDefaultVideoMIME, name: "media-video.mp4"}}, + } { + t.Run(test.name, func(t *testing.T) { + got, err := wecomAttachmentDescriptor(test.message) + if err != nil || got != test.want { + t.Fatalf("descriptor = %+v, %v; want %+v", got, err, test.want) + } + }) + } + if _, err := wecomAttachmentDescriptor(inboundXML{MsgType: "location", MediaID: "media"}); !errors.Is(err, ErrInvalid) { + t.Fatalf("unsupported descriptor = %v", err) + } + for _, test := range []struct { + format, mime, suffix string + }{ + {format: "mp3", mime: "audio/mpeg", suffix: ".mp3"}, + {format: "wav", mime: "audio/wav", suffix: ".wav"}, + {format: "m4a", mime: "audio/mp4", suffix: ".m4a"}, + {format: "ogg", mime: "audio/ogg", suffix: ".ogg"}, + {format: "unknown", mime: wecomDefaultVoiceMIME, suffix: wecomDefaultVoiceSuffix}, + } { + if mime, suffix := voiceMIME(test.format); mime != test.mime || suffix != test.suffix { + t.Fatalf("voiceMIME(%q) = %q/%q, want %q/%q", test.format, mime, suffix, test.mime, test.suffix) + } + } + if got := mediaName(" report.pdf ", "media-file", ".bin"); got != "report.pdf" { + t.Fatalf("mediaName kept name = %q", got) + } + long := strings.Repeat("x", wecomMediaContentRunes) + if got := wecomMediaContent(attachment.Reference{Kind: attachment.KindDocument, Name: long}); got != "[wecom document attachment]" { + t.Fatalf("long media content = %q", got) + } + if got, err := normalizeAttachmentBytes(1); err != nil || got != 1 { + t.Fatalf("normalizeAttachmentBytes = %d, %v", got, err) + } + if _, err := normalizeAttachmentBytes(maximumAttachmentBytes + 1); !errors.Is(err, ErrInvalid) { + t.Fatalf("oversized attachment limit = %v", err) + } +} + +func TestHTTPMediaDownloaderValidatesRequestAndResponseBoundaries(t *testing.T) { + normalized, err := normalizeMediaDownloadRequest(MediaDownloadRequest{ + TenantID: " tenant-a ", BindingID: " binding-1 ", CorpID: " corp ", AgentID: " 1 ", AppSecret: " secret ", + MediaID: " media-1 ", Kind: attachment.KindDocument, MIMEType: " APPLICATION/PDF ", + }) + if err != nil || normalized.TenantID != "tenant-a" || normalized.MIMEType != "application/pdf" || normalized.MaximumBytes != defaultAttachmentBytes { + t.Fatalf("normalized request = %+v, %v", normalized, err) + } + for _, test := range []struct { + name string + mutate func(*MediaDownloadRequest) + }{ + {name: "missing tenant", mutate: func(request *MediaDownloadRequest) { request.TenantID = "" }}, + {name: "invalid kind", mutate: func(request *MediaDownloadRequest) { request.Kind = "sticker" }}, + {name: "control value", mutate: func(request *MediaDownloadRequest) { request.MediaID = "bad\nmedia" }}, + {name: "oversized maximum", mutate: func(request *MediaDownloadRequest) { request.MaximumBytes = maximumAttachmentBytes + 1 }}, + } { + t.Run(test.name, func(t *testing.T) { + request := normalized + test.mutate(&request) + if _, err := normalizeMediaDownloadRequest(request); !errors.Is(err, ErrAttachment) { + t.Fatalf("normalizeMediaDownloadRequest accepted %+v: %v", request, err) + } + }) + } + + canceled, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := (&HTTPMediaDownloader{}).Download(canceled, normalized); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled download = %v", err) + } + if _, err := readWeComMediaResponse(context.Background(), nil, normalized); !errors.Is(err, ErrAttachment) { + t.Fatalf("nil response = %v", err) + } + if _, err := readWeComMediaResponse(context.Background(), &http.Response{}, normalized); !errors.Is(err, ErrAttachment) { + t.Fatalf("nil body = %v", err) + } + for _, test := range []struct { + name string + status int + contentType string + body string + kind attachment.Kind + maximum int64 + }{ + {name: "bad status", status: http.StatusBadGateway, contentType: "application/octet-stream", body: "x", kind: attachment.KindDocument, maximum: 8}, + {name: "empty body", status: http.StatusOK, contentType: "application/octet-stream", kind: attachment.KindDocument, maximum: 8}, + {name: "too large", status: http.StatusOK, contentType: "application/octet-stream", body: "123456", kind: attachment.KindDocument, maximum: 5}, + {name: "json suffix", status: http.StatusOK, contentType: "application/problem+json", body: "{}", kind: attachment.KindDocument, maximum: 8}, + {name: "media mismatch", status: http.StatusOK, contentType: "image/png", body: "png", kind: attachment.KindDocument, maximum: 8}, + } { + t.Run(test.name, func(t *testing.T) { + response := &http.Response{StatusCode: test.status, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(test.body))} + response.Header.Set("Content-Type", test.contentType) + request := normalized + request.Kind = test.kind + request.MaximumBytes = test.maximum + if data, err := readWeComMediaResponse(context.Background(), response, request); data != nil || !errors.Is(err, ErrAttachment) { + t.Fatalf("readWeComMediaResponse = %q, %v", data, err) + } + }) + } + if providerJSONError("not a media type", nil) { + t.Fatal("non-JSON invalid content type was rejected as provider JSON") + } + for _, test := range []struct { + kind attachment.Kind + contentType string + want bool + }{ + {kind: attachment.KindImage, contentType: "image/png", want: true}, + {kind: attachment.KindVideo, contentType: "video/mp4", want: true}, + {kind: attachment.KindAudio, contentType: "audio/mpeg", want: true}, + {kind: attachment.KindDocument, contentType: "application/pdf", want: true}, + {kind: attachment.KindDocument, contentType: "audio/mpeg"}, + {kind: attachment.Kind("unknown"), contentType: "application/octet-stream", want: true}, + {kind: attachment.Kind("unknown"), contentType: "text/plain"}, + } { + if got := downloadContentTypeMatches(test.kind, test.contentType); got != test.want { + t.Fatalf("downloadContentTypeMatches(%q, %q) = %t, want %t", test.kind, test.contentType, got, test.want) + } + } +} + +func TestHandlerBuildsGroupMediaAndIngestFailures(t *testing.T) { + target := staticTestTarget(t) + principal, err := gateway.NewChannelPrincipal(target) + if err != nil { + t.Fatal(err) + } + state := callbackState{principal: principal, agentID: "1", appSecret: "app-secret"} + data := []byte("document") + downloader := &fakeWeComMediaDownloader{data: data} + handler := &Handler{attachments: attachmentmemory.New(), mediaDownloader: downloader, maxAttachmentBytes: defaultAttachmentBytes} + + inbound, err := handler.buildInboundMessage(context.Background(), state, inboundXML{ + MsgID: "message-file", FromUserName: "user-1", ChatID: "chat-1", MsgType: "file", AgentID: "1", MediaID: "media-file", FileName: "brief.pdf", + }) + if err != nil { + t.Fatal(err) + } + if inbound.ConversationKind != channels.ConversationGroup || inbound.ExternalChatID != "chat-1" || inbound.Content != "[wecom document attachment: brief.pdf]" || len(inbound.Attachments) != 1 { + t.Fatalf("group media inbound = %+v", inbound) + } + if downloader.request.BindingID != target.BindingID || downloader.request.CorpID != target.ProviderAccountID || downloader.request.AppSecret != "app-secret" { + t.Fatalf("download request = %+v", downloader.request) + } + + if _, err := handler.buildInboundMessage(context.Background(), state, inboundXML{MsgID: "message-bad", FromUserName: "user-1", MsgType: "location", MediaID: "media"}); !errors.Is(err, ErrAttachment) { + t.Fatalf("invalid media build = %v", err) + } + canceled, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := handler.ingestAttachment(canceled, state, inboundXML{MsgID: "message-cancel", FromUserName: "user-1", MsgType: "image", MediaID: "media-image"}); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled ingest = %v", err) + } + for _, test := range []struct { + name string + downloader MediaDownloader + want error + }{ + {name: "download canceled", downloader: wecomMediaDownloaderFunc(func(context.Context, MediaDownloadRequest) (io.ReadCloser, error) { return nil, context.Canceled }), want: context.Canceled}, + {name: "download redacted", downloader: wecomMediaDownloaderFunc(func(context.Context, MediaDownloadRequest) (io.ReadCloser, error) { + return nil, errors.New("provider secret") + }), want: ErrAttachment}, + {name: "nil reader", downloader: wecomMediaDownloaderFunc(func(context.Context, MediaDownloadRequest) (io.ReadCloser, error) { return nil, nil }), want: ErrAttachment}, + {name: "read error", downloader: wecomMediaDownloaderFunc(func(context.Context, MediaDownloadRequest) (io.ReadCloser, error) { return failingReadCloser{}, nil }), want: ErrAttachment}, + {name: "empty body", downloader: &fakeWeComMediaDownloader{}, want: ErrAttachment}, + {name: "too large", downloader: &fakeWeComMediaDownloader{data: []byte("123456")}, want: ErrAttachment}, + } { + t.Run(test.name, func(t *testing.T) { + h := &Handler{attachments: attachmentmemory.New(), mediaDownloader: test.downloader, maxAttachmentBytes: 5} + _, err := h.ingestAttachment(context.Background(), state, inboundXML{MsgID: "message-" + test.name, FromUserName: "user-1", MsgType: "image", MediaID: "media-image"}) + if !errors.Is(err, test.want) || strings.Contains(err.Error(), "secret") { + t.Fatalf("ingest error = %v, want %v", err, test.want) + } + }) + } + storeCanceled, cancelStore := context.WithCancel(context.Background()) + downloader = &fakeWeComMediaDownloader{data: data} + cancelingStore := wecomAttachmentStoreFunc(func(context.Context, string, attachment.Upload, io.Reader) (attachment.Reference, error) { + cancelStore() + return attachment.Reference{}, context.Canceled + }) + h := &Handler{attachments: cancelingStore, mediaDownloader: downloader, maxAttachmentBytes: defaultAttachmentBytes} + if _, err := h.ingestAttachment(storeCanceled, state, inboundXML{MsgID: "message-store-cancel", FromUserName: "user-1", MsgType: "image", MediaID: "media-image"}); !errors.Is(err, context.Canceled) { + t.Fatalf("store cancellation = %v", err) + } + h = &Handler{attachments: wecomAttachmentStoreFunc(func(context.Context, string, attachment.Upload, io.Reader) (attachment.Reference, error) { + return attachment.Reference{}, errors.New("write failed") + }), mediaDownloader: downloader, maxAttachmentBytes: defaultAttachmentBytes} + if _, err := h.ingestAttachment(context.Background(), state, inboundXML{MsgID: "message-store", FromUserName: "user-1", MsgType: "image", MediaID: "media-image"}); !errors.Is(err, ErrAttachment) { + t.Fatalf("store failure = %v", err) + } +} + +func TestHandlerAttachmentIDIncludesBindingScope(t *testing.T) { + firstTarget := staticTestTarget(t) + secondTarget := staticTestTarget(t) + if firstTarget.BindingID == secondTarget.BindingID { + t.Fatal("test requires distinct binding IDs") + } + firstPrincipal, err := gateway.NewChannelPrincipal(firstTarget) + if err != nil { + t.Fatal(err) + } + secondPrincipal, err := gateway.NewChannelPrincipal(secondTarget) + if err != nil { + t.Fatal(err) + } + data := []byte("wecom-image") + handler := &Handler{attachments: attachmentmemory.New(), mediaDownloader: &fakeWeComMediaDownloader{data: data}, maxAttachmentBytes: defaultAttachmentBytes} + message := inboundXML{MsgID: "message-same", FromUserName: "user-1", MsgType: "image", AgentID: "1", MediaID: "media-same"} + + first, err := handler.ingestAttachment(context.Background(), callbackState{principal: firstPrincipal, agentID: "1", appSecret: "app-secret"}, message) + if err != nil { + t.Fatal(err) + } + second, err := handler.ingestAttachment(context.Background(), callbackState{principal: secondPrincipal, agentID: "1", appSecret: "app-secret"}, message) + if err != nil { + t.Fatal(err) + } + if first.ID == second.ID { + t.Fatalf("attachment IDs collided across bindings: %q", first.ID) + } + if first.ID != attachmentID(firstTarget.BindingID, "message-same", 0, "media-same") || second.ID != attachmentID(secondTarget.BindingID, "message-same", 0, "media-same") { + t.Fatalf("binding-scoped IDs = %q and %q", first.ID, second.ID) + } +} + +func TestProviderNativeMediaErrorBranches(t *testing.T) { + data := []byte("png") + reference := wecomReplyReference(t, attachment.KindImage, "image/png", "chart.png", data) + base := storage.ReplyOutbox{ + TenantID: "tenant-a", EventID: "event-image", ReplyID: "reply-image", SegmentIndex: 0, SegmentCount: 1, + Kind: storage.ReplyKindImage, Payload: "caption", Attachment: reference, Fallback: "[image attachment: chart.png]", + ReplyTarget: storage.ReplyTarget{ConversationKind: "direct", ReceiverID: "user-1"}, + } + for _, test := range []struct { + name string + reader *providerAttachmentReader + class string + retryable bool + }{ + {name: "missing attachment", reader: &providerAttachmentReader{err: storage.ErrNotFound}, class: "invalid"}, + {name: "storage unavailable", reader: &providerAttachmentReader{err: errors.New("storage down")}, class: "unavailable", retryable: true}, + {name: "tampered attachment", reader: &providerAttachmentReader{content: attachment.Content{Data: []byte("tampered")}}, class: "invalid"}, + } { + t.Run(test.name, func(t *testing.T) { + provider := &Provider{CorpID: "corp", AgentID: "1", AppSecret: "secret", Attachments: test.reader, token: "cached", tokenExpiry: time.Now().Add(time.Hour)} + _, err := provider.Deliver(context.Background(), base) + assertDeliveryErrorClass(t, err, test.class, test.retryable) + }) + } + for _, test := range []struct { + name string + status int + body string + class string + retryable bool + }{ + {name: "upload unavailable", status: http.StatusBadGateway, class: "unavailable", retryable: true}, + {name: "upload malformed", status: http.StatusOK, body: "not-json", class: "provider_error", retryable: true}, + {name: "upload rejected", status: http.StatusOK, body: `{"errcode":40003}`, class: "provider_error"}, + {name: "missing media id", status: http.StatusOK, body: `{"errcode":0}`, class: "provider_error", retryable: true}, + } { + t.Run(test.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/cgi-bin/media/upload" { + t.Fatalf("path = %s", r.URL.Path) + } + w.WriteHeader(test.status) + _, _ = io.WriteString(w, test.body) + })) + defer server.Close() + provider := &Provider{CorpID: "corp", AgentID: "1", AppSecret: "secret", BaseURL: server.URL, HTTPClient: server.Client(), Attachments: &providerAttachmentReader{content: attachment.Content{Data: data}}, token: "cached", tokenExpiry: time.Now().Add(time.Hour)} + _, err := provider.Deliver(context.Background(), base) + assertDeliveryErrorClass(t, err, test.class, test.retryable) + }) + } + transport := &Provider{ + CorpID: "corp", AgentID: "1", AppSecret: "secret", + HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { return nil, context.DeadlineExceeded })}, + Attachments: &providerAttachmentReader{content: attachment.Content{Data: data}}, + token: "cached", tokenExpiry: time.Now().Add(time.Hour), + } + _, err := transport.Deliver(context.Background(), base) + assertDeliveryErrorClass(t, err, "timeout", true) + + fallback := base + fallback.Kind = storage.ReplyKindAudio + fallback.Attachment = wecomReplyReference(t, attachment.KindAudio, "audio/mpeg", "voice.mp3", []byte("mp3")) + var payload map[string]any + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/cgi-bin/message/send" { + t.Fatalf("path = %s", r.URL.Path) + } + _ = json.NewDecoder(r.Body).Decode(&payload) + _, _ = io.WriteString(w, `{"errcode":0,"msgid":"audio-fallback"}`) + })) + defer server.Close() + receipt, err := (&Provider{CorpID: "corp", AgentID: "1", AppSecret: "secret", BaseURL: server.URL, HTTPClient: server.Client(), token: "cached", tokenExpiry: time.Now().Add(time.Hour)}).Deliver(context.Background(), fallback) + if err != nil || receipt != "audio-fallback" { + t.Fatalf("audio fallback = %q, %v", receipt, err) + } + text, ok := payload["text"].(map[string]any) + if payload["msgtype"] != "text" || !ok || text["content"] != fallback.Fallback { + t.Fatalf("audio fallback payload = %#v", payload) + } +} + func TestProviderValidatesBeforeReceiptReplay(t *testing.T) { p := &Provider{CorpID: "corp", AgentID: "1", AppSecret: "secret", receipts: map[string]string{"tenant\x00reply\x000": "m-1"}} value := storage.ReplyOutbox{TenantID: "tenant", ReplyID: "reply", SegmentIndex: 0, Payload: "", ReplyTarget: storage.ReplyTarget{ConversationKind: "direct", ReceiverID: "user"}} @@ -429,6 +977,75 @@ func TestHandlerAcceptsEncryptedTextWithRequestAndTraceIDs(t *testing.T) { } } +//nolint:gocyclo // Keeps the encrypted callback-to-attachment contract visible in one scenario. +func TestHandlerAcceptsEncryptedNativeMediaAsAttachment(t *testing.T) { + dispatcher := &callbackDispatchStub{requests: make(chan gateway.DispatchRequest, 1)} + app := dynamicTestApp(t, "t_01ARZ3NDEKTSV4RRFFQ69G5FAV") + binding := dynamicTestBinding(t, "media-route", "env/wecom", app.AppID) + data := []byte("wecom-image") + downloader := &fakeWeComMediaDownloader{data: data} + handler, err := New(Config{ + Candidates: &dynamicCandidateConsumer{binding: binding}, + Tenants: dynamicTenantRepository{value: dynamicTestTenant(t)}, + Apps: dynamicAppRepository{value: app}, + Credentials: dynamicCredentials{values: map[string]Credentials{binding.SecretRef: {CallbackToken: "token", EncodingAESKey: base64.RawStdEncoding.EncodeToString(bytes.Repeat([]byte{1}, 32)), AppSecret: "app-secret"}}}, + Dispatcher: dispatcher, + Attachments: attachmentmemory.New(), + MediaDownloader: downloader, + MaxBodyBytes: 1 << 20, + ExecutionTimeout: time.Minute, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = handler.Close() }) + + request := callbackXMLRequestAtPath(t, "/wecom/callback/media-route", []byte("message-mediauser-1image1media-imagehttps://provider.example/token")) + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != http.StatusOK || response.Body.String() != "success" { + t.Fatalf("media callback response = %d %q", response.Code, response.Body.String()) + } + select { + case request := <-dispatcher.requests: + if request.Principal.TenantID() != binding.TenantID { + t.Fatalf("dispatch principal tenant = %q, want %q", request.Principal.TenantID(), binding.TenantID) + } + if request.Message.ContentType != gateway.ContentTypeMedia || request.Message.Content != "[wecom image attachment: media-image.jpg]" || len(request.Message.Attachments) != 1 { + t.Fatalf("media dispatch message = %+v", request.Message) + } + reference := request.Message.Attachments[0] + digest := sha256.Sum256(data) + if reference.ID != attachmentID(binding.BindingID, "message-media", 0, "media-image") || reference.Kind != attachment.KindImage || reference.MIMEType != "image/jpeg" || reference.Name != "media-image.jpg" || reference.Provider != "wecom" || reference.ProviderID != "media-image" || reference.Size != int64(len(data)) || reference.SHA256 != hex.EncodeToString(digest[:]) { + t.Fatalf("media attachment reference = %+v", reference) + } + if downloader.request.MediaID != "media-image" || downloader.request.TenantID != binding.TenantID || downloader.request.BindingID != binding.BindingID || + downloader.request.CorpID != "corp" || downloader.request.AgentID != "1" || downloader.request.AppSecret != "app-secret" || + downloader.request.Kind != attachment.KindImage || downloader.request.MIMEType != "image/jpeg" || downloader.request.MaximumBytes != defaultAttachmentBytes { + t.Fatalf("download request = %+v", downloader.request) + } + case <-time.After(time.Second): + t.Fatal("media callback did not dispatch") + } +} + +func TestHandlerRejectsMediaWithoutConfiguredAttachmentBoundary(t *testing.T) { + dispatcher := &callbackDispatchStub{requests: make(chan gateway.DispatchRequest, 1)} + handler := newCallbackTestHandler(t, dispatcher) + t.Cleanup(func() { _ = handler.Close() }) + + response := httptest.NewRecorder() + handler.ServeHTTP(response, callbackXMLRequestAtPath(t, "/", []byte("message-mediauser-1image1media-image"))) + if response.Code != http.StatusBadRequest { + t.Fatalf("media without attachment boundary response = %d", response.Code) + } + select { + case request := <-dispatcher.requests: + t.Fatalf("unsupported media reached dispatch: %+v", request) + default: + } +} + func TestHandlerCloseCancelsAndJoinsAcceptedDrain(t *testing.T) { dispatcher := &callbackDispatchStub{requests: make(chan gateway.DispatchRequest, 1), canceled: make(chan struct{})} handler := newCallbackTestHandler(t, dispatcher) @@ -717,6 +1334,31 @@ func TestHandlerRejectsMalformedMessages(t *testing.T) { } } +func TestHandlerRejectsUnknownOrIncompleteMediaCallbacks(t *testing.T) { + dispatcher := &callbackDispatchStub{requests: make(chan gateway.DispatchRequest, 1)} + handler := newCallbackTestHandler(t, dispatcher) + handler.attachments = attachmentmemory.New() + handler.mediaDownloader = &fakeWeComMediaDownloader{data: []byte("media")} + handler.maxAttachmentBytes = defaultAttachmentBytes + t.Cleanup(func() { _ = handler.Close() }) + + for _, message := range [][]byte{ + []byte("unknown-mediauser-1location1media-id"), + []byte("missing-media-iduser-1file1report.pdf"), + } { + response := httptest.NewRecorder() + handler.ServeHTTP(response, callbackXMLRequestAtPath(t, "/", message)) + if response.Code != http.StatusBadRequest { + t.Fatalf("malformed media response = %d", response.Code) + } + } + select { + case request := <-dispatcher.requests: + t.Fatalf("malformed media reached dispatch: %+v", request) + default: + } +} + func TestDynamicHandlerAnswersVerifiedChallenge(t *testing.T) { app := dynamicTestApp(t, "t_01ARZ3NDEKTSV4RRFFQ69G5FAV") binding := dynamicTestBinding(t, "challenge-key", "env/wecom", app.AppID) @@ -759,7 +1401,8 @@ func TestBindingProviderUsesActiveWeComBindingAndCachesProvider(t *testing.T) { } lookup := &bindingLookupStub{binding: binding} credentials := &credentialResolverStub{credentials: Credentials{AppSecret: "app-secret"}} - provider := &BindingProvider{Bindings: lookup, Credentials: credentials} + reader := &providerAttachmentReader{} + provider := &BindingProvider{Bindings: lookup, Credentials: credentials, Attachments: reader} value := storage.ReplyOutbox{TenantID: binding.TenantID, ReplyTarget: storage.ReplyTarget{BindingID: binding.BindingID, ConversationKind: "direct", ReceiverID: "user-1"}} first, err := provider.provider(context.Background(), value) if err != nil { @@ -772,6 +1415,9 @@ func TestBindingProviderUsesActiveWeComBindingAndCachesProvider(t *testing.T) { if first != second || first.CorpID != "corp" || first.AgentID != "1" { t.Fatalf("binding provider = %+v, cached=%t", first, first == second) } + if first.Attachments != reader { + t.Fatalf("binding provider attachment reader = %#v", first.Attachments) + } if lookup.calls != 2 || credentials.calls != 2 { t.Fatalf("lookup=%d credentials=%d", lookup.calls, credentials.calls) } @@ -974,6 +1620,79 @@ func (roundTrip roundTripFunc) RoundTrip(request *http.Request) (*http.Response, return roundTrip(request) } +type fakeWeComMediaDownloader struct { + data []byte + err error + request MediaDownloadRequest + mediaID string +} + +func (downloader *fakeWeComMediaDownloader) Download(ctx context.Context, request MediaDownloadRequest) (io.ReadCloser, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + downloader.request = request + downloader.mediaID = request.MediaID + if downloader.err != nil { + return nil, downloader.err + } + return io.NopCloser(bytes.NewReader(downloader.data)), nil +} + +type wecomMediaDownloaderFunc func(context.Context, MediaDownloadRequest) (io.ReadCloser, error) + +func (function wecomMediaDownloaderFunc) Download(ctx context.Context, request MediaDownloadRequest) (io.ReadCloser, error) { + return function(ctx, request) +} + +type failingReadCloser struct{} + +func (failingReadCloser) Read([]byte) (int, error) { + return 0, errors.New("read failed") +} + +func (failingReadCloser) Close() error { return nil } + +type wecomAttachmentStoreFunc func(context.Context, string, attachment.Upload, io.Reader) (attachment.Reference, error) + +func (function wecomAttachmentStoreFunc) PutAttachment(ctx context.Context, tenantID string, upload attachment.Upload, content io.Reader) (attachment.Reference, error) { + return function(ctx, tenantID, upload, content) +} + +func (wecomAttachmentStoreFunc) Load(context.Context, string, string, attachment.Reference) (attachment.Content, error) { + return attachment.Content{}, storage.ErrNotFound +} + +func (wecomAttachmentStoreFunc) BindAttachments(context.Context, string, string, []attachment.Reference) error { + return storage.ErrNotFound +} + +func (wecomAttachmentStoreFunc) CleanupAttachments(context.Context, string, time.Time) (int, error) { + return 0, nil +} + +func (wecomAttachmentStoreFunc) Close() error { return nil } + +type providerAttachmentReader struct { + content attachment.Content + err error + tenantID string + eventID string + reference attachment.Reference + calls int +} + +func (reader *providerAttachmentReader) Load(_ context.Context, tenantID, eventID string, reference attachment.Reference) (attachment.Content, error) { + reader.calls++ + reader.tenantID = tenantID + reader.eventID = eventID + reader.reference = reference + if reader.err != nil { + return attachment.Content{}, reader.err + } + return reader.content.Clone(), nil +} + func assertDeliveryErrorClass(t *testing.T, err error, class string, retryable bool) { t.Helper() var deliveryErr *outbox.DeliveryError @@ -982,6 +1701,16 @@ func assertDeliveryErrorClass(t *testing.T, err error, class string, retryable b } } +func wecomReplyReference(t *testing.T, kind attachment.Kind, contentType, name string, data []byte) attachment.Reference { + t.Helper() + digest := sha256.Sum256(data) + reference := attachment.Reference{ID: "attachment-" + string(kind), Kind: kind, MIMEType: contentType, Name: name, Size: int64(len(data)), SHA256: hex.EncodeToString(digest[:])} + if _, err := reference.Normalize(); err != nil { + t.Fatalf("attachment reference = %v", err) + } + return reference +} + func requestBodyWithTrailingXML(t *testing.T, request *http.Request) string { t.Helper() body, err := io.ReadAll(request.Body) @@ -1039,6 +1768,11 @@ func callbackVerificationRequest(token, path, ciphertext string) *http.Request { func callbackTestRequestAtPath(t *testing.T, path, messageID, userID, content string) *http.Request { t.Helper() plain := []byte("" + messageID + "" + userID + "text1" + content + "") + return callbackXMLRequestAtPath(t, path, plain) +} + +func callbackXMLRequestAtPath(t *testing.T, path string, plain []byte) *http.Request { + t.Helper() ciphertext := encryptCallbackTestPayload(t, bytes.Repeat([]byte{1}, 32), "receive", plain) request := callbackVerificationRequest("token", path, ciphertext) request.Method = http.MethodPost @@ -1171,6 +1905,24 @@ func newDynamicVerifyFixture(t *testing.T) (*Handler, *verifyCandidateConsumer, return handler, consumer, ciphertext, callbackVerificationRequest("token", "/wecom/callback/verify-route", ciphertext) } +func staticTestTarget(t *testing.T) channels.RoutingTarget { + t.Helper() + app := dynamicTestApp(t, "t_01ARZ3NDEKTSV4RRFFQ69G5FAV") + binding := dynamicTestBinding(t, "static-route", "env/wecom", app.AppID) + target, err := channels.ResolveCandidateRoutingTarget( + context.Background(), + &dynamicCandidateConsumer{binding: binding}, + dynamicTenantRepository{value: dynamicTestTenant(t)}, + dynamicAppRepository{value: app}, + channels.CandidateBindingContext{Channel: channels.ChannelWeCom}, + func(context.Context, channels.Binding) error { return nil }, + ) + if err != nil { + t.Fatal(err) + } + return target +} + type dynamicCredentials struct{ values map[string]Credentials } func (resolver dynamicCredentials) Resolve(_ context.Context, scope channels.SecretScope) (Credentials, error) { diff --git a/trpcservice/gateway/dispatch.go b/trpcservice/gateway/dispatch.go index 55650b7f..863cdabb 100644 --- a/trpcservice/gateway/dispatch.go +++ b/trpcservice/gateway/dispatch.go @@ -8,6 +8,7 @@ import ( "strings" "time" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/audit" "github.com/XnLemon/trpc-agent-service/trpcservice/channels" "github.com/XnLemon/trpc-agent-service/trpcservice/metrics" @@ -96,6 +97,9 @@ type DispatchConfig struct { AuditWriter audit.Writer // HandoffStore durably reserves and finalizes execution audit facts. HandoffStore audit.HandoffStore + // Attachments loads verified tenant-owned media only when an inbound message + // contains attachment references. Text-only dispatches remain independent of it. + Attachments attachment.Reader } // Dispatcher resolves a fixed plan, acquires a Runner lease, and translates @@ -110,6 +114,7 @@ type Dispatcher struct { materializer *outbox.Materializer auditWriter audit.Writer handoffStore audit.HandoffStore + attachments attachment.Reader } type durableExecution struct { @@ -135,6 +140,11 @@ func NewDispatcher(config DispatchConfig) (*Dispatcher, error) { if config.Observability == nil { config.Observability = observability.NewNoopProvider() } + if config.Attachments == nil { + if reader, ok := config.RuntimeStore.(attachment.Reader); ok { + config.Attachments = reader + } + } config.AuditWriter = metrics.WrapAuditWriter(config.AuditWriter, config.Observability) if config.Materializer == nil && config.RuntimeStore != nil { materializer, err := outbox.NewMaterializer(outbox.MaterializerConfig{Store: config.RuntimeStore, Observability: config.Observability}) @@ -143,7 +153,7 @@ func NewDispatcher(config DispatchConfig) (*Dispatcher, error) { } config.Materializer = materializer } - return &Dispatcher{resolver: config.Resolver, registry: config.Registry, drainTimeout: config.DrainTimeout, telemetry: config.Observability, metrics: metrics.New(config.Observability), runtimeStore: config.RuntimeStore, materializer: config.Materializer, auditWriter: config.AuditWriter, handoffStore: config.HandoffStore}, nil + return &Dispatcher{resolver: config.Resolver, registry: config.Registry, drainTimeout: config.DrainTimeout, telemetry: config.Observability, metrics: metrics.New(config.Observability), runtimeStore: config.RuntimeStore, materializer: config.Materializer, auditWriter: config.AuditWriter, handoffStore: config.HandoffStore, attachments: config.Attachments}, nil } // Ready reports whether both plan resolution and Runner acquisition are ready. @@ -193,6 +203,32 @@ func (dispatcher *Dispatcher) Dispatch(ctx context.Context, request DispatchRequ finishWithError(err) return nil, err } + attachmentEventID := "" + if durable != nil { + attachmentEventID = durable.eventID + } + if len(message.Attachments) > 0 { + binder, ok := dispatcher.attachments.(attachment.Binder) + if !ok || attachmentEventID == "" { + dispatcher.failDurable(durable, ErrExecution) + finishWithError(ErrExecution) + return nil, ErrExecution + } + if err := binder.BindAttachments(ctx, request.Principal.TenantID(), attachmentEventID, message.Attachments); err != nil { + dispatcher.failDurable(durable, err) + finishWithError(err) + return nil, ErrExecution + } + } + userMessage, err := buildUserMessage(ctx, dispatcher.attachments, request.Principal.TenantID(), attachmentEventID, message) + if err != nil { + dispatcher.failDurable(durable, err) + finishWithError(err) + if IsContextCancellation(err) { + return nil, err + } + return nil, ErrExecution + } if planApp.CanaryRevision != nil && planSnapshot.Revision().Revision == *planApp.CanaryRevision { selectedRevision := planSnapshot.Revision().Revision if err := dispatcher.writeExecutionAuditRevision(ctx, request.Principal, message, identity, requestID, traceID, audit.EventCanarySelected, "", &selectedRevision); err != nil { @@ -244,7 +280,7 @@ func (dispatcher *Dispatcher) Dispatch(ctx context.Context, request DispatchRequ runnerStarted := time.Now() runnerCtx, _, finishRunner := observability.StartOperation(ctx, dispatcher.telemetry, observability.OperationRunnerExecution, "runner") _ = dispatcher.metrics.Request(runnerCtx, map[string]string{"component": "runner", "operation": observability.OperationRunnerExecution, "status": "started"}) - runnerEvents, err := runnerValue.Run(runnerCtx, identity.UserID, identity.SessionID, trpcmodel.NewUserMessage(message.Content), trpcagent.WithRequestID(requestID)) + runnerEvents, err := runnerValue.Run(runnerCtx, identity.UserID, identity.SessionID, userMessage, trpcagent.WithRequestID(requestID)) if err != nil { finishRunner(err) _ = dispatcher.metrics.Operation(runnerCtx, runnerStarted, map[string]string{"component": "runner", "operation": observability.OperationRunnerExecution}, err) diff --git a/trpcservice/gateway/dispatch_test.go b/trpcservice/gateway/dispatch_test.go index 0a1c8a67..34e3e9c6 100644 --- a/trpcservice/gateway/dispatch_test.go +++ b/trpcservice/gateway/dispatch_test.go @@ -10,6 +10,7 @@ import ( "time" "github.com/XnLemon/trpc-agent-service/trpcservice/agent" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/audit" "github.com/XnLemon/trpc-agent-service/trpcservice/channels" "github.com/XnLemon/trpc-agent-service/trpcservice/runtime" @@ -84,6 +85,25 @@ type claimStoreStub struct { getErr, createErr, recordErr, transitionErr error } +type dispatchAttachmentStore struct { + bindFn func(context.Context, string, string, []attachment.Reference) error + loadFn func(context.Context, string, string, attachment.Reference) (attachment.Content, error) +} + +func (s dispatchAttachmentStore) BindAttachments(ctx context.Context, tenantID, eventID string, references []attachment.Reference) error { + if s.bindFn != nil { + return s.bindFn(ctx, tenantID, eventID, references) + } + return nil +} + +func (s dispatchAttachmentStore) Load(ctx context.Context, tenantID, eventID string, reference attachment.Reference) (attachment.Content, error) { + if s.loadFn != nil { + return s.loadFn(ctx, tenantID, eventID, reference) + } + return attachment.Content{}, errors.New("attachment unavailable") +} + func (s *claimStoreStub) GetSession(context.Context, string, string) (runtimestorage.Session, error) { if s.getErr != nil { return runtimestorage.Session{}, s.getErr @@ -807,6 +827,59 @@ func TestDispatcherDurableChannelClaimSuppressesDuplicateRunner(t *testing.T) { } } +func TestDispatcherBindsStoredAttachmentBeforePassingVerifiedContentToRunner(t *testing.T) { + fixture := newGatewayFixture(t) + target := newTrustedRoutingTarget(t, fixture) + principal, err := NewChannelPrincipal(target) + if err != nil { + t.Fatal(err) + } + resolver, err := NewPlanResolver(PlanResolverConfig{ + Tenants: fixture.tenants, Apps: fixture.apps, Models: fixture.models, Backends: fixture.backends, + ModelCatalog: fixture.modelCatalog, BackendCatalog: fixture.backendCatalog, + }) + if err != nil { + t.Fatal(err) + } + data := []byte("image") + store := inmemory.New() + t.Cleanup(func() { _ = store.Close() }) + reference, err := store.PutAttachment(context.Background(), principal.TenantID(), attachment.Upload{ID: "attachment-1", Kind: attachment.KindImage, MIMEType: "image/png", Size: int64(len(data)), Provider: "telegram", ProviderID: "file-1"}, strings.NewReader(string(data))) + if err != nil { + t.Fatalf("PutAttachment = %v", err) + } + if _, err := store.Load(context.Background(), principal.TenantID(), "unbound-event", reference); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("unbound Load = %v", err) + } + var captured trpcmodel.Message + runnerValue := &testRunner{runFn: func(_ context.Context, _ string, _ string, message trpcmodel.Message, _ ...trpcagent.RunOption) (<-chan *trpcevent.Event, error) { + captured = message + events := make(chan *trpcevent.Event, 1) + events <- &trpcevent.Event{Response: &trpcmodel.Response{Done: true}} + close(events) + return events, nil + }} + registry, err := NewRunnerRegistry(RunnerRegistryConfig{Factory: func(context.Context, runtime.ExecutionPlan) (Runner, error) { return runnerValue, nil }}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = registry.Close() }) + dispatcher, err := NewDispatcher(DispatchConfig{Resolver: resolver, Registry: registry, RuntimeStore: store}) + if err != nil { + t.Fatal(err) + } + stream, err := dispatcher.Dispatch(context.Background(), DispatchRequest{Principal: principal, Message: InboundMessage{Content: "describe", ExternalMessageID: "attachment-message", ExternalUserID: "user-1", ConversationKind: channels.ConversationDirect, ExternalPeerID: "peer-1", Attachments: []attachment.Reference{reference}}}) + if err != nil { + t.Fatalf("Dispatch = %v", err) + } + if events := collectDispatchEvents(stream); len(events) != 1 || !events[0].Done { + t.Fatalf("dispatch events = %+v", events) + } + if captured.Content != "describe" || len(captured.ContentParts) != 1 || captured.ContentParts[0].Type != trpcmodel.ContentTypeImage || captured.ContentParts[0].Image == nil || string(captured.ContentParts[0].Image.Data) != string(data) { + t.Fatalf("Runner message = %+v", captured) + } +} + func TestDispatcherMaterializesDurableChannelReplyAndWorkerCompletesLifecycle(t *testing.T) { fixture := newGatewayFixture(t) target := newTrustedRoutingTarget(t, fixture) @@ -1114,6 +1187,114 @@ func TestDispatcherDurableDispatchFailurePaths(t *testing.T) { _ = registry.Close() } +func TestDispatcherDurableAttachmentFailurePaths(t *testing.T) { + fixture := newGatewayFixture(t) + target := newTrustedRoutingTarget(t, fixture) + principal, err := NewChannelPrincipal(target) + if err != nil { + t.Fatal(err) + } + resolver, err := NewPlanResolver(PlanResolverConfig{Tenants: fixture.tenants, Apps: fixture.apps, Models: fixture.models, Backends: fixture.backends, ModelCatalog: fixture.modelCatalog, BackendCatalog: fixture.backendCatalog}) + if err != nil { + t.Fatal(err) + } + reference := testAttachmentReference(t, attachment.KindImage, "image/png", []byte("image")) + newDispatcher := func(t *testing.T, store runtimestorage.RuntimeStore, attachments attachment.Reader) (*Dispatcher, *RunnerRegistry, *atomic.Int32) { + t.Helper() + var runnerCalls atomic.Int32 + registry, err := NewRunnerRegistry(RunnerRegistryConfig{Factory: func(context.Context, runtime.ExecutionPlan) (Runner, error) { + return &testRunner{runFn: func(context.Context, string, string, trpcmodel.Message, ...trpcagent.RunOption) (<-chan *trpcevent.Event, error) { + runnerCalls.Add(1) + return nil, nil + }}, nil + }}) + if err != nil { + t.Fatal(err) + } + dispatcher, err := NewDispatcher(DispatchConfig{ + Resolver: resolver, Registry: registry, RuntimeStore: store, Attachments: attachments, DrainTimeout: time.Millisecond, + }) + if err != nil { + t.Fatal(err) + } + return dispatcher, registry, &runnerCalls + } + for _, tt := range []struct { + name string + attachments attachment.Reader + want error + }{ + { + name: "missing binder", + attachments: attachmentReaderFunc(func(context.Context, string, string, attachment.Reference) (attachment.Content, error) { + return attachment.Content{}, errors.New("unexpected load") + }), + want: ErrExecution, + }, + { + name: "binder failure", + attachments: dispatchAttachmentStore{bindFn: func(context.Context, string, string, []attachment.Reference) error { + return errors.New("bind failed") + }}, + want: ErrExecution, + }, + { + name: "load failure", + attachments: dispatchAttachmentStore{loadFn: func(context.Context, string, string, attachment.Reference) (attachment.Content, error) { + return attachment.Content{}, errors.New("load failed") + }}, + want: ErrExecution, + }, + { + name: "load cancellation", + attachments: dispatchAttachmentStore{loadFn: func(context.Context, string, string, attachment.Reference) (attachment.Content, error) { + return attachment.Content{}, context.Canceled + }}, + want: context.Canceled, + }, + } { + t.Run(tt.name, func(t *testing.T) { + store := inmemory.New() + dispatcher, registry, runnerCalls := newDispatcher(t, store, tt.attachments) + defer func() { _ = registry.Close() }() + message := InboundMessage{ + Content: "caption", ExternalMessageID: "attachment-" + strings.ReplaceAll(tt.name, " ", "-"), ExternalUserID: "user-1", + ConversationKind: channels.ConversationDirect, ExternalPeerID: "peer-1", Attachments: []attachment.Reference{reference}, + } + stream, err := dispatcher.Dispatch(context.Background(), DispatchRequest{Principal: principal, Message: message}) + if !errors.Is(err, tt.want) || stream != nil { + t.Fatalf("Dispatch() stream=%v err=%v, want %v", stream, err, tt.want) + } + if runnerCalls.Load() != 0 { + t.Fatal("Runner started after attachment preparation failed") + } + assertDurableMessageStatus(t, store, principal, target, message, runtimestorage.EventFailed) + }) + } +} + +func assertDurableMessageStatus(t *testing.T, store runtimestorage.RuntimeStore, principal Principal, target channels.RoutingTarget, message InboundMessage, status string) { + t.Helper() + identity, err := dispatchRunnerIdentity(principal, message) + if err != nil { + t.Fatal(err) + } + reply, err := replyTarget(target, message) + if err != nil { + t.Fatal(err) + } + event, duplicate, err := store.RecordMessage(context.Background(), runtimestorage.MessageEventInput{ + TenantID: principal.TenantID(), EventID: "probe-" + message.ExternalMessageID, SessionID: identity.SessionID, + BindingID: target.BindingID, ExternalMessageID: message.ExternalMessageID, ReplyTarget: reply, + }) + if err != nil { + t.Fatal(err) + } + if !duplicate || event.Status != status { + t.Fatalf("durable message duplicate=%v status=%q, want %q", duplicate, event.Status, status) + } +} + func TestDispatcherRedactsRunnerErrors(t *testing.T) { runnerValue := &testRunner{} runnerValue.runFn = func(context.Context, string, string, trpcmodel.Message, ...trpcagent.RunOption) (<-chan *trpcevent.Event, error) { diff --git a/trpcservice/gateway/gateway.go b/trpcservice/gateway/gateway.go index 06cc5f36..2f4c28a0 100644 --- a/trpcservice/gateway/gateway.go +++ b/trpcservice/gateway/gateway.go @@ -7,6 +7,7 @@ import ( "fmt" "strings" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/channels" ) @@ -48,9 +49,10 @@ const ( // ConversationGroup identifies a group conversation. ConversationGroup = channels.ConversationGroup - maxPrincipalIDRunes = 256 - maxMessageRunes = 64 * 1024 - maxExternalIDRunes = 1024 + maxPrincipalIDRunes = 256 + maxMessageRunes = 64 * 1024 + maxExternalIDRunes = 1024 + maxInboundAttachments = 10 ) // PrincipalKind distinguishes the two independent authentication paths. @@ -152,22 +154,42 @@ type InboundMessage struct { ExternalPeerID string ExternalChatID string ExternalThreadID string + // Attachments contains verified, tenant-owned media references. It never + // contains provider URLs, credentials, or direct fetch instructions. + Attachments []attachment.Reference } // Normalize validates the message without consulting untrusted route hints. func (m InboundMessage) Normalize() (InboundMessage, error) { clone := m clone.Content = strings.TrimSpace(clone.Content) + clone.Attachments = append([]attachment.Reference(nil), clone.Attachments...) if clone.ContentType == "" { clone.ContentType = ContentTypeText + if len(clone.Attachments) > 0 { + clone.ContentType = ContentTypeMedia + } } switch clone.ContentType { case ContentTypeText, ContentTypeMedia, ContentTypeRich: default: return InboundMessage{}, fmt.Errorf("%w: unsupported content type", ErrInvalid) } - if n := len([]rune(clone.Content)); n < 1 || n > maxMessageRunes { - return InboundMessage{}, fmt.Errorf("%w: content must contain 1-%d characters", ErrInvalid, maxMessageRunes) + if n := len([]rune(clone.Content)); n > maxMessageRunes { + return InboundMessage{}, fmt.Errorf("%w: content must contain at most %d characters", ErrInvalid, maxMessageRunes) + } + if clone.Content == "" && len(clone.Attachments) == 0 { + return InboundMessage{}, fmt.Errorf("%w: content or attachment is required", ErrInvalid) + } + if len(clone.Attachments) > maxInboundAttachments { + return InboundMessage{}, fmt.Errorf("%w: too many attachments", ErrInvalid) + } + for index, value := range clone.Attachments { + normalized, err := value.Normalize() + if err != nil { + return InboundMessage{}, fmt.Errorf("%w: attachment %d: %v", ErrInvalid, index, err) + } + clone.Attachments[index] = normalized } if clone.ExternalMessageID != "" { if err := validateExternalID(clone.ExternalMessageID, "external message ID"); err != nil { diff --git a/trpcservice/gateway/http.go b/trpcservice/gateway/http.go index a8323b7d..8aa6164d 100644 --- a/trpcservice/gateway/http.go +++ b/trpcservice/gateway/http.go @@ -386,7 +386,12 @@ func (handler *HTTPHandler) decodeMessage(writer http.ResponseWriter, request *h if err := decoder.Decode(&trailing); err != io.EOF { return InboundMessage{}, fmt.Errorf("%w: request JSON has trailing data", ErrInvalid) } - message := InboundMessage(input) + message := InboundMessage{ + Content: input.Content, ContentType: input.ContentType, + ExternalMessageID: input.ExternalMessageID, ExternalUserID: input.ExternalUserID, + ConversationKind: input.ConversationKind, ExternalPeerID: input.ExternalPeerID, + ExternalChatID: input.ExternalChatID, ExternalThreadID: input.ExternalThreadID, + } return message.Normalize() } diff --git a/trpcservice/gateway/message.go b/trpcservice/gateway/message.go new file mode 100644 index 00000000..f083a2d0 --- /dev/null +++ b/trpcservice/gateway/message.go @@ -0,0 +1,51 @@ +package gateway + +import ( + "context" + "fmt" + "strings" + + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" + trpcmodel "trpc.group/trpc-go/trpc-agent-go/model" +) + +func buildUserMessage(ctx context.Context, reader attachment.Reader, tenantID, eventID string, inbound InboundMessage) (trpcmodel.Message, error) { + message := trpcmodel.Message{Role: trpcmodel.RoleUser, Content: inbound.Content} + if len(inbound.Attachments) == 0 { + return message, nil + } + if reader == nil { + return trpcmodel.Message{}, fmt.Errorf("attachment reader is required") + } + for index, reference := range inbound.Attachments { + content, err := reader.Load(ctx, tenantID, eventID, reference) + if err != nil { + return trpcmodel.Message{}, fmt.Errorf("load attachment %d: %w", index, err) + } + if err := content.Validate(reference); err != nil { + return trpcmodel.Message{}, fmt.Errorf("validate attachment %d: %w", index, err) + } + message.ContentParts = append(message.ContentParts, contentPart(reference, content)) + } + return message, nil +} + +func contentPart(reference attachment.Reference, content attachment.Content) trpcmodel.ContentPart { + switch reference.Kind { + case attachment.KindImage: + return trpcmodel.ContentPart{Type: trpcmodel.ContentTypeImage, Image: &trpcmodel.Image{Data: content.Data, Detail: "auto", Format: mediaSubtype(reference.MIMEType)}} + case attachment.KindAudio: + return trpcmodel.ContentPart{Type: trpcmodel.ContentTypeAudio, Audio: &trpcmodel.Audio{Data: content.Data, Format: mediaSubtype(reference.MIMEType)}} + case attachment.KindVideo: + return trpcmodel.ContentPart{Type: trpcmodel.ContentTypeVideo, Video: &trpcmodel.Video{Data: content.Data, Format: mediaSubtype(reference.MIMEType)}} + default: + return trpcmodel.ContentPart{Type: trpcmodel.ContentTypeFile, File: &trpcmodel.File{Name: reference.Name, Data: content.Data, MimeType: reference.MIMEType}} + } +} + +func mediaSubtype(contentType string) string { + if _, subtype, found := strings.Cut(contentType, "/"); found { + return subtype + } + return contentType +} diff --git a/trpcservice/gateway/message_test.go b/trpcservice/gateway/message_test.go new file mode 100644 index 00000000..99b17974 --- /dev/null +++ b/trpcservice/gateway/message_test.go @@ -0,0 +1,112 @@ +package gateway + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "testing" + + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" + trpcmodel "trpc.group/trpc-go/trpc-agent-go/model" +) + +func TestInboundMessageNormalizesAttachmentOnlyInput(t *testing.T) { + reference := testAttachmentReference(t, attachment.KindImage, "image/png", []byte("image")) + message, err := (InboundMessage{Attachments: []attachment.Reference{reference}, ExternalMessageID: "message", ExternalUserID: "user", ConversationKind: ConversationDirect, ExternalPeerID: "peer"}).Normalize() + if err != nil { + t.Fatalf("Normalize = %v", err) + } + if message.ContentType != ContentTypeMedia || len(message.Attachments) != 1 { + t.Fatalf("normalized message = %+v", message) + } +} + +func TestInboundMessageRejectsAttachmentBoundaries(t *testing.T) { + reference := testAttachmentReference(t, attachment.KindImage, "image/png", []byte("image")) + attachments := make([]attachment.Reference, maxInboundAttachments+1) + for index := range attachments { + attachments[index] = reference + } + if _, err := (InboundMessage{Attachments: attachments, ExternalMessageID: "message", ExternalUserID: "user", ConversationKind: ConversationDirect, ExternalPeerID: "peer"}).Normalize(); !errors.Is(err, ErrInvalid) { + t.Fatalf("too many attachments error = %v", err) + } + invalid := reference + invalid.ID = "" + if _, err := (InboundMessage{Attachments: []attachment.Reference{invalid}, ExternalMessageID: "message", ExternalUserID: "user", ConversationKind: ConversationDirect, ExternalPeerID: "peer"}).Normalize(); !errors.Is(err, ErrInvalid) { + t.Fatalf("invalid attachment error = %v", err) + } +} + +func TestBuildUserMessageUsesVerifiedContentParts(t *testing.T) { + data := []byte("image") + reference := testAttachmentReference(t, attachment.KindImage, "image/png", data) + reader := attachmentReaderFunc(func(context.Context, string, string, attachment.Reference) (attachment.Content, error) { + return attachment.Content{Data: data}, nil + }) + message, err := buildUserMessage(context.Background(), reader, "tenant", "event", InboundMessage{Content: "describe", Attachments: []attachment.Reference{reference}}) + if err != nil { + t.Fatalf("buildUserMessage = %v", err) + } + if message.Content != "describe" || len(message.ContentParts) != 1 || message.ContentParts[0].Type != trpcmodel.ContentTypeImage || string(message.ContentParts[0].Image.Data) != string(data) { + t.Fatalf("message = %+v", message) + } +} + +func TestContentPartCoversAttachmentFamilies(t *testing.T) { + for _, test := range []struct { + name string + ref attachment.Reference + want trpcmodel.ContentType + }{ + {name: "audio", ref: testAttachmentReference(t, attachment.KindAudio, "audio/mpeg", []byte("mp3")), want: trpcmodel.ContentTypeAudio}, + {name: "video", ref: testAttachmentReference(t, attachment.KindVideo, "video/mp4", []byte("mp4")), want: trpcmodel.ContentTypeVideo}, + {name: "document", ref: testAttachmentReference(t, attachment.KindDocument, "application/pdf", []byte("pdf")), want: trpcmodel.ContentTypeFile}, + } { + t.Run(test.name, func(t *testing.T) { + part := contentPart(test.ref, attachment.Content{Data: []byte("content")}) + if part.Type != test.want { + t.Fatalf("part type = %q, want %q", part.Type, test.want) + } + }) + } + if got := mediaSubtype("application"); got != "application" { + t.Fatalf("mediaSubtype without slash = %q", got) + } +} + +func TestBuildUserMessageRejectsUnavailableOrTamperedAttachment(t *testing.T) { + reference := testAttachmentReference(t, attachment.KindDocument, "application/pdf", []byte("document")) + message := InboundMessage{Attachments: []attachment.Reference{reference}} + if _, err := buildUserMessage(context.Background(), nil, "tenant", "event", message); err == nil { + t.Fatal("buildUserMessage accepted nil reader") + } + reader := attachmentReaderFunc(func(context.Context, string, string, attachment.Reference) (attachment.Content, error) { + return attachment.Content{Data: []byte("tampered")}, nil + }) + if _, err := buildUserMessage(context.Background(), reader, "tenant", "event", message); err == nil { + t.Fatal("buildUserMessage accepted tampered content") + } + reader = attachmentReaderFunc(func(context.Context, string, string, attachment.Reference) (attachment.Content, error) { + return attachment.Content{}, errors.New("unavailable") + }) + if _, err := buildUserMessage(context.Background(), reader, "tenant", "event", message); err == nil { + t.Fatal("buildUserMessage accepted reader error") + } +} + +type attachmentReaderFunc func(context.Context, string, string, attachment.Reference) (attachment.Content, error) + +func (function attachmentReaderFunc) Load(ctx context.Context, tenantID, eventID string, reference attachment.Reference) (attachment.Content, error) { + return function(ctx, tenantID, eventID, reference) +} + +func testAttachmentReference(t *testing.T, kind attachment.Kind, contentType string, data []byte) attachment.Reference { + t.Helper() + digest := sha256.Sum256(data) + reference := attachment.Reference{ID: "attachment-1", Kind: kind, MIMEType: contentType, Size: int64(len(data)), SHA256: hex.EncodeToString(digest[:])} + if _, err := reference.Normalize(); err != nil { + t.Fatalf("test reference = %v", err) + } + return reference +} diff --git a/trpcservice/runtime/storage/attachment.go b/trpcservice/runtime/storage/attachment.go new file mode 100644 index 00000000..8a68a6cb --- /dev/null +++ b/trpcservice/runtime/storage/attachment.go @@ -0,0 +1,20 @@ +package storage + +import ( + "context" + "io" + "time" + + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" +) + +// AttachmentStore persists tenant-scoped attachment data and lifecycle metadata. +// An attachment is first stored, then bound to a durable message event by the +// Gateway before it can be read by a Runner. +type AttachmentStore interface { + PutAttachment(context.Context, string, attachment.Upload, io.Reader) (attachment.Reference, error) + Load(context.Context, string, string, attachment.Reference) (attachment.Content, error) + BindAttachments(context.Context, string, string, []attachment.Reference) error + CleanupAttachments(context.Context, string, time.Time) (int, error) + Close() error +} diff --git a/trpcservice/runtime/storage/capabilities.go b/trpcservice/runtime/storage/capabilities.go index 57e73683..691be994 100644 --- a/trpcservice/runtime/storage/capabilities.go +++ b/trpcservice/runtime/storage/capabilities.go @@ -210,6 +210,7 @@ type RuntimeCapabilities interface { AuditStore VectorStore ObjectStore + AttachmentStore } // SessionStore is the tenant-scoped session runtime contract. diff --git a/trpcservice/runtime/storage/inmemory/attachment.go b/trpcservice/runtime/storage/inmemory/attachment.go new file mode 100644 index 00000000..b36cc7c0 --- /dev/null +++ b/trpcservice/runtime/storage/inmemory/attachment.go @@ -0,0 +1,149 @@ +package inmemory + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "io" + "strings" + "time" + + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" + runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" +) + +type storedAttachment struct { + reference attachment.Reference + eventID string + expiresAt time.Time +} + +// PutAttachment stores exactly one bounded attachment and its unbound lifecycle record. +func (s *Store) PutAttachment(ctx context.Context, tenantID string, upload attachment.Upload, content io.Reader) (attachment.Reference, error) { + if err := check(ctx); err != nil { + return attachment.Reference{}, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || content == nil { + return attachment.Reference{}, runtimestorage.ErrInvalid + } + normalized, err := upload.Normalize(time.Now().UTC()) + if err != nil { + return attachment.Reference{}, err + } + data, err := io.ReadAll(io.LimitReader(content, normalized.Size+1)) + if err != nil { + return attachment.Reference{}, runtimestorage.ErrStorage + } + if int64(len(data)) != normalized.Size { + return attachment.Reference{}, attachment.ErrInvalid + } + digest := sha256.Sum256(data) + reference := attachment.Reference{ID: normalized.ID, Kind: normalized.Kind, MIMEType: normalized.MIMEType, Name: normalized.Name, Size: normalized.Size, SHA256: hex.EncodeToString(digest[:]), Provider: normalized.Provider, ProviderID: normalized.ProviderID} + if _, err := reference.Normalize(); err != nil { + return attachment.Reference{}, err + } + s.mu.Lock() + defer s.mu.Unlock() + key := key(tenantID, reference.ID) + if existing, ok := s.attachments[key]; ok { + if existing.reference != reference { + return attachment.Reference{}, runtimestorage.ErrConflict + } + return existing.reference, nil + } + if object, ok := s.objects[key]; ok && (object.ContentType != reference.MIMEType || object.Size != reference.Size || object.ETag != reference.SHA256) { + return attachment.Reference{}, runtimestorage.ErrConflict + } + s.objectData[key] = append([]byte(nil), data...) + s.objects[key] = runtimestorage.ObjectInfo{TenantID: tenantID, ObjectKey: reference.ID, ContentType: reference.MIMEType, Size: reference.Size, ETag: reference.SHA256, CreatedAt: time.Now().UTC()} + s.attachments[key] = storedAttachment{reference: reference, expiresAt: normalized.ExpiresAt} + return reference, nil +} + +// BindAttachments associates unexpired attachment references with an existing event. +func (s *Store) BindAttachments(ctx context.Context, tenantID, eventID string, references []attachment.Reference) error { + if err := check(ctx); err != nil { + return err + } + if runtimestorage.ValidateTenant(tenantID) != nil || eventID == "" { + return runtimestorage.ErrInvalid + } + s.mu.Lock() + defer s.mu.Unlock() + if _, ok := s.events[key(tenantID, eventID)]; !ok { + return runtimestorage.ErrNotFound + } + now := time.Now().UTC() + for _, reference := range references { + normalized, err := reference.Normalize() + if err != nil { + return err + } + stored, ok := s.attachments[key(tenantID, normalized.ID)] + if !ok { + return runtimestorage.ErrNotFound + } + if stored.reference != normalized || !stored.expiresAt.After(now) || stored.eventID != "" && stored.eventID != eventID { + return runtimestorage.ErrConflict + } + } + for _, reference := range references { + stored := s.attachments[key(tenantID, reference.ID)] + stored.eventID = eventID + s.attachments[key(tenantID, reference.ID)] = stored + } + return nil +} + +// Load returns attachment data only when its reference belongs to eventID. +func (s *Store) Load(ctx context.Context, tenantID, eventID string, reference attachment.Reference) (attachment.Content, error) { + if err := check(ctx); err != nil { + return attachment.Content{}, err + } + normalized, err := reference.Normalize() + if err != nil { + return attachment.Content{}, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || eventID == "" { + return attachment.Content{}, runtimestorage.ErrInvalid + } + s.mu.RLock() + stored, ok := s.attachments[key(tenantID, normalized.ID)] + data := append([]byte(nil), s.objectData[key(tenantID, normalized.ID)]...) + s.mu.RUnlock() + if !ok || stored.reference != normalized || stored.eventID != eventID || !stored.expiresAt.After(time.Now().UTC()) { + return attachment.Content{}, runtimestorage.ErrNotFound + } + result := attachment.Content{Data: data} + if err := result.Validate(normalized); err != nil { + return attachment.Content{}, err + } + return result, nil +} + +// CleanupAttachments removes expired unbound attachments and attachments whose +// bound message has reached a terminal state. +func (s *Store) CleanupAttachments(ctx context.Context, tenantID string, before time.Time) (int, error) { + if err := check(ctx); err != nil { + return 0, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || before.IsZero() { + return 0, runtimestorage.ErrInvalid + } + s.mu.Lock() + defer s.mu.Unlock() + removed := 0 + for storageKey, value := range s.attachments { + event, found := s.events[key(tenantID, value.eventID)] + terminal := value.eventID == "" || found && (event.Status == runtimestorage.EventCompleted || event.Status == runtimestorage.EventFailed) + if terminal && !value.expiresAt.After(before) && strings.HasPrefix(storageKey, key(tenantID)) { + delete(s.attachments, storageKey) + delete(s.objects, storageKey) + delete(s.objectData, storageKey) + removed++ + } + } + return removed, nil +} + +var _ runtimestorage.AttachmentStore = (*Store)(nil) diff --git a/trpcservice/runtime/storage/inmemory/attachment_test.go b/trpcservice/runtime/storage/inmemory/attachment_test.go new file mode 100644 index 00000000..85e233fc --- /dev/null +++ b/trpcservice/runtime/storage/inmemory/attachment_test.go @@ -0,0 +1,244 @@ +package inmemory + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "strings" + "testing" + "time" + + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" + runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" +) + +func TestAttachmentStoreScopesBindsAndCleansUp(t *testing.T) { + store := New() + t.Cleanup(func() { _ = store.Close() }) + data := []byte("document") + digest := sha256.Sum256(data) + stored, err := store.PutAttachment(context.Background(), "tenant-a", attachment.Upload{ID: "attachment-1", Kind: attachment.KindDocument, MIMEType: "application/pdf", Size: int64(len(data))}, strings.NewReader(string(data))) + if err != nil { + t.Fatalf("PutAttachment = %v", err) + } + reference := attachment.Reference{ID: "attachment-1", Kind: attachment.KindDocument, MIMEType: "application/pdf", Size: int64(len(data)), SHA256: hex.EncodeToString(digest[:])} + if stored != reference { + t.Fatalf("stored reference = %+v", stored) + } + if err := store.BindAttachments(context.Background(), "tenant-a", "event-1", []attachment.Reference{reference}); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("binding unknown event = %v", err) + } + if _, err := store.CreateSession(context.Background(), "tenant-a", "session-1", nil); err != nil { + t.Fatalf("CreateSession = %v", err) + } + if _, _, err := store.RecordMessage(context.Background(), runtimestorage.MessageEventInput{TenantID: "tenant-a", EventID: "event-1", SessionID: "session-1", BindingID: "binding-1", ExternalMessageID: "external-1"}); err != nil { + t.Fatalf("RecordMessage = %v", err) + } + if err := store.BindAttachments(context.Background(), "tenant-a", "event-1", []attachment.Reference{reference}); err != nil { + t.Fatalf("BindAttachments = %v", err) + } + content, err := store.Load(context.Background(), "tenant-a", "event-1", reference) + if err != nil || string(content.Data) != string(data) { + t.Fatalf("Load = %q, %v", content.Data, err) + } + if _, err := store.Load(context.Background(), "tenant-b", "event-1", reference); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("cross-tenant Load = %v", err) + } + reference.Size++ + if _, err := store.Load(context.Background(), "tenant-a", "event-1", reference); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("size mismatch Load = %v", err) + } + if removed, err := store.CleanupAttachments(context.Background(), "tenant-a", time.Now().UTC().Add(attachment.DefaultRetention)); err != nil || removed != 0 { + t.Fatalf("bound cleanup = %d, %v", removed, err) + } + if err := store.DeleteSession(context.Background(), "tenant-a", "session-1"); err != nil { + t.Fatalf("DeleteSession = %v", err) + } + if removed, err := store.CleanupAttachments(context.Background(), "tenant-a", time.Now().UTC().Add(attachment.DefaultRetention+time.Second)); err != nil || removed != 1 { + t.Fatalf("orphaned cleanup = %d, %v", removed, err) + } +} + +func TestDeleteSessionDoesNotUnbindAnotherTenantAttachment(t *testing.T) { + store := New() + t.Cleanup(func() { _ = store.Close() }) + for _, tenantID := range []string{"tenant-a", "tenant-b"} { + if _, err := store.CreateSession(context.Background(), tenantID, "session-1", nil); err != nil { + t.Fatal(err) + } + if _, _, err := store.RecordMessage(context.Background(), runtimestorage.MessageEventInput{TenantID: tenantID, EventID: "event-shared", SessionID: "session-1", BindingID: "binding-1", ExternalMessageID: "external-1"}); err != nil { + t.Fatal(err) + } + } + data := []byte("tenant-b-data") + digest := sha256.Sum256(data) + reference, err := store.PutAttachment(context.Background(), "tenant-b", attachment.Upload{ID: "attachment-b", Kind: attachment.KindDocument, MIMEType: "application/pdf", Size: int64(len(data))}, strings.NewReader(string(data))) + if err != nil { + t.Fatal(err) + } + if reference.SHA256 != hex.EncodeToString(digest[:]) { + t.Fatalf("reference digest = %q", reference.SHA256) + } + if err := store.BindAttachments(context.Background(), "tenant-b", "event-shared", []attachment.Reference{reference}); err != nil { + t.Fatal(err) + } + if err := store.DeleteSession(context.Background(), "tenant-a", "session-1"); err != nil { + t.Fatal(err) + } + if content, err := store.Load(context.Background(), "tenant-b", "event-shared", reference); err != nil || string(content.Data) != string(data) { + t.Fatalf("tenant-b attachment after tenant-a delete = %q, %v", content.Data, err) + } +} + +func TestAttachmentStoreRejectsInvalidInputsAndConflicts(t *testing.T) { + store := New() + t.Cleanup(func() { _ = store.Close() }) + ctx := context.Background() + data := []byte("document") + upload := attachment.Upload{ID: "attachment-1", Kind: attachment.KindDocument, MIMEType: "application/pdf", Size: int64(len(data))} + + canceled, cancel := context.WithCancel(ctx) + cancel() + if _, err := store.PutAttachment(canceled, "tenant-a", upload, strings.NewReader(string(data))); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled PutAttachment = %v", err) + } + if _, err := store.PutAttachment(ctx, "", upload, strings.NewReader(string(data))); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("invalid tenant PutAttachment = %v", err) + } + if _, err := store.PutAttachment(ctx, "tenant-a", upload, nil); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("nil content PutAttachment = %v", err) + } + if _, err := store.PutAttachment(ctx, "tenant-a", attachment.Upload{ID: "bad", Kind: attachment.KindImage, MIMEType: "application/pdf", Size: 1}, strings.NewReader("x")); !errors.Is(err, attachment.ErrInvalid) { + t.Fatalf("invalid upload PutAttachment = %v", err) + } + if _, err := store.PutAttachment(ctx, "tenant-a", upload, failingAttachmentReader{}); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("reader failure PutAttachment = %v", err) + } + if _, err := store.PutAttachment(ctx, "tenant-a", upload, strings.NewReader("short")); !errors.Is(err, attachment.ErrInvalid) { + t.Fatalf("size mismatch PutAttachment = %v", err) + } + + reference, err := store.PutAttachment(ctx, "tenant-a", upload, strings.NewReader(string(data))) + if err != nil { + t.Fatal(err) + } + conflict := upload + conflict.Size = int64(len("documenz")) + if _, err := store.PutAttachment(ctx, "tenant-a", conflict, strings.NewReader("documenz")); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("conflicting attachment PutAttachment = %v", err) + } + + store.mu.Lock() + delete(store.attachments, key("tenant-a", reference.ID)) + store.mu.Unlock() + if _, err := store.PutAttachment(ctx, "tenant-a", conflict, strings.NewReader("documenz")); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("object conflict PutAttachment = %v", err) + } +} + +type failingAttachmentReader struct{} + +func (failingAttachmentReader) Read([]byte) (int, error) { + return 0, errors.New("read failed") +} + +func TestAttachmentStoreBindFailures(t *testing.T) { + store, reference := newAttachmentStoreBindingFixture(t) + ctx := context.Background() + canceled, cancel := context.WithCancel(ctx) + cancel() + + if err := store.BindAttachments(canceled, "tenant-a", "event-1", []attachment.Reference{reference}); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled BindAttachments = %v", err) + } + if err := store.BindAttachments(ctx, "", "event-1", []attachment.Reference{reference}); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("invalid tenant BindAttachments = %v", err) + } + if err := store.BindAttachments(ctx, "tenant-a", "", []attachment.Reference{reference}); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("invalid event BindAttachments = %v", err) + } + badReference := reference + badReference.ID = "bad\nid" + if err := store.BindAttachments(ctx, "tenant-a", "event-1", []attachment.Reference{badReference}); !errors.Is(err, attachment.ErrInvalid) { + t.Fatalf("invalid reference BindAttachments = %v", err) + } + missing := reference + missing.ID = "attachment-missing" + if err := store.BindAttachments(ctx, "tenant-a", "event-1", []attachment.Reference{missing}); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("missing attachment BindAttachments = %v", err) + } + if err := store.BindAttachments(ctx, "tenant-a", "event-1", []attachment.Reference{reference}); err != nil { + t.Fatal(err) + } + if err := store.BindAttachments(ctx, "tenant-a", "event-2", []attachment.Reference{reference}); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("rebind conflict = %v", err) + } +} + +func TestAttachmentStoreLoadFailures(t *testing.T) { + store, reference := newAttachmentStoreBindingFixture(t) + ctx := context.Background() + if err := store.BindAttachments(ctx, "tenant-a", "event-1", []attachment.Reference{reference}); err != nil { + t.Fatal(err) + } + canceled, cancel := context.WithCancel(ctx) + cancel() + + if _, err := store.Load(canceled, "tenant-a", "event-1", reference); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled Load = %v", err) + } + if _, err := store.Load(ctx, "", "event-1", reference); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("invalid tenant Load = %v", err) + } + if _, err := store.Load(ctx, "tenant-a", "", reference); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("invalid event Load = %v", err) + } + badReference := reference + badReference.ID = "bad\nid" + if _, err := store.Load(ctx, "tenant-a", "event-1", badReference); !errors.Is(err, attachment.ErrInvalid) { + t.Fatalf("invalid reference Load = %v", err) + } + if _, err := store.Load(ctx, "tenant-a", "event-2", reference); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("wrong event Load = %v", err) + } +} + +func TestAttachmentStoreCleanupFailures(t *testing.T) { + store, _ := newAttachmentStoreBindingFixture(t) + ctx := context.Background() + canceled, cancel := context.WithCancel(ctx) + cancel() + + if removed, err := store.CleanupAttachments(canceled, "tenant-a", time.Now()); !errors.Is(err, context.Canceled) || removed != 0 { + t.Fatalf("canceled CleanupAttachments = %d, %v", removed, err) + } + if removed, err := store.CleanupAttachments(ctx, "", time.Now()); !errors.Is(err, runtimestorage.ErrInvalid) || removed != 0 { + t.Fatalf("invalid tenant CleanupAttachments = %d, %v", removed, err) + } + if removed, err := store.CleanupAttachments(ctx, "tenant-a", time.Time{}); !errors.Is(err, runtimestorage.ErrInvalid) || removed != 0 { + t.Fatalf("zero before CleanupAttachments = %d, %v", removed, err) + } +} + +func newAttachmentStoreBindingFixture(t *testing.T) (*Store, attachment.Reference) { + t.Helper() + store := New() + t.Cleanup(func() { _ = store.Close() }) + ctx := context.Background() + data := []byte("document") + reference, err := store.PutAttachment(ctx, "tenant-a", attachment.Upload{ID: "attachment-1", Kind: attachment.KindDocument, MIMEType: "application/pdf", Size: int64(len(data))}, strings.NewReader(string(data))) + if err != nil { + t.Fatal(err) + } + if _, err := store.CreateSession(ctx, "tenant-a", "session-1", nil); err != nil { + t.Fatal(err) + } + if _, _, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", EventID: "event-1", SessionID: "session-1", BindingID: "binding-1", ExternalMessageID: "external-1"}); err != nil { + t.Fatal(err) + } + if _, _, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", EventID: "event-2", SessionID: "session-1", BindingID: "binding-1", ExternalMessageID: "external-2"}); err != nil { + t.Fatal(err) + } + return store, reference +} diff --git a/trpcservice/runtime/storage/inmemory/inmemory.go b/trpcservice/runtime/storage/inmemory/inmemory.go index 70e880c6..f7ed20f6 100644 --- a/trpcservice/runtime/storage/inmemory/inmemory.go +++ b/trpcservice/runtime/storage/inmemory/inmemory.go @@ -31,6 +31,7 @@ type Store struct { vectors map[string]runtimestorage.VectorRecord objects map[string]runtimestorage.ObjectInfo objectData map[string][]byte + attachments map[string]storedAttachment indexQueue chan runtimestorage.MemoryRecord indexDone chan struct{} indexMu *sync.RWMutex @@ -136,7 +137,7 @@ func newStore(lifecycle *backendLifecycle) *Store { memories: map[string]runtimestorage.MemoryRecord{}, summaries: map[string]runtimestorage.SummaryRecord{}, knowledge: map[string]runtimestorage.KnowledgeDocument{}, artifacts: map[string]runtimestorage.ArtifactRecord{}, audits: map[string][]runtimestorage.AuditRecord{}, vectors: map[string]runtimestorage.VectorRecord{}, - objects: map[string]runtimestorage.ObjectInfo{}, objectData: map[string][]byte{}, + objects: map[string]runtimestorage.ObjectInfo{}, objectData: map[string][]byte{}, attachments: map[string]storedAttachment{}, indexQueue: make(chan runtimestorage.MemoryRecord, 128), indexDone: lifecycle.done, indexMu: lifecycle.indexMu, lifecycle: lifecycle, closeOnce: &sync.Once{}, } go store.indexWorker() @@ -253,6 +254,12 @@ func (s *Store) DeleteSession(ctx context.Context, tenantID, sessionID string) e delete(s.replies, replyKey) } } + for attachmentKey, attachmentValue := range s.attachments { + if attachmentValue.eventID == event.EventID && strings.HasPrefix(attachmentKey, key(tenantID)) { + attachmentValue.eventID = "" + s.attachments[attachmentKey] = attachmentValue + } + } } return nil } @@ -438,30 +445,18 @@ func (s *Store) EnqueueReply(ctx context.Context, value runtimestorage.ReplyOutb if err := check(ctx); err != nil { return runtimestorage.ReplyOutbox{}, err } - if err := runtimestorage.ValidateTenant(value.TenantID); err != nil || value.ReplyID == "" || value.EventID == "" || value.SegmentIndex < 0 || value.SegmentCount <= value.SegmentIndex || runtimestorage.ValidateReplyTarget(value.ReplyTarget) != nil { - return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid - } - if value.Status == "" { - value.Status = runtimestorage.ReplyPending - } - if value.Status != runtimestorage.ReplyPending { + value, err := prepareReply(value) + if err != nil { return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid } - now := time.Now().UTC() - value.CreatedAt = now - value.UpdatedAt = now s.mu.Lock() defer s.mu.Unlock() - event, ok := s.events[key(value.TenantID, value.EventID)] - if !ok { - return runtimestorage.ReplyOutbox{}, runtimestorage.ErrNotFound - } - if event.ReplyTarget != value.ReplyTarget { - return runtimestorage.ReplyOutbox{}, runtimestorage.ErrConflict + if err := s.validateReplyEvent(value); err != nil { + return runtimestorage.ReplyOutbox{}, err } k := replyKey(value.TenantID, value.ReplyID, value.SegmentIndex) if existing, ok := s.replies[k]; ok { - if existing.EventID != value.EventID || existing.SegmentCount != value.SegmentCount || existing.Payload != value.Payload || existing.ReplyTarget != value.ReplyTarget { + if !sameReplyContract(existing, value) { return runtimestorage.ReplyOutbox{}, runtimestorage.ErrConflict } return cloneReply(existing), nil @@ -470,6 +465,54 @@ func (s *Store) EnqueueReply(ctx context.Context, value runtimestorage.ReplyOutb return cloneReply(value), nil } +func prepareReply(value runtimestorage.ReplyOutbox) (runtimestorage.ReplyOutbox, error) { + normalized, err := runtimestorage.NormalizeReplyOutbox(value) + if err != nil { + return runtimestorage.ReplyOutbox{}, err + } + if err := validateReplySegment(normalized); err != nil { + return runtimestorage.ReplyOutbox{}, err + } + if normalized.Status == "" { + normalized.Status = runtimestorage.ReplyPending + } + if normalized.Status != runtimestorage.ReplyPending { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + now := time.Now().UTC() + normalized.CreatedAt = now + normalized.UpdatedAt = now + return normalized, nil +} + +func validateReplySegment(value runtimestorage.ReplyOutbox) error { + if runtimestorage.ValidateTenant(value.TenantID) != nil || value.ReplyID == "" || value.EventID == "" || value.SegmentIndex < 0 || value.SegmentCount <= value.SegmentIndex || runtimestorage.ValidateReplyTarget(value.ReplyTarget) != nil { + return runtimestorage.ErrInvalid + } + return nil +} + +func (s *Store) validateReplyEvent(value runtimestorage.ReplyOutbox) error { + event, ok := s.events[key(value.TenantID, value.EventID)] + if !ok { + return runtimestorage.ErrNotFound + } + if event.ReplyTarget != value.ReplyTarget { + return runtimestorage.ErrConflict + } + return nil +} + +func sameReplyContract(existing, value runtimestorage.ReplyOutbox) bool { + return existing.EventID == value.EventID && + existing.SegmentCount == value.SegmentCount && + existing.Payload == value.Payload && + existing.Kind == value.Kind && + existing.Attachment == value.Attachment && + existing.Fallback == value.Fallback && + existing.ReplyTarget == value.ReplyTarget +} + // EnqueueReplies validates a complete reply before committing any new segment. // This prevents a failed multi-segment materialization from exposing a // deliverable prefix to a worker. @@ -493,6 +536,11 @@ func (s *Store) enqueueReplies(ctx context.Context, correlation runtimestorage.R if err := check(ctx); err != nil { return nil, err } + var err error + values, err = normalizeReplyBatch(values) + if err != nil { + return nil, err + } first, _, err := validateReplyBatch(values) if err != nil { return nil, err @@ -562,9 +610,21 @@ func validateReplyBatch(values []runtimestorage.ReplyOutbox) (runtimestorage.Rep return first, seen, nil } +func normalizeReplyBatch(values []runtimestorage.ReplyOutbox) ([]runtimestorage.ReplyOutbox, error) { + normalized := make([]runtimestorage.ReplyOutbox, 0, len(values)) + for _, value := range values { + reply, err := runtimestorage.NormalizeReplyOutbox(value) + if err != nil { + return nil, runtimestorage.ErrInvalid + } + normalized = append(normalized, reply) + } + return normalized, nil +} + func validateExistingReplies(replies map[string]runtimestorage.ReplyOutbox, values []runtimestorage.ReplyOutbox) error { for _, value := range values { - if existing, ok := replies[replyKey(value.TenantID, value.ReplyID, value.SegmentIndex)]; ok && (existing.EventID != value.EventID || existing.SegmentCount != value.SegmentCount || existing.Payload != value.Payload || existing.ReplyTarget != value.ReplyTarget) { + if existing, ok := replies[replyKey(value.TenantID, value.ReplyID, value.SegmentIndex)]; ok && (existing.EventID != value.EventID || existing.SegmentCount != value.SegmentCount || existing.Payload != value.Payload || existing.Kind != value.Kind || existing.Attachment != value.Attachment || existing.Fallback != value.Fallback || existing.ReplyTarget != value.ReplyTarget) { return runtimestorage.ErrConflict } } diff --git a/trpcservice/runtime/storage/inmemory/inmemory_test.go b/trpcservice/runtime/storage/inmemory/inmemory_test.go index c666f0af..8b126f07 100644 --- a/trpcservice/runtime/storage/inmemory/inmemory_test.go +++ b/trpcservice/runtime/storage/inmemory/inmemory_test.go @@ -2,11 +2,14 @@ package inmemory_test import ( "context" + "crypto/sha256" + "encoding/hex" "errors" "sync" "testing" "time" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/inmemory" ) @@ -126,6 +129,36 @@ func TestStoreReplyCorrelationNormalizesTraceParentAtPersistenceBoundary(t *test } } +func TestStorePersistsMediaReplyContractAndDetectsConflicts(t *testing.T) { + store := inmemory.New() + seedEvent(t, store, "tenant-a", "session-media-reply", "event-media-reply") + reference := mediaReplyReference(t, attachment.KindImage, "image/png", []byte("png")) + reply := runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", ReplyID: "reply-media", EventID: "event-media-reply", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Attachment: reference, Fallback: "[image attachment: chart.png]", + } + first, err := store.EnqueueReply(context.Background(), reply) + if err != nil { + t.Fatal(err) + } + if first.Kind != runtimestorage.ReplyKindImage || first.Attachment != reference || first.Fallback != reply.Fallback { + t.Fatalf("stored media reply = %+v", first) + } + if second, err := store.EnqueueReply(context.Background(), reply); err != nil || second.Attachment != reference { + t.Fatalf("idempotent media reply = %+v err=%v", second, err) + } + conflict := reply + conflict.Fallback = "[image attachment: changed]" + if _, err := store.EnqueueReply(context.Background(), conflict); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("media conflict = %v", err) + } + invalid := reply + invalid.Fallback = "" + if _, err := store.EnqueueReply(context.Background(), invalid); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("invalid media fallback = %v", err) + } +} + func TestStoreDuplicateMessageAndConcurrentSequence(t *testing.T) { store := inmemory.New() if _, err := store.CreateSession(context.Background(), "tenant-a", "session-1", nil); err != nil { @@ -165,6 +198,16 @@ func TestStoreDuplicateMessageAndConcurrentSequence(t *testing.T) { } } +func mediaReplyReference(t *testing.T, kind attachment.Kind, contentType string, data []byte) attachment.Reference { + t.Helper() + digest := sha256.Sum256(data) + reference := attachment.Reference{ID: "attachment-media", Kind: kind, MIMEType: contentType, Name: "chart.png", Size: int64(len(data)), SHA256: hex.EncodeToString(digest[:])} + if _, err := reference.Normalize(); err != nil { + t.Fatalf("test attachment = %v", err) + } + return reference +} + func TestStorePersistsFirstReplyTargetForDuplicateMessage(t *testing.T) { store := inmemory.New() if _, err := store.CreateSession(context.Background(), "tenant-a", "session-1", nil); err != nil { diff --git a/trpcservice/runtime/storage/postgres/attachment.go b/trpcservice/runtime/storage/postgres/attachment.go new file mode 100644 index 00000000..7c7efb0c --- /dev/null +++ b/trpcservice/runtime/storage/postgres/attachment.go @@ -0,0 +1,278 @@ +package postgres + +import ( + "bytes" + "context" + "crypto/sha256" + "database/sql" + "encoding/hex" + "errors" + "fmt" + "io" + "time" + + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" + runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" + pgstorage "github.com/XnLemon/trpc-agent-service/trpcservice/storage/postgres" +) + +const attachmentColumns = "tenant_id,attachment_id,kind,mime_type,name,size,sha256,provider,provider_id,event_id,expires_at" + +type storedAttachment struct { + reference attachment.Reference + eventID sql.NullString + expiresAt time.Time +} + +// PutAttachment writes immutable bytes and their lifecycle metadata in one transaction. +func (s *Store) PutAttachment(ctx context.Context, tenantID string, upload attachment.Upload, content io.Reader) (attachment.Reference, error) { + if err := checkCapability(ctx, s); err != nil { + return attachment.Reference{}, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || content == nil { + return attachment.Reference{}, runtimestorage.ErrInvalid + } + reference, data, expiresAt, err := prepareAttachment(upload, content) + if err != nil { + return attachment.Reference{}, err + } + return s.persistAttachment(ctx, tenantID, reference, data, expiresAt) +} + +func prepareAttachment(upload attachment.Upload, content io.Reader) (attachment.Reference, []byte, time.Time, error) { + normalized, err := upload.Normalize(time.Now().UTC()) + if err != nil { + return attachment.Reference{}, nil, time.Time{}, err + } + data, err := io.ReadAll(io.LimitReader(content, normalized.Size+1)) + if err != nil { + return attachment.Reference{}, nil, time.Time{}, runtimestorage.ErrStorage + } + if int64(len(data)) != normalized.Size { + return attachment.Reference{}, nil, time.Time{}, attachment.ErrInvalid + } + digest := sha256.Sum256(data) + reference := attachment.Reference{ID: normalized.ID, Kind: normalized.Kind, MIMEType: normalized.MIMEType, Name: normalized.Name, Size: normalized.Size, SHA256: hex.EncodeToString(digest[:]), Provider: normalized.Provider, ProviderID: normalized.ProviderID} + if _, err := reference.Normalize(); err != nil { + return attachment.Reference{}, nil, time.Time{}, err + } + return reference, data, normalized.ExpiresAt, nil +} + +func (s *Store) persistAttachment(ctx context.Context, tenantID string, reference attachment.Reference, data []byte, expiresAt time.Time) (attachment.Reference, error) { + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return attachment.Reference{}, runtimestorage.ErrStorage + } + defer func() { _ = tx.Rollback() }() + stored, found, err := loadAttachment(ctx, tx, tenantID, reference.ID, true) + if err != nil { + return attachment.Reference{}, err + } + if found { + if stored.reference != reference { + return attachment.Reference{}, runtimestorage.ErrConflict + } + return commitAttachment(tx, stored.reference) + } + if err := ensureAttachmentObject(ctx, tx, tenantID, reference, data); err != nil { + return attachment.Reference{}, err + } + if err := insertAttachmentMetadata(ctx, tx, tenantID, reference, expiresAt); err != nil { + return attachment.Reference{}, err + } + return commitAttachment(tx, reference) +} + +func ensureAttachmentObject(ctx context.Context, tx *sql.Tx, tenantID string, reference attachment.Reference, data []byte) error { + if _, err := tx.ExecContext(ctx, "INSERT INTO public.runtime_object (tenant_id,object_key,content_type,content,size,etag) VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT (tenant_id,object_key) DO NOTHING", tenantID, reference.ID, reference.MIMEType, data, reference.Size, reference.SHA256); err != nil { + return pgstorage.MapError(ctx, err, runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid) + } + var objectType, objectETag string + var objectSize int64 + if err := tx.QueryRowContext(ctx, "SELECT content_type,size,etag FROM public.runtime_object WHERE tenant_id=$1 AND object_key=$2", tenantID, reference.ID).Scan(&objectType, &objectSize, &objectETag); err != nil { + return pgstorage.MapError(ctx, err, runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid) + } + if objectType != reference.MIMEType || objectSize != reference.Size || objectETag != reference.SHA256 { + return runtimestorage.ErrConflict + } + return nil +} + +func insertAttachmentMetadata(ctx context.Context, tx *sql.Tx, tenantID string, reference attachment.Reference, expiresAt time.Time) error { + result, err := tx.ExecContext(ctx, "INSERT INTO public.runtime_attachment (tenant_id,attachment_id,kind,mime_type,name,size,sha256,provider,provider_id,expires_at) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10) ON CONFLICT (tenant_id,attachment_id) DO NOTHING", tenantID, reference.ID, reference.Kind, reference.MIMEType, reference.Name, reference.Size, reference.SHA256, reference.Provider, reference.ProviderID, expiresAt) + if err != nil { + return pgstorage.MapError(ctx, err, runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid) + } + created, err := result.RowsAffected() + if err != nil { + return runtimestorage.ErrStorage + } + if created != 0 { + return nil + } + stored, found, err := loadAttachment(ctx, tx, tenantID, reference.ID, true) + if err != nil { + return err + } + if !found || stored.reference != reference { + return runtimestorage.ErrConflict + } + return nil +} + +func commitAttachment(tx *sql.Tx, reference attachment.Reference) (attachment.Reference, error) { + if err := tx.Commit(); err != nil { + return attachment.Reference{}, runtimestorage.ErrStorage + } + return reference, nil +} + +// BindAttachments atomically binds references to an existing message event. +func (s *Store) BindAttachments(ctx context.Context, tenantID, eventID string, references []attachment.Reference) error { + if err := checkCapability(ctx, s); err != nil { + return err + } + if runtimestorage.ValidateTenant(tenantID) != nil || eventID == "" { + return runtimestorage.ErrInvalid + } + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return runtimestorage.ErrStorage + } + defer func() { _ = tx.Rollback() }() + var present bool + if err := tx.QueryRowContext(ctx, "SELECT EXISTS(SELECT 1 FROM public.message_event WHERE tenant_id=$1 AND event_id=$2)", tenantID, eventID).Scan(&present); err != nil || !present { + if err == nil { + return runtimestorage.ErrNotFound + } + return pgstorage.MapError(ctx, err, runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid) + } + for _, reference := range references { + normalized, err := reference.Normalize() + if err != nil { + return err + } + stored, found, err := loadAttachment(ctx, tx, tenantID, normalized.ID, true) + if err != nil { + return err + } + if !found { + return runtimestorage.ErrNotFound + } + if stored.reference != normalized || !stored.expiresAt.After(time.Now().UTC()) || stored.eventID.Valid && stored.eventID.String != eventID { + return runtimestorage.ErrConflict + } + } + for _, reference := range references { + if _, err := tx.ExecContext(ctx, "UPDATE public.runtime_attachment SET event_id=$3 WHERE tenant_id=$1 AND attachment_id=$2 AND (event_id IS NULL OR event_id=$3)", tenantID, reference.ID, eventID); err != nil { + return pgstorage.MapError(ctx, err, runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid) + } + } + if err := tx.Commit(); err != nil { + return runtimestorage.ErrStorage + } + return nil +} + +// Load returns data only for a reference previously bound to eventID. +func (s *Store) Load(ctx context.Context, tenantID, eventID string, reference attachment.Reference) (attachment.Content, error) { + if err := checkCapability(ctx, s); err != nil { + return attachment.Content{}, err + } + normalized, err := reference.Normalize() + if err != nil { + return attachment.Content{}, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || eventID == "" { + return attachment.Content{}, runtimestorage.ErrInvalid + } + stored, found, err := loadAttachment(ctx, s.db, tenantID, normalized.ID, false) + if err != nil { + return attachment.Content{}, err + } + if !found || stored.reference != normalized || !stored.eventID.Valid || stored.eventID.String != eventID || !stored.expiresAt.After(time.Now().UTC()) { + return attachment.Content{}, runtimestorage.ErrNotFound + } + var data []byte + if err := s.db.QueryRowContext(ctx, "SELECT content FROM public.runtime_object WHERE tenant_id=$1 AND object_key=$2", tenantID, normalized.ID).Scan(&data); err != nil { + return attachment.Content{}, pgstorage.MapError(ctx, err, runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid) + } + content := attachment.Content{Data: bytes.Clone(data)} + if err := content.Validate(normalized); err != nil { + return attachment.Content{}, err + } + return content, nil +} + +// CleanupAttachments deletes expired unbound attachments and attachments whose +// bound message has reached a terminal state. +func (s *Store) CleanupAttachments(ctx context.Context, tenantID string, before time.Time) (int, error) { + if err := checkCapability(ctx, s); err != nil { + return 0, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || before.IsZero() { + return 0, runtimestorage.ErrInvalid + } + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return 0, runtimestorage.ErrStorage + } + defer func() { _ = tx.Rollback() }() + rows, err := tx.QueryContext(ctx, "DELETE FROM public.runtime_attachment AS a WHERE a.tenant_id=$1 AND a.expires_at <= $2 AND (a.event_id IS NULL OR EXISTS (SELECT 1 FROM public.message_event AS e WHERE e.tenant_id=a.tenant_id AND e.event_id=a.event_id AND e.status IN ('completed','failed'))) RETURNING a.attachment_id", tenantID, before.UTC()) + if err != nil { + return 0, pgstorage.MapError(ctx, err, runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid) + } + defer func() { _ = rows.Close() }() + ids := make([]string, 0) + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return 0, runtimestorage.ErrStorage + } + ids = append(ids, id) + } + if err := rows.Err(); err != nil { + return 0, runtimestorage.ErrStorage + } + for _, id := range ids { + if _, err := tx.ExecContext(ctx, "DELETE FROM public.runtime_object WHERE tenant_id=$1 AND object_key=$2", tenantID, id); err != nil { + return 0, pgstorage.MapError(ctx, err, runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid) + } + } + if err := tx.Commit(); err != nil { + return 0, runtimestorage.ErrStorage + } + return len(ids), nil +} + +type attachmentQuerier interface { + QueryRowContext(context.Context, string, ...any) *sql.Row +} + +func loadAttachment(ctx context.Context, query attachmentQuerier, tenantID, id string, lock bool) (storedAttachment, bool, error) { + var stored storedAttachment + var storedTenantID string + var kind string + statement := "SELECT " + attachmentColumns + " FROM public.runtime_attachment WHERE tenant_id=$1 AND attachment_id=$2" + if lock { + statement += " FOR UPDATE" + } + err := query.QueryRowContext(ctx, statement, tenantID, id).Scan(&storedTenantID, &stored.reference.ID, &kind, &stored.reference.MIMEType, &stored.reference.Name, &stored.reference.Size, &stored.reference.SHA256, &stored.reference.Provider, &stored.reference.ProviderID, &stored.eventID, &stored.expiresAt) + if errors.Is(err, sql.ErrNoRows) { + return storedAttachment{}, false, nil + } + if err != nil { + return storedAttachment{}, false, pgstorage.MapError(ctx, err, runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid) + } + if storedTenantID != tenantID { + return storedAttachment{}, false, runtimestorage.ErrStorage + } + stored.reference.Kind = attachment.Kind(kind) + if _, err := stored.reference.Normalize(); err != nil { + return storedAttachment{}, false, fmt.Errorf("%w: stored attachment metadata is invalid", runtimestorage.ErrStorage) + } + return stored, true, nil +} + +var _ runtimestorage.AttachmentStore = (*Store)(nil) diff --git a/trpcservice/runtime/storage/postgres/attachment_test.go b/trpcservice/runtime/storage/postgres/attachment_test.go new file mode 100644 index 00000000..3face151 --- /dev/null +++ b/trpcservice/runtime/storage/postgres/attachment_test.go @@ -0,0 +1,632 @@ +package postgres_test + +import ( + "context" + "crypto/sha256" + "database/sql" + "encoding/hex" + "errors" + "io" + "regexp" + "strings" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" + runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" + runtimepostgres "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/postgres" +) + +var runtimeAttachmentColumns = []string{"tenant_id", "attachment_id", "kind", "mime_type", "name", "size", "sha256", "provider", "provider_id", "event_id", "expires_at"} + +func TestPostgresAttachmentLifecycleContracts(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + store := runtimepostgres.New(db) + data := []byte("document") + reference := attachmentReference(data) + when := time.Now().UTC() + + mock.ExpectBegin() + mock.ExpectQuery(regexp.QuoteMeta("SELECT tenant_id,attachment_id,kind,mime_type,name,size,sha256,provider,provider_id,event_id,expires_at FROM public.runtime_attachment WHERE tenant_id=$1 AND attachment_id=$2 FOR UPDATE")). + WithArgs("tenant-a", reference.ID).WillReturnError(sql.ErrNoRows) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO public.runtime_object (tenant_id,object_key,content_type,content,size,etag) VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT (tenant_id,object_key) DO NOTHING")). + WithArgs("tenant-a", reference.ID, reference.MIMEType, data, reference.Size, reference.SHA256).WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectQuery(regexp.QuoteMeta("SELECT content_type,size,etag FROM public.runtime_object WHERE tenant_id=$1 AND object_key=$2")). + WithArgs("tenant-a", reference.ID).WillReturnRows(sqlmock.NewRows([]string{"content_type", "size", "etag"}).AddRow(reference.MIMEType, reference.Size, reference.SHA256)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO public.runtime_attachment (tenant_id,attachment_id,kind,mime_type,name,size,sha256,provider,provider_id,expires_at) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10) ON CONFLICT (tenant_id,attachment_id) DO NOTHING")). + WithArgs("tenant-a", reference.ID, reference.Kind, reference.MIMEType, reference.Name, reference.Size, reference.SHA256, reference.Provider, reference.ProviderID, sqlmock.AnyArg()).WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + stored, err := store.PutAttachment(context.Background(), "tenant-a", attachment.Upload{ID: reference.ID, Kind: reference.Kind, MIMEType: reference.MIMEType, Size: reference.Size}, strings.NewReader(string(data))) + if err != nil || stored != reference { + t.Fatalf("PutAttachment = %+v, %v", stored, err) + } + + mock.ExpectBegin() + mock.ExpectQuery(regexp.QuoteMeta("SELECT EXISTS(SELECT 1 FROM public.message_event WHERE tenant_id=$1 AND event_id=$2)")). + WithArgs("tenant-a", "event-a").WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + mock.ExpectQuery(regexp.QuoteMeta("SELECT tenant_id,attachment_id,kind,mime_type,name,size,sha256,provider,provider_id,event_id,expires_at FROM public.runtime_attachment WHERE tenant_id=$1 AND attachment_id=$2 FOR UPDATE")). + WithArgs("tenant-a", reference.ID).WillReturnRows(runtimeAttachmentRow(reference, nil, when.Add(time.Hour))) + mock.ExpectExec(regexp.QuoteMeta("UPDATE public.runtime_attachment SET event_id=$3 WHERE tenant_id=$1 AND attachment_id=$2 AND (event_id IS NULL OR event_id=$3)")). + WithArgs("tenant-a", reference.ID, "event-a").WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + if err := store.BindAttachments(context.Background(), "tenant-a", "event-a", []attachment.Reference{reference}); err != nil { + t.Fatalf("BindAttachments = %v; expectations = %v", err, mock.ExpectationsWereMet()) + } + + mock.ExpectQuery(regexp.QuoteMeta("SELECT tenant_id,attachment_id,kind,mime_type,name,size,sha256,provider,provider_id,event_id,expires_at FROM public.runtime_attachment WHERE tenant_id=$1 AND attachment_id=$2")). + WithArgs("tenant-a", reference.ID).WillReturnRows(runtimeAttachmentRow(reference, "event-a", when.Add(time.Hour))) + mock.ExpectQuery(regexp.QuoteMeta("SELECT content FROM public.runtime_object WHERE tenant_id=$1 AND object_key=$2")). + WithArgs("tenant-a", reference.ID).WillReturnRows(sqlmock.NewRows([]string{"content"}).AddRow(data)) + content, err := store.Load(context.Background(), "tenant-a", "event-a", reference) + if err != nil || string(content.Data) != string(data) { + t.Fatalf("Load = %q, %v", content.Data, err) + } + + mock.ExpectBegin() + mock.ExpectQuery(regexp.QuoteMeta("DELETE FROM public.runtime_attachment AS a WHERE a.tenant_id=$1 AND a.expires_at <= $2 AND (a.event_id IS NULL OR EXISTS (SELECT 1 FROM public.message_event AS e WHERE e.tenant_id=a.tenant_id AND e.event_id=a.event_id AND e.status IN ('completed','failed'))) RETURNING a.attachment_id")). + WithArgs("tenant-a", when).WillReturnRows(sqlmock.NewRows([]string{"attachment_id"}).AddRow(reference.ID)) + mock.ExpectExec(regexp.QuoteMeta("DELETE FROM public.runtime_object WHERE tenant_id=$1 AND object_key=$2")).WithArgs("tenant-a", reference.ID).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + removed, err := store.CleanupAttachments(context.Background(), "tenant-a", when) + if err != nil || removed != 1 { + t.Fatalf("CleanupAttachments = %d, %v", removed, err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestPostgresAttachmentConcurrentPutKeepsExactReference(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + data := []byte("document") + reference := attachmentReference(data) + when := time.Now().UTC().Add(time.Hour) + mock.ExpectBegin() + mock.ExpectQuery("FROM public.runtime_attachment WHERE tenant_id=\\$1 AND attachment_id=\\$2 FOR UPDATE").WithArgs("tenant-a", reference.ID).WillReturnError(sql.ErrNoRows) + mock.ExpectExec("INSERT INTO public.runtime_object").WithArgs("tenant-a", reference.ID, reference.MIMEType, data, reference.Size, reference.SHA256).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery("SELECT content_type,size,etag FROM public.runtime_object").WithArgs("tenant-a", reference.ID).WillReturnRows(sqlmock.NewRows([]string{"content_type", "size", "etag"}).AddRow(reference.MIMEType, reference.Size, reference.SHA256)) + mock.ExpectExec("INSERT INTO public.runtime_attachment").WithArgs("tenant-a", reference.ID, reference.Kind, reference.MIMEType, reference.Name, reference.Size, reference.SHA256, reference.Provider, reference.ProviderID, sqlmock.AnyArg()).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery("FROM public.runtime_attachment WHERE tenant_id=\\$1 AND attachment_id=\\$2 FOR UPDATE").WithArgs("tenant-a", reference.ID).WillReturnRows(runtimeAttachmentRow(reference, nil, when)) + mock.ExpectCommit() + stored, err := runtimepostgres.New(db).PutAttachment(context.Background(), "tenant-a", attachment.Upload{ID: reference.ID, Kind: reference.Kind, MIMEType: reference.MIMEType, Size: reference.Size}, strings.NewReader(string(data))) + if err != nil || stored != reference { + t.Fatalf("PutAttachment concurrent result = %+v, %v; expectations = %v", stored, err, mock.ExpectationsWereMet()) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestPostgresAttachmentPutValidationAndPrepareErrors(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + store := runtimepostgres.New(db) + data := []byte("document") + reference := attachmentReference(data) + upload := attachmentUpload(reference) + + for _, test := range []struct { + name string + tenant string + upload attachment.Upload + content io.Reader + want error + }{ + {name: "invalid tenant", tenant: "", upload: upload, content: strings.NewReader(string(data)), want: runtimestorage.ErrInvalid}, + {name: "nil content", tenant: "tenant-a", upload: upload, want: runtimestorage.ErrInvalid}, + {name: "invalid upload", tenant: "tenant-a", upload: attachment.Upload{ID: "", Kind: reference.Kind, MIMEType: reference.MIMEType, Size: reference.Size}, content: strings.NewReader(string(data)), want: attachment.ErrInvalid}, + {name: "reader error", tenant: "tenant-a", upload: upload, content: errReader{}, want: runtimestorage.ErrStorage}, + {name: "size mismatch", tenant: "tenant-a", upload: attachment.Upload{ID: reference.ID, Kind: reference.Kind, MIMEType: reference.MIMEType, Size: reference.Size + 1}, content: strings.NewReader(string(data)), want: attachment.ErrInvalid}, + } { + t.Run(test.name, func(t *testing.T) { + _, err := store.PutAttachment(context.Background(), test.tenant, test.upload, test.content) + if !errors.Is(err, test.want) { + t.Fatalf("PutAttachment error = %v, want %v", err, test.want) + } + }) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestPostgresAttachmentPutFailureBoundaries(t *testing.T) { + data := []byte("document") + reference := attachmentReference(data) + expiresAt := time.Now().UTC().Add(time.Hour) + for _, test := range []struct { + name string + prepare func(sqlmock.Sqlmock) + want error + }{ + { + name: "begin failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin().WillReturnError(errors.New("begin failed")) + }, + want: runtimestorage.ErrStorage, + }, + { + name: "existing reference commits", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentLookup(mock, reference.ID, true).WillReturnRows(runtimeAttachmentRow(reference, nil, expiresAt)) + mock.ExpectCommit() + }, + }, + { + name: "existing reference conflict", + prepare: func(mock sqlmock.Sqlmock) { + conflict := reference + conflict.Name = "other.pdf" + mock.ExpectBegin() + expectAttachmentLookup(mock, reference.ID, true).WillReturnRows(runtimeAttachmentRow(conflict, nil, expiresAt)) + mock.ExpectRollback() + }, + want: runtimestorage.ErrConflict, + }, + { + name: "initial lookup failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentLookup(mock, reference.ID, true).WillReturnError(errors.New("lookup failed")) + mock.ExpectRollback() + }, + want: runtimestorage.ErrStorage, + }, + { + name: "object insert failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentLookup(mock, reference.ID, true).WillReturnError(sql.ErrNoRows) + expectObjectInsert(mock, reference, data).WillReturnError(errors.New("object insert failed")) + mock.ExpectRollback() + }, + want: runtimestorage.ErrStorage, + }, + { + name: "object metadata conflict", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentLookup(mock, reference.ID, true).WillReturnError(sql.ErrNoRows) + expectObjectInsert(mock, reference, data).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery("SELECT content_type,size,etag FROM public.runtime_object"). + WithArgs("tenant-a", reference.ID).WillReturnRows(sqlmock.NewRows([]string{"content_type", "size", "etag"}).AddRow("application/json", reference.Size, reference.SHA256)) + mock.ExpectRollback() + }, + want: runtimestorage.ErrConflict, + }, + { + name: "metadata rows affected failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentLookup(mock, reference.ID, true).WillReturnError(sql.ErrNoRows) + expectObjectPersisted(mock, reference, data) + expectMetadataInsert(mock, reference).WillReturnResult(sqlmock.NewErrorResult(errors.New("rows failed"))) + mock.ExpectRollback() + }, + want: runtimestorage.ErrStorage, + }, + { + name: "metadata race conflict", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentLookup(mock, reference.ID, true).WillReturnError(sql.ErrNoRows) + expectObjectPersisted(mock, reference, data) + expectMetadataInsert(mock, reference).WillReturnResult(sqlmock.NewResult(0, 0)) + expectAttachmentLookup(mock, reference.ID, true).WillReturnError(sql.ErrNoRows) + mock.ExpectRollback() + }, + want: runtimestorage.ErrConflict, + }, + { + name: "commit failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentLookup(mock, reference.ID, true).WillReturnError(sql.ErrNoRows) + expectObjectPersisted(mock, reference, data) + expectMetadataInsert(mock, reference).WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit().WillReturnError(errors.New("commit failed")) + }, + want: runtimestorage.ErrStorage, + }, + } { + t.Run(test.name, func(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + test.prepare(mock) + got, err := runtimepostgres.New(db).PutAttachment(context.Background(), "tenant-a", attachmentUpload(reference), strings.NewReader(string(data))) + if test.want == nil { + if err != nil || got != reference { + t.Fatalf("PutAttachment = %+v, %v", got, err) + } + } else if !errors.Is(err, test.want) { + t.Fatalf("PutAttachment error = %v, want %v", err, test.want) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestPostgresAttachmentBindFailureBoundaries(t *testing.T) { + data := []byte("document") + reference := attachmentReference(data) + expiresAt := time.Now().UTC().Add(time.Hour) + for _, test := range []struct { + name string + references []attachment.Reference + prepare func(sqlmock.Sqlmock) + want error + }{ + { + name: "begin failure", + references: []attachment.Reference{reference}, + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin().WillReturnError(errors.New("begin failed")) + }, + want: runtimestorage.ErrStorage, + }, + { + name: "missing event", + references: []attachment.Reference{reference}, + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectMessageEvent(mock, "event-a").WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(false)) + mock.ExpectRollback() + }, + want: runtimestorage.ErrNotFound, + }, + { + name: "event lookup failure", + references: []attachment.Reference{reference}, + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectMessageEvent(mock, "event-a").WillReturnError(errors.New("event lookup failed")) + mock.ExpectRollback() + }, + want: runtimestorage.ErrStorage, + }, + { + name: "invalid reference", + references: []attachment.Reference{{ + ID: "attachment-1", Kind: attachment.KindDocument, MIMEType: "image/png", Size: reference.Size, SHA256: reference.SHA256, + }}, + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectMessageEvent(mock, "event-a").WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + mock.ExpectRollback() + }, + want: attachment.ErrInvalid, + }, + { + name: "missing attachment", + references: []attachment.Reference{reference}, + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectMessageEvent(mock, "event-a").WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + expectAttachmentLookup(mock, reference.ID, true).WillReturnError(sql.ErrNoRows) + mock.ExpectRollback() + }, + want: runtimestorage.ErrNotFound, + }, + { + name: "expired attachment", + references: []attachment.Reference{reference}, + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectMessageEvent(mock, "event-a").WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + expectAttachmentLookup(mock, reference.ID, true).WillReturnRows(runtimeAttachmentRow(reference, nil, time.Now().UTC().Add(-time.Hour))) + mock.ExpectRollback() + }, + want: runtimestorage.ErrConflict, + }, + { + name: "already bound elsewhere", + references: []attachment.Reference{reference}, + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectMessageEvent(mock, "event-a").WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + expectAttachmentLookup(mock, reference.ID, true).WillReturnRows(runtimeAttachmentRow(reference, "event-b", expiresAt)) + mock.ExpectRollback() + }, + want: runtimestorage.ErrConflict, + }, + { + name: "update failure", + references: []attachment.Reference{reference}, + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectMessageEvent(mock, "event-a").WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + expectAttachmentLookup(mock, reference.ID, true).WillReturnRows(runtimeAttachmentRow(reference, nil, expiresAt)) + expectAttachmentBind(mock, reference.ID, "event-a").WillReturnError(errors.New("bind failed")) + mock.ExpectRollback() + }, + want: runtimestorage.ErrStorage, + }, + { + name: "commit failure", + references: []attachment.Reference{reference}, + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectMessageEvent(mock, "event-a").WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + expectAttachmentLookup(mock, reference.ID, true).WillReturnRows(runtimeAttachmentRow(reference, nil, expiresAt)) + expectAttachmentBind(mock, reference.ID, "event-a").WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit().WillReturnError(errors.New("commit failed")) + }, + want: runtimestorage.ErrStorage, + }, + } { + t.Run(test.name, func(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + test.prepare(mock) + err = runtimepostgres.New(db).BindAttachments(context.Background(), "tenant-a", "event-a", test.references) + if !errors.Is(err, test.want) { + t.Fatalf("BindAttachments error = %v, want %v", err, test.want) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestPostgresAttachmentLoadFailureBoundaries(t *testing.T) { + data := []byte("document") + reference := attachmentReference(data) + expiresAt := time.Now().UTC().Add(time.Hour) + for _, test := range []struct { + name string + reference attachment.Reference + tenant string + eventID string + prepare func(sqlmock.Sqlmock) + want error + }{ + {name: "invalid tenant", tenant: "", eventID: "event-a", reference: reference, want: runtimestorage.ErrInvalid}, + {name: "invalid event", tenant: "tenant-a", eventID: "", reference: reference, want: runtimestorage.ErrInvalid}, + {name: "invalid reference", tenant: "tenant-a", eventID: "event-a", reference: attachment.Reference{ID: "bad", Kind: attachment.KindDocument, MIMEType: "application/pdf", Size: 1, SHA256: "bad"}, want: attachment.ErrInvalid}, + { + name: "lookup failure", tenant: "tenant-a", eventID: "event-a", reference: reference, + prepare: func(mock sqlmock.Sqlmock) { + expectAttachmentLookup(mock, reference.ID, false).WillReturnError(errors.New("lookup failed")) + }, + want: runtimestorage.ErrStorage, + }, + { + name: "missing attachment", tenant: "tenant-a", eventID: "event-a", reference: reference, + prepare: func(mock sqlmock.Sqlmock) { + expectAttachmentLookup(mock, reference.ID, false).WillReturnError(sql.ErrNoRows) + }, + want: runtimestorage.ErrNotFound, + }, + { + name: "unbound attachment", tenant: "tenant-a", eventID: "event-a", reference: reference, + prepare: func(mock sqlmock.Sqlmock) { + expectAttachmentLookup(mock, reference.ID, false).WillReturnRows(runtimeAttachmentRow(reference, nil, expiresAt)) + }, + want: runtimestorage.ErrNotFound, + }, + { + name: "wrong event", tenant: "tenant-a", eventID: "event-a", reference: reference, + prepare: func(mock sqlmock.Sqlmock) { + expectAttachmentLookup(mock, reference.ID, false).WillReturnRows(runtimeAttachmentRow(reference, "event-b", expiresAt)) + }, + want: runtimestorage.ErrNotFound, + }, + { + name: "expired", tenant: "tenant-a", eventID: "event-a", reference: reference, + prepare: func(mock sqlmock.Sqlmock) { + expectAttachmentLookup(mock, reference.ID, false).WillReturnRows(runtimeAttachmentRow(reference, "event-a", time.Now().UTC().Add(-time.Hour))) + }, + want: runtimestorage.ErrNotFound, + }, + { + name: "object read failure", tenant: "tenant-a", eventID: "event-a", reference: reference, + prepare: func(mock sqlmock.Sqlmock) { + expectAttachmentLookup(mock, reference.ID, false).WillReturnRows(runtimeAttachmentRow(reference, "event-a", expiresAt)) + mock.ExpectQuery("SELECT content FROM public.runtime_object").WithArgs("tenant-a", reference.ID).WillReturnError(errors.New("object read failed")) + }, + want: runtimestorage.ErrStorage, + }, + { + name: "content digest mismatch", tenant: "tenant-a", eventID: "event-a", reference: reference, + prepare: func(mock sqlmock.Sqlmock) { + expectAttachmentLookup(mock, reference.ID, false).WillReturnRows(runtimeAttachmentRow(reference, "event-a", expiresAt)) + mock.ExpectQuery("SELECT content FROM public.runtime_object").WithArgs("tenant-a", reference.ID).WillReturnRows(sqlmock.NewRows([]string{"content"}).AddRow([]byte("mismatch"))) + }, + want: attachment.ErrInvalid, + }, + } { + t.Run(test.name, func(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + if test.prepare != nil { + test.prepare(mock) + } + _, err = runtimepostgres.New(db).Load(context.Background(), test.tenant, test.eventID, test.reference) + if !errors.Is(err, test.want) { + t.Fatalf("Load error = %v, want %v", err, test.want) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestPostgresAttachmentCleanupFailureBoundaries(t *testing.T) { + now := time.Now().UTC() + for _, test := range []struct { + name string + prepare func(sqlmock.Sqlmock) + want error + }{ + { + name: "begin failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin().WillReturnError(errors.New("begin failed")) + }, + want: runtimestorage.ErrStorage, + }, + { + name: "delete query failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentCleanup(mock, now).WillReturnError(errors.New("delete failed")) + mock.ExpectRollback() + }, + want: runtimestorage.ErrStorage, + }, + { + name: "scan failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentCleanup(mock, now).WillReturnRows(sqlmock.NewRows([]string{"attachment_id"}).AddRow(nil)) + mock.ExpectRollback() + }, + want: runtimestorage.ErrStorage, + }, + { + name: "rows failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentCleanup(mock, now).WillReturnRows(sqlmock.NewRows([]string{"attachment_id"}).AddRow("attachment-1").RowError(0, errors.New("rows failed"))) + mock.ExpectRollback() + }, + want: runtimestorage.ErrStorage, + }, + { + name: "object delete failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentCleanup(mock, now).WillReturnRows(sqlmock.NewRows([]string{"attachment_id"}).AddRow("attachment-1")) + mock.ExpectExec("DELETE FROM public.runtime_object").WithArgs("tenant-a", "attachment-1").WillReturnError(errors.New("object delete failed")) + mock.ExpectRollback() + }, + want: runtimestorage.ErrStorage, + }, + { + name: "commit failure", + prepare: func(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + expectAttachmentCleanup(mock, now).WillReturnRows(sqlmock.NewRows([]string{"attachment_id"}).AddRow("attachment-1")) + mock.ExpectExec("DELETE FROM public.runtime_object").WithArgs("tenant-a", "attachment-1").WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit().WillReturnError(errors.New("commit failed")) + }, + want: runtimestorage.ErrStorage, + }, + } { + t.Run(test.name, func(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + test.prepare(mock) + _, err = runtimepostgres.New(db).CleanupAttachments(context.Background(), "tenant-a", now) + if !errors.Is(err, test.want) { + t.Fatalf("CleanupAttachments error = %v, want %v", err, test.want) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestPostgresAttachmentMethodsFailClosed(t *testing.T) { + var nilStore *runtimepostgres.Store + data := []byte("document") + reference := attachmentReference(data) + if _, err := nilStore.PutAttachment(context.Background(), "tenant-a", attachment.Upload{ID: reference.ID, Kind: reference.Kind, MIMEType: reference.MIMEType, Size: reference.Size}, strings.NewReader(string(data))); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("nil PutAttachment = %v", err) + } + if _, err := nilStore.Load(context.Background(), "tenant-a", "event-a", reference); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("nil Load = %v", err) + } + if err := nilStore.BindAttachments(context.Background(), "tenant-a", "event-a", []attachment.Reference{reference}); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("nil BindAttachments = %v", err) + } + if _, err := nilStore.CleanupAttachments(context.Background(), "tenant-a", time.Now()); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("nil CleanupAttachments = %v", err) + } +} + +func attachmentReference(data []byte) attachment.Reference { + digest := sha256.Sum256(data) + return attachment.Reference{ID: "attachment-1", Kind: attachment.KindDocument, MIMEType: "application/pdf", Size: int64(len(data)), SHA256: hex.EncodeToString(digest[:])} +} + +func attachmentUpload(reference attachment.Reference) attachment.Upload { + return attachment.Upload{ID: reference.ID, Kind: reference.Kind, MIMEType: reference.MIMEType, Name: reference.Name, Size: reference.Size, Provider: reference.Provider, ProviderID: reference.ProviderID} +} + +func runtimeAttachmentRow(reference attachment.Reference, eventID any, expiresAt time.Time) *sqlmock.Rows { + return sqlmock.NewRows(runtimeAttachmentColumns).AddRow("tenant-a", reference.ID, string(reference.Kind), reference.MIMEType, reference.Name, reference.Size, reference.SHA256, reference.Provider, reference.ProviderID, eventID, expiresAt) +} + +func expectAttachmentLookup(mock sqlmock.Sqlmock, id string, lock bool) *sqlmock.ExpectedQuery { + statement := "FROM public.runtime_attachment WHERE tenant_id=\\$1 AND attachment_id=\\$2" + if lock { + statement += " FOR UPDATE" + } + return mock.ExpectQuery(statement).WithArgs("tenant-a", id) +} + +func expectObjectInsert(mock sqlmock.Sqlmock, reference attachment.Reference, data []byte) *sqlmock.ExpectedExec { + return mock.ExpectExec("INSERT INTO public.runtime_object").WithArgs("tenant-a", reference.ID, reference.MIMEType, data, reference.Size, reference.SHA256) +} + +func expectObjectPersisted(mock sqlmock.Sqlmock, reference attachment.Reference, data []byte) { + expectObjectInsert(mock, reference, data).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery("SELECT content_type,size,etag FROM public.runtime_object"). + WithArgs("tenant-a", reference.ID). + WillReturnRows(sqlmock.NewRows([]string{"content_type", "size", "etag"}).AddRow(reference.MIMEType, reference.Size, reference.SHA256)) +} + +func expectMetadataInsert(mock sqlmock.Sqlmock, reference attachment.Reference) *sqlmock.ExpectedExec { + return mock.ExpectExec("INSERT INTO public.runtime_attachment"). + WithArgs("tenant-a", reference.ID, reference.Kind, reference.MIMEType, reference.Name, reference.Size, reference.SHA256, reference.Provider, reference.ProviderID, sqlmock.AnyArg()) +} + +func expectMessageEvent(mock sqlmock.Sqlmock, eventID string) *sqlmock.ExpectedQuery { + return mock.ExpectQuery("SELECT EXISTS").WithArgs("tenant-a", eventID) +} + +func expectAttachmentBind(mock sqlmock.Sqlmock, attachmentID, eventID string) *sqlmock.ExpectedExec { + return mock.ExpectExec("UPDATE public.runtime_attachment SET event_id").WithArgs("tenant-a", attachmentID, eventID) +} + +func expectAttachmentCleanup(mock sqlmock.Sqlmock, before time.Time) *sqlmock.ExpectedQuery { + return mock.ExpectQuery("DELETE FROM public.runtime_attachment").WithArgs("tenant-a", before) +} + +type errReader struct{} + +func (errReader) Read([]byte) (int, error) { + return 0, errors.New("read failed") +} diff --git a/trpcservice/runtime/storage/postgres/repository.go b/trpcservice/runtime/storage/postgres/repository.go index 91022a4a..abe1ff60 100644 --- a/trpcservice/runtime/storage/postgres/repository.go +++ b/trpcservice/runtime/storage/postgres/repository.go @@ -17,7 +17,8 @@ import ( type Store struct{ db *sql.DB } const eventColumns = "tenant_id,event_id,session_id,binding_id,external_message_id,idempotency_key,event_seq,status,fencing_token,lease_owner,lease_expires_at,reply_id,segment_count,reply_conversation_kind,reply_receiver_id,reply_thread_id,created_at,updated_at" -const replyColumns = "tenant_id,reply_id,event_id,segment_index,segment_count,payload,reply_binding_id,reply_conversation_kind,reply_receiver_id,reply_thread_id,status,attempts,fencing_token,lease_owner,lease_expires_at,provider_message_id,last_error_class,created_at,updated_at" +const replyColumns = "tenant_id,reply_id,event_id,segment_index,segment_count,payload,reply_kind,attachment_id,attachment_kind,attachment_mime_type,attachment_name,attachment_size,attachment_sha256,attachment_provider,attachment_provider_id,fallback,reply_binding_id,reply_conversation_kind,reply_receiver_id,reply_thread_id,status,attempts,fencing_token,lease_owner,lease_expires_at,provider_message_id,last_error_class,created_at,updated_at" +const insertReplyStatement = "INSERT INTO public.reply_outbox (tenant_id,reply_id,event_id,segment_index,segment_count,payload,reply_kind,attachment_id,attachment_kind,attachment_mime_type,attachment_name,attachment_size,attachment_sha256,attachment_provider,attachment_provider_id,fallback,reply_binding_id,reply_conversation_kind,reply_receiver_id,reply_thread_id,status) SELECT $1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,'pending' WHERE EXISTS (SELECT 1 FROM public.message_event WHERE tenant_id=$1 AND event_id=$3 AND ((reply_conversation_kind='' AND reply_receiver_id='' AND reply_thread_id='' AND $17='' AND $18='' AND $19='' AND $20='') OR (binding_id=$17 AND reply_conversation_kind=$18 AND reply_receiver_id=$19 AND reply_thread_id=$20))) ON CONFLICT (tenant_id,reply_id,segment_index) DO UPDATE SET updated_at=public.reply_outbox.updated_at WHERE public.reply_outbox.event_id=EXCLUDED.event_id AND public.reply_outbox.segment_count=EXCLUDED.segment_count AND public.reply_outbox.payload=EXCLUDED.payload AND public.reply_outbox.reply_kind=EXCLUDED.reply_kind AND public.reply_outbox.attachment_id=EXCLUDED.attachment_id AND public.reply_outbox.attachment_kind=EXCLUDED.attachment_kind AND public.reply_outbox.attachment_mime_type=EXCLUDED.attachment_mime_type AND public.reply_outbox.attachment_name=EXCLUDED.attachment_name AND public.reply_outbox.attachment_size=EXCLUDED.attachment_size AND public.reply_outbox.attachment_sha256=EXCLUDED.attachment_sha256 AND public.reply_outbox.attachment_provider=EXCLUDED.attachment_provider AND public.reply_outbox.attachment_provider_id=EXCLUDED.attachment_provider_id AND public.reply_outbox.fallback=EXCLUDED.fallback AND public.reply_outbox.reply_binding_id=EXCLUDED.reply_binding_id AND public.reply_outbox.reply_conversation_kind=EXCLUDED.reply_conversation_kind AND public.reply_outbox.reply_receiver_id=EXCLUDED.reply_receiver_id AND public.reply_outbox.reply_thread_id=EXCLUDED.reply_thread_id" // New creates a PostgreSQL runtime store over db. func New(db *sql.DB) *Store { return &Store{db: db} } @@ -326,6 +327,11 @@ func (s *Store) EnqueueReply(ctx context.Context, value runtimestorage.ReplyOutb if err := check(ctx); err != nil { return runtimestorage.ReplyOutbox{}, err } + var normalizeErr error + value, normalizeErr = runtimestorage.NormalizeReplyOutbox(value) + if normalizeErr != nil { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } if runtimestorage.ValidateTenant(value.TenantID) != nil || value.ReplyID == "" || value.EventID == "" || value.SegmentIndex < 0 || value.SegmentCount <= value.SegmentIndex || runtimestorage.ValidateReplyTarget(value.ReplyTarget) != nil { return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid } @@ -345,7 +351,7 @@ func (s *Store) EnqueueReply(ctx context.Context, value runtimestorage.ReplyOutb } } var result runtimestorage.ReplyOutbox - err := s.db.QueryRowContext(ctx, "INSERT INTO public.reply_outbox (tenant_id,reply_id,event_id,segment_index,segment_count,payload,reply_binding_id,reply_conversation_kind,reply_receiver_id,reply_thread_id,status) SELECT $1,$2,$3,$4,$5,$6,$7,$8,$9,$10,'pending' WHERE EXISTS (SELECT 1 FROM public.message_event WHERE tenant_id=$1 AND event_id=$3 AND ((reply_conversation_kind='' AND reply_receiver_id='' AND reply_thread_id='' AND $7='' AND $8='' AND $9='' AND $10='') OR (binding_id=$7 AND reply_conversation_kind=$8 AND reply_receiver_id=$9 AND reply_thread_id=$10))) ON CONFLICT (tenant_id,reply_id,segment_index) DO UPDATE SET updated_at=public.reply_outbox.updated_at WHERE public.reply_outbox.event_id=EXCLUDED.event_id AND public.reply_outbox.segment_count=EXCLUDED.segment_count AND public.reply_outbox.payload=EXCLUDED.payload AND public.reply_outbox.reply_binding_id=EXCLUDED.reply_binding_id AND public.reply_outbox.reply_conversation_kind=EXCLUDED.reply_conversation_kind AND public.reply_outbox.reply_receiver_id=EXCLUDED.reply_receiver_id AND public.reply_outbox.reply_thread_id=EXCLUDED.reply_thread_id RETURNING "+replyColumns, value.TenantID, value.ReplyID, value.EventID, value.SegmentIndex, value.SegmentCount, value.Payload, value.ReplyTarget.BindingID, value.ReplyTarget.ConversationKind, value.ReplyTarget.ReceiverID, value.ReplyTarget.ThreadID).Scan(replyArgs(&result)...) + err := s.db.QueryRowContext(ctx, insertReplyStatement+" RETURNING "+replyColumns, replyInsertArgs(value)...).Scan(replyArgs(&result)...) if err != nil { if errors.Is(err, sql.ErrNoRows) { if _, lookupErr := s.GetMessage(ctx, value.TenantID, value.EventID); lookupErr != nil { @@ -365,6 +371,11 @@ func (s *Store) EnqueueReplies(ctx context.Context, values []runtimestorage.Repl if err := check(ctx); err != nil { return nil, err } + var err error + values, err = normalizeReplyBatch(values) + if err != nil { + return nil, err + } if err := validateReplyBatch(values); err != nil { return nil, err } @@ -403,6 +414,11 @@ func (s *Store) EnqueueRepliesWithCorrelation(ctx context.Context, correlation r return nil, runtimestorage.ErrInvalid } correlation.TraceParent = observability.NormalizeTraceParent(correlation.TraceParent) + var normalizeErr error + values, normalizeErr = normalizeReplyBatch(values) + if normalizeErr != nil { + return nil, normalizeErr + } first, err := validateReplyBatchForCorrelation(correlation, values) if err != nil { return nil, err @@ -474,11 +490,23 @@ func validateReplyBatch(values []runtimestorage.ReplyOutbox) error { return nil } +func normalizeReplyBatch(values []runtimestorage.ReplyOutbox) ([]runtimestorage.ReplyOutbox, error) { + normalized := make([]runtimestorage.ReplyOutbox, 0, len(values)) + for _, value := range values { + reply, err := runtimestorage.NormalizeReplyOutbox(value) + if err != nil { + return nil, runtimestorage.ErrInvalid + } + normalized = append(normalized, reply) + } + return normalized, nil +} + func (s *Store) insertReplySegments(ctx context.Context, tx *sql.Tx, values []runtimestorage.ReplyOutbox) ([]runtimestorage.ReplyOutbox, error) { result := make([]runtimestorage.ReplyOutbox, 0, len(values)) for _, value := range values { var row runtimestorage.ReplyOutbox - err := tx.QueryRowContext(ctx, "INSERT INTO public.reply_outbox (tenant_id,reply_id,event_id,segment_index,segment_count,payload,reply_binding_id,reply_conversation_kind,reply_receiver_id,reply_thread_id,status) SELECT $1,$2,$3,$4,$5,$6,$7,$8,$9,$10,'pending' WHERE EXISTS (SELECT 1 FROM public.message_event WHERE tenant_id=$1 AND event_id=$3 AND ((reply_conversation_kind='' AND reply_receiver_id='' AND reply_thread_id='' AND $7='' AND $8='' AND $9='' AND $10='') OR (binding_id=$7 AND reply_conversation_kind=$8 AND reply_receiver_id=$9 AND reply_thread_id=$10))) ON CONFLICT (tenant_id,reply_id,segment_index) DO UPDATE SET updated_at=public.reply_outbox.updated_at WHERE public.reply_outbox.event_id=EXCLUDED.event_id AND public.reply_outbox.segment_count=EXCLUDED.segment_count AND public.reply_outbox.payload=EXCLUDED.payload AND public.reply_outbox.reply_binding_id=EXCLUDED.reply_binding_id AND public.reply_outbox.reply_conversation_kind=EXCLUDED.reply_conversation_kind AND public.reply_outbox.reply_receiver_id=EXCLUDED.reply_receiver_id AND public.reply_outbox.reply_thread_id=EXCLUDED.reply_thread_id RETURNING "+replyColumns, value.TenantID, value.ReplyID, value.EventID, value.SegmentIndex, value.SegmentCount, value.Payload, value.ReplyTarget.BindingID, value.ReplyTarget.ConversationKind, value.ReplyTarget.ReceiverID, value.ReplyTarget.ThreadID).Scan(replyArgs(&row)...) + err := tx.QueryRowContext(ctx, insertReplyStatement+" RETURNING "+replyColumns, replyInsertArgs(value)...).Scan(replyArgs(&row)...) if err != nil { if errors.Is(err, sql.ErrNoRows) { if _, lookupErr := lookupMessage(ctx, tx, value.TenantID, value.EventID); lookupErr != nil { @@ -605,8 +633,23 @@ func check(ctx context.Context) error { func eventArgs(value *runtimestorage.MessageEvent) []any { return []any{&value.TenantID, &value.EventID, &value.SessionID, &value.BindingID, &value.ExternalMessageID, &value.IdempotencyKey, &value.EventSeq, &value.Status, &value.FencingToken, &value.LeaseOwner, &value.LeaseExpiresAt, &value.ReplyID, &value.SegmentCount, &value.ReplyTarget.ConversationKind, &value.ReplyTarget.ReceiverID, &value.ReplyTarget.ThreadID, &value.CreatedAt, &value.UpdatedAt} } +func replyInsertArgs(value runtimestorage.ReplyOutbox) []any { + return []any{ + value.TenantID, value.ReplyID, value.EventID, value.SegmentIndex, value.SegmentCount, value.Payload, + value.Kind, value.Attachment.ID, value.Attachment.Kind, value.Attachment.MIMEType, value.Attachment.Name, + value.Attachment.Size, value.Attachment.SHA256, value.Attachment.Provider, value.Attachment.ProviderID, value.Fallback, + value.ReplyTarget.BindingID, value.ReplyTarget.ConversationKind, value.ReplyTarget.ReceiverID, value.ReplyTarget.ThreadID, + } +} func replyArgs(value *runtimestorage.ReplyOutbox) []any { - return []any{&value.TenantID, &value.ReplyID, &value.EventID, &value.SegmentIndex, &value.SegmentCount, &value.Payload, &value.ReplyTarget.BindingID, &value.ReplyTarget.ConversationKind, &value.ReplyTarget.ReceiverID, &value.ReplyTarget.ThreadID, &value.Status, &value.Attempts, &value.FencingToken, &value.LeaseOwner, &value.LeaseExpiresAt, &value.ProviderMessageID, &value.LastErrorClass, &value.CreatedAt, &value.UpdatedAt} + return []any{ + &value.TenantID, &value.ReplyID, &value.EventID, &value.SegmentIndex, &value.SegmentCount, &value.Payload, + &value.Kind, &value.Attachment.ID, &value.Attachment.Kind, &value.Attachment.MIMEType, &value.Attachment.Name, + &value.Attachment.Size, &value.Attachment.SHA256, &value.Attachment.Provider, &value.Attachment.ProviderID, &value.Fallback, + &value.ReplyTarget.BindingID, &value.ReplyTarget.ConversationKind, &value.ReplyTarget.ReceiverID, &value.ReplyTarget.ThreadID, + &value.Status, &value.Attempts, &value.FencingToken, &value.LeaseOwner, &value.LeaseExpiresAt, + &value.ProviderMessageID, &value.LastErrorClass, &value.CreatedAt, &value.UpdatedAt, + } } func cloneSession(value runtimestorage.Session) runtimestorage.Session { if value.State != nil { diff --git a/trpcservice/runtime/storage/postgres/repository_test.go b/trpcservice/runtime/storage/postgres/repository_test.go index 617ab885..ecebafc7 100644 --- a/trpcservice/runtime/storage/postgres/repository_test.go +++ b/trpcservice/runtime/storage/postgres/repository_test.go @@ -2,20 +2,23 @@ package postgres_test import ( "context" + "crypto/sha256" "database/sql" "database/sql/driver" + "encoding/hex" "errors" "testing" "time" "github.com/DATA-DOG/go-sqlmock" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" runtimepostgres "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/postgres" "github.com/jackc/pgx/v5/pgconn" ) var eventColumns = []string{"tenant_id", "event_id", "session_id", "binding_id", "external_message_id", "idempotency_key", "event_seq", "status", "fencing_token", "lease_owner", "lease_expires_at", "reply_id", "segment_count", "reply_conversation_kind", "reply_receiver_id", "reply_thread_id", "created_at", "updated_at"} -var replyColumns = []string{"tenant_id", "reply_id", "event_id", "segment_index", "segment_count", "payload", "reply_binding_id", "reply_conversation_kind", "reply_receiver_id", "reply_thread_id", "status", "attempts", "fencing_token", "lease_owner", "lease_expires_at", "provider_message_id", "last_error_class", "created_at", "updated_at"} +var replyColumns = []string{"tenant_id", "reply_id", "event_id", "segment_index", "segment_count", "payload", "reply_kind", "attachment_id", "attachment_kind", "attachment_mime_type", "attachment_name", "attachment_size", "attachment_sha256", "attachment_provider", "attachment_provider_id", "fallback", "reply_binding_id", "reply_conversation_kind", "reply_receiver_id", "reply_thread_id", "status", "attempts", "fencing_token", "lease_owner", "lease_expires_at", "provider_message_id", "last_error_class", "created_at", "updated_at"} var historyColumns = []string{"tenant_id", "session_id", "event_id", "payload", "history_seq", "created_at"} func eventRow(when time.Time) *sqlmock.Rows { @@ -23,7 +26,25 @@ func eventRow(when time.Time) *sqlmock.Rows { } func replyRow(when time.Time) *sqlmock.Rows { - return sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply-1", "event-1", 0, 1, "payload", "", "", "", "", "pending", 0, int64(0), "", nil, "", "", when, when) + return sqlmock.NewRows(replyColumns).AddRow(replyValues("reply-1", "event-1", 0, 1, "payload", "pending", 0, int64(0), "", nil, "", "", when)...) +} + +func replyValues(replyID, eventID string, segmentIndex, segmentCount int, payload, status string, attempts, fencingToken any, leaseOwner string, leaseExpiresAt any, providerID, errorClass string, when time.Time) []driver.Value { + return []driver.Value{"tenant-a", replyID, eventID, segmentIndex, segmentCount, payload, "text", "", "", "", "", int64(0), "", "", "", "", "", "", "", "", status, attempts, fencingToken, leaseOwner, leaseExpiresAt, providerID, errorClass, when, when} +} + +func replyInsertArgs(replyID, eventID string, segmentIndex, segmentCount int, payload string) []driver.Value { + return []driver.Value{"tenant-a", replyID, eventID, segmentIndex, segmentCount, payload, "text", "", "", "", "", int64(0), "", "", "", "", "", "", "", ""} +} + +func mediaReplyReference(t *testing.T, kind attachment.Kind, contentType string, data []byte) attachment.Reference { + t.Helper() + digest := sha256.Sum256(data) + reference := attachment.Reference{ID: "attachment-media", Kind: kind, MIMEType: contentType, Name: "chart.png", Size: int64(len(data)), SHA256: hex.EncodeToString(digest[:]), Provider: "telegram", ProviderID: "file-id"} + if _, err := reference.Normalize(); err != nil { + t.Fatalf("test attachment = %v", err) + } + return reference } func TestGetSessionUsesExplicitTenantPredicateAndDefensiveState(t *testing.T) { @@ -158,7 +179,7 @@ func TestRuntimeStoreCoversMessageAndReplyLifecycle(t *testing.T) { if _, err := store.GetMessage(context.Background(), "tenant-a", "event-1"); err != nil { t.Fatal(err) } - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply-1", "event-1", 0, 1, "payload", "", "", "", "").WillReturnRows(replyRow(when)) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply-1", "event-1", 0, 1, "payload")...).WillReturnRows(replyRow(when)) if _, err := store.EnqueueReply(context.Background(), runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply-1", EventID: "event-1", SegmentIndex: 0, SegmentCount: 1, Payload: "payload"}); err != nil { t.Fatal(err) } @@ -166,11 +187,11 @@ func TestRuntimeStoreCoversMessageAndReplyLifecycle(t *testing.T) { if _, err := store.GetReply(context.Background(), "tenant-a", "reply-1", 0); err != nil { t.Fatal(err) } - mock.ExpectQuery("UPDATE public.reply_outbox SET status='sending'").WithArgs("tenant-a", "reply-1", 0, "worker-a", int64(3)).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply-1", "event-1", 0, 1, "payload", "", "", "", "", "sending", 1, int64(1), "worker-a", when.Add(time.Minute), "", "", when, when)) + mock.ExpectQuery("UPDATE public.reply_outbox SET status='sending'").WithArgs("tenant-a", "reply-1", 0, "worker-a", int64(3)).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow(replyValues("reply-1", "event-1", 0, 1, "payload", "sending", 1, int64(1), "worker-a", when.Add(time.Minute), "", "", when)...)) if _, err := store.ClaimReply(context.Background(), "tenant-a", "reply-1", 0, "worker-a", 3*time.Second); err != nil { t.Fatal(err) } - mock.ExpectQuery("UPDATE public.reply_outbox SET status=\\$5").WithArgs("tenant-a", "reply-1", 0, "sending", "sent", "worker-a", int64(0), "provider-1", "", int64(1)).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply-1", "event-1", 0, 1, "payload", "", "", "", "", "sent", 2, int64(2), "worker-a", nil, "provider-1", "", when, when)) + mock.ExpectQuery("UPDATE public.reply_outbox SET status=\\$5").WithArgs("tenant-a", "reply-1", 0, "sending", "sent", "worker-a", int64(0), "provider-1", "", int64(1)).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow(replyValues("reply-1", "event-1", 0, 1, "payload", "sent", 2, int64(2), "worker-a", nil, "provider-1", "", when)...)) if _, err := store.TransitionReply(context.Background(), runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply-1", SegmentIndex: 0, From: "sending", To: "sent", Owner: "worker-a", FencingToken: 1, ProviderID: "provider-1"}); err != nil { t.Fatal(err) } @@ -207,7 +228,7 @@ func TestEnqueueReplyRejectsLegacyTargetForRoutedEvent(t *testing.T) { } defer func() { _ = db.Close() }() when := time.Now().UTC() - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply-1", "event-1", 0, 1, "payload", "", "", "", "").WillReturnError(sql.ErrNoRows) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply-1", "event-1", 0, 1, "payload")...).WillReturnError(sql.ErrNoRows) mock.ExpectQuery("SELECT tenant_id,event_id,session_id,binding_id,external_message_id").WithArgs("tenant-a", "event-1").WillReturnRows(sqlmock.NewRows(eventColumns).AddRow("tenant-a", "event-1", "session-1", "binding-1", "external-1", "", int64(2), "completed", int64(1), "", nil, "reply-1", 1, "direct", "user-1", "", when, when)) _, err = runtimepostgres.New(db).EnqueueReply(context.Background(), runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply-1", EventID: "event-1", SegmentCount: 1, Payload: "payload"}) if !errors.Is(err, runtimestorage.ErrConflict) { @@ -218,6 +239,39 @@ func TestEnqueueReplyRejectsLegacyTargetForRoutedEvent(t *testing.T) { } } +func TestEnqueueReplyPersistsMediaReplyContract(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + defer func() { _ = db.Close() }() + when := time.Now().UTC() + reference := mediaReplyReference(t, attachment.KindImage, "image/png", []byte("png")) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs( + "tenant-a", "reply-media", "event-media", 0, 1, "caption", + runtimestorage.ReplyKindImage, reference.ID, reference.Kind, reference.MIMEType, reference.Name, + reference.Size, reference.SHA256, reference.Provider, reference.ProviderID, "[image attachment: chart.png]", + "", "", "", "", + ).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow( + "tenant-a", "reply-media", "event-media", 0, 1, "caption", + "image", reference.ID, reference.Kind, reference.MIMEType, reference.Name, reference.Size, reference.SHA256, reference.Provider, reference.ProviderID, "[image attachment: chart.png]", + "", "", "", "", "pending", 0, int64(0), "", nil, "", "", when, when, + )) + got, err := runtimepostgres.New(db).EnqueueReply(context.Background(), runtimestorage.ReplyOutbox{ + TenantID: "tenant-a", ReplyID: "reply-media", EventID: "event-media", SegmentIndex: 0, SegmentCount: 1, + Kind: runtimestorage.ReplyKindImage, Payload: "caption", Attachment: reference, Fallback: "[image attachment: chart.png]", + }) + if err != nil { + t.Fatal(err) + } + if got.Kind != runtimestorage.ReplyKindImage || got.Attachment != reference || got.Fallback != "[image attachment: chart.png]" { + t.Fatalf("media reply = %+v", got) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + func TestRuntimeStoreCoversEventHistoryAndMessageLifecycle(t *testing.T) { db, mock, err := sqlmock.New() if err != nil { @@ -320,7 +374,7 @@ func TestRuntimeStoreListReplyCandidates(t *testing.T) { defer func() { _ = db.Close() }() store := runtimepostgres.New(db) when := time.Now().UTC() - mock.ExpectQuery("SELECT tenant_id,reply_id,event_id,segment_index").WithArgs("tenant-a").WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply-1", "event-1", 0, 1, "payload", "", "", "", "", "pending", 0, int64(0), "", nil, "", "", when, when)) + mock.ExpectQuery("SELECT tenant_id,reply_id,event_id,segment_index").WithArgs("tenant-a").WillReturnRows(sqlmock.NewRows(replyColumns).AddRow(replyValues("reply-1", "event-1", 0, 1, "payload", "pending", 0, int64(0), "", nil, "", "", when)...)) values, err := store.ListReplyCandidates(context.Background(), "tenant-a") if err != nil || len(values) != 1 || values[0].ReplyID != "reply-1" { t.Fatalf("reply candidates = %+v err=%v", values, err) @@ -515,7 +569,7 @@ func TestRuntimeStorePostgresErrorBranches(t *testing.T) { if _, err := store.GetMessage(ctx, "tenant-a", "event-error"); !errors.Is(err, runtimestorage.ErrStorage) { t.Fatalf("get message error = %v", err) } - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply-error", "event", 0, 1, "", "", "", "", "").WillReturnError(errors.New("enqueue failed")) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply-error", "event", 0, 1, "")...).WillReturnError(errors.New("enqueue failed")) if _, err := store.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply-error", EventID: "event", SegmentIndex: 0, SegmentCount: 1}); !errors.Is(err, runtimestorage.ErrStorage) { t.Fatalf("enqueue error = %v", err) } @@ -569,7 +623,7 @@ func TestRuntimeStoreTransitionValidationAndLease(t *testing.T) { t.Fatalf("illegal transition = %v", err) } when := time.Now().UTC() - mock.ExpectQuery("UPDATE public.reply_outbox SET status=\\$5").WithArgs("tenant-a", "reply-lease", 0, "pending", "sending", "worker", int64(2), "", "", int64(0)).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply-lease", "event", 0, 1, "payload", "", "", "", "", "sending", 1, int64(1), "worker", when.Add(time.Minute), "", "", when, when)) + mock.ExpectQuery("UPDATE public.reply_outbox SET status=\\$5").WithArgs("tenant-a", "reply-lease", 0, "pending", "sending", "worker", int64(2), "", "", int64(0)).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow(replyValues("reply-lease", "event", 0, 1, "payload", "sending", 1, int64(1), "worker", when.Add(time.Minute), "", "", when)...)) if _, err := store.TransitionReply(context.Background(), runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply-lease", SegmentIndex: 0, From: "pending", To: "sending", Owner: "worker", LeaseDuration: 2 * time.Second}); err != nil { t.Fatal(err) } @@ -587,8 +641,8 @@ func TestRuntimeStoreEnqueueRepliesRollsBackPartialMaterialization(t *testing.T) store := runtimepostgres.New(db) when := time.Now().UTC() mock.ExpectBegin() - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply-batch", "event", 0, 2, "first", "", "", "", "").WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply-batch", "event", 0, 2, "first", "", "", "", "", "pending", 0, int64(0), "", nil, "", "", when, when)) - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply-batch", "event", 1, 2, "second", "", "", "", "").WillReturnError(errors.New("second insert failed")) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply-batch", "event", 0, 2, "first")...).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow(replyValues("reply-batch", "event", 0, 2, "first", "pending", 0, int64(0), "", nil, "", "", when)...)) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply-batch", "event", 1, 2, "second")...).WillReturnError(errors.New("second insert failed")) mock.ExpectRollback() _, err = store.EnqueueReplies(context.Background(), []runtimestorage.ReplyOutbox{ {TenantID: "tenant-a", ReplyID: "reply-batch", EventID: "event", SegmentIndex: 0, SegmentCount: 2, Payload: "first"}, @@ -611,7 +665,7 @@ func TestRuntimeStoreEnqueueRepliesMapsMissingEvent(t *testing.T) { db.SetMaxOpenConns(1) store := runtimepostgres.New(db) mock.ExpectBegin() - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply-missing-event", "event-missing", 0, 1, "payload", "", "", "", "").WillReturnError(sql.ErrNoRows) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply-missing-event", "event-missing", 0, 1, "payload")...).WillReturnError(sql.ErrNoRows) mock.ExpectQuery("SELECT tenant_id,event_id,session_id,binding_id,external_message_id").WithArgs("tenant-a", "event-missing").WillReturnError(sql.ErrNoRows) mock.ExpectRollback() _, err = store.EnqueueReplies(context.Background(), []runtimestorage.ReplyOutbox{{ @@ -634,8 +688,8 @@ func TestRuntimeStoreEnqueueRepliesCommitsCompleteBatch(t *testing.T) { store := runtimepostgres.New(db) when := time.Now().UTC() mock.ExpectBegin() - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply-batch", "event", 0, 2, "first", "", "", "", "").WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply-batch", "event", 0, 2, "first", "", "", "", "", "pending", 0, int64(0), "", nil, "", "", when, when)) - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply-batch", "event", 1, 2, "second", "", "", "", "").WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply-batch", "event", 1, 2, "second", "", "", "", "", "pending", 0, int64(0), "", nil, "", "", when, when)) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply-batch", "event", 0, 2, "first")...).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow(replyValues("reply-batch", "event", 0, 2, "first", "pending", 0, int64(0), "", nil, "", "", when)...)) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply-batch", "event", 1, 2, "second")...).WillReturnRows(sqlmock.NewRows(replyColumns).AddRow(replyValues("reply-batch", "event", 1, 2, "second", "pending", 0, int64(0), "", nil, "", "", when)...)) mock.ExpectCommit() rows, err := store.EnqueueReplies(context.Background(), []runtimestorage.ReplyOutbox{ {TenantID: "tenant-a", ReplyID: "reply-batch", EventID: "event", SegmentIndex: 0, SegmentCount: 2, Payload: "first"}, @@ -659,7 +713,7 @@ func TestRuntimeStoreEnqueueRepliesWithCorrelationIsAtomic(t *testing.T) { when := time.Now().UTC() mock.ExpectBegin() mock.ExpectQuery("INSERT INTO public.runtime_reply_correlation").WithArgs("tenant-a", "event", "request", "trace", "").WillReturnRows(sqlmock.NewRows([]string{"tenant_id"}).AddRow("tenant-a")) - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply", "event", 0, 1, "payload", "", "", "", "").WillReturnRows(replyRow(when)) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply", "event", 0, 1, "payload")...).WillReturnRows(replyRow(when)) mock.ExpectCommit() rows, err := store.EnqueueRepliesWithCorrelation(context.Background(), runtimestorage.ReplyCorrelation{TenantID: "tenant-a", EventID: "event", RequestID: "request", TraceID: "trace"}, []runtimestorage.ReplyOutbox{{TenantID: "tenant-a", ReplyID: "reply", EventID: "event", SegmentIndex: 0, SegmentCount: 1, Payload: "payload"}}) if err != nil || len(rows) != 1 { @@ -680,7 +734,7 @@ func TestRuntimeStoreEnqueueRepliesWithCorrelationNormalizesTraceParent(t *testi when := time.Now().UTC() mock.ExpectBegin() mock.ExpectQuery("INSERT INTO public.runtime_reply_correlation").WithArgs("tenant-a", "event", "request", "trace", "").WillReturnRows(sqlmock.NewRows([]string{"tenant_id"}).AddRow("tenant-a")) - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply", "event", 0, 1, "payload", "", "", "", "").WillReturnRows(replyRow(when)) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply", "event", 0, 1, "payload")...).WillReturnRows(replyRow(when)) mock.ExpectCommit() correlation := runtimestorage.ReplyCorrelation{TenantID: "tenant-a", EventID: "event", RequestID: "request", TraceID: "trace", TraceParent: "malformed"} if _, err := store.EnqueueRepliesWithCorrelation(context.Background(), correlation, []runtimestorage.ReplyOutbox{{TenantID: "tenant-a", ReplyID: "reply", EventID: "event", SegmentIndex: 0, SegmentCount: 1, Payload: "payload"}}); err != nil { @@ -738,14 +792,14 @@ func TestRuntimeStoreEnqueueRepliesWithCorrelationFailureBoundaries(t *testing.T {name: "segment failure", setup: func(mock sqlmock.Sqlmock, when time.Time) { mock.ExpectBegin() mock.ExpectQuery("INSERT INTO public.runtime_reply_correlation").WithArgs("tenant-a", "event", "request", "trace", "").WillReturnRows(sqlmock.NewRows([]string{"tenant_id"}).AddRow("tenant-a")) - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply", "event", 0, 1, "payload", "", "", "", "").WillReturnError(errors.New("segment failed")) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply", "event", 0, 1, "payload")...).WillReturnError(errors.New("segment failed")) mock.ExpectRollback() _ = when }}, {name: "commit failure", setup: func(mock sqlmock.Sqlmock, when time.Time) { mock.ExpectBegin() mock.ExpectQuery("INSERT INTO public.runtime_reply_correlation").WithArgs("tenant-a", "event", "request", "trace", "").WillReturnRows(sqlmock.NewRows([]string{"tenant_id"}).AddRow("tenant-a")) - mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs("tenant-a", "reply", "event", 0, 1, "payload", "", "", "", "").WillReturnRows(replyRow(when)) + mock.ExpectQuery("INSERT INTO public.reply_outbox").WithArgs(replyInsertArgs("reply", "event", 0, 1, "payload")...).WillReturnRows(replyRow(when)) mock.ExpectCommit().WillReturnError(errors.New("commit failed")) }}, } { @@ -949,11 +1003,11 @@ func TestRuntimeStoreListReplyCandidatesErrorBranches(t *testing.T) { t.Fatalf("candidate query error = %v", err) } when := time.Now().UTC() - mock.ExpectQuery("SELECT tenant_id,reply_id,event_id,segment_index").WithArgs("tenant-scan-error").WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply", "event", 0, 1, "payload", "", "", "", "", "pending", "bad-attempts", int64(0), "", nil, "", "", when, when)) + mock.ExpectQuery("SELECT tenant_id,reply_id,event_id,segment_index").WithArgs("tenant-scan-error").WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply", "event", 0, 1, "payload", "text", "", "", "", "", int64(0), "", "", "", "", "", "", "", "", "pending", "bad-attempts", int64(0), "", nil, "", "", when, when)) if _, err := store.ListReplyCandidates(context.Background(), "tenant-scan-error"); !errors.Is(err, runtimestorage.ErrStorage) { t.Fatalf("candidate scan error = %v", err) } - mock.ExpectQuery("SELECT tenant_id,reply_id,event_id,segment_index").WithArgs("tenant-rows-error").WillReturnRows(sqlmock.NewRows(replyColumns).AddRow("tenant-a", "reply", "event", 0, 1, "payload", "", "", "", "", "pending", 0, int64(0), "", nil, "", "", when, when).AddRow("tenant-a", "reply-2", "event", 0, 1, "payload", "", "", "", "", "pending", 0, int64(0), "", nil, "", "", when, when).RowError(1, errors.New("candidate rows failed"))) + mock.ExpectQuery("SELECT tenant_id,reply_id,event_id,segment_index").WithArgs("tenant-rows-error").WillReturnRows(sqlmock.NewRows(replyColumns).AddRow(replyValues("reply", "event", 0, 1, "payload", "pending", 0, int64(0), "", nil, "", "", when)...).AddRow(replyValues("reply-2", "event", 0, 1, "payload", "pending", 0, int64(0), "", nil, "", "", when)...).RowError(1, errors.New("candidate rows failed"))) if _, err := store.ListReplyCandidates(context.Background(), "tenant-rows-error"); !errors.Is(err, runtimestorage.ErrStorage) { t.Fatalf("candidate rows error = %v", err) } diff --git a/trpcservice/runtime/storage/storage.go b/trpcservice/runtime/storage/storage.go index decca625..ee5c0bb3 100644 --- a/trpcservice/runtime/storage/storage.go +++ b/trpcservice/runtime/storage/storage.go @@ -8,6 +8,7 @@ import ( "strings" "time" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" pgstorage "github.com/XnLemon/trpc-agent-service/trpcservice/storage/postgres" ) @@ -32,6 +33,23 @@ func ValidateEmbedding(values []float64) bool { const maxReplyTargetIDRunes = 1024 +// ReplyKind identifies the durable representation a channel provider should +// attempt before falling back to text. +type ReplyKind string + +const ( + // ReplyKindText identifies the legacy text reply path. + ReplyKindText ReplyKind = "text" + // ReplyKindImage identifies an image attachment reply. + ReplyKindImage ReplyKind = "image" + // ReplyKindVideo identifies a video attachment reply. + ReplyKindVideo ReplyKind = "video" + // ReplyKindAudio identifies an audio attachment reply. + ReplyKindAudio ReplyKind = "audio" + // ReplyKindDocument identifies a document attachment reply. + ReplyKindDocument ReplyKind = "document" +) + // ReplyTarget is the trusted, durable destination for a channel reply. A zero // target is retained only for rows created before per-message routing existed. type ReplyTarget struct { @@ -166,7 +184,10 @@ type ReplyOutbox struct { EventID string SegmentIndex int SegmentCount int + Kind ReplyKind Payload string + Attachment attachment.Reference + Fallback string ReplyTarget ReplyTarget Status string Attempts int @@ -291,6 +312,56 @@ func validReplyTargetID(value string) bool { return true } +// NormalizeReplyOutbox validates a reply's protocol-neutral media contract and +// returns a canonical copy. A zero Kind preserves the historical text-only path. +func NormalizeReplyOutbox(value ReplyOutbox) (ReplyOutbox, error) { + value.Kind = normalizedReplyKind(value.Kind) + switch value.Kind { + case ReplyKindText: + if value.Attachment != (attachment.Reference{}) || value.Fallback != "" { + return ReplyOutbox{}, ErrInvalid + } + return value, nil + case ReplyKindImage, ReplyKindVideo, ReplyKindAudio, ReplyKindDocument: + reference, err := value.Attachment.Normalize() + if err != nil { + return ReplyOutbox{}, err + } + if replyKindForAttachment(reference.Kind) != value.Kind { + return ReplyOutbox{}, ErrInvalid + } + if !ValidateText(value.Fallback, 4096, true) { + return ReplyOutbox{}, ErrInvalid + } + value.Attachment = reference + return value, nil + default: + return ReplyOutbox{}, ErrInvalid + } +} + +func normalizedReplyKind(kind ReplyKind) ReplyKind { + if kind == "" { + return ReplyKindText + } + return ReplyKind(strings.ToLower(strings.TrimSpace(string(kind)))) +} + +func replyKindForAttachment(kind attachment.Kind) ReplyKind { + switch kind { + case attachment.KindImage: + return ReplyKindImage + case attachment.KindVideo: + return ReplyKindVideo + case attachment.KindAudio: + return ReplyKindAudio + case attachment.KindDocument: + return ReplyKindDocument + default: + return "" + } +} + // ValidateTransition reports whether a reply transition is legal. func ValidateTransition(from, to string) bool { switch from { diff --git a/trpcservice/runtime/storage/storage_test.go b/trpcservice/runtime/storage/storage_test.go index ff672be6..a79cbbc2 100644 --- a/trpcservice/runtime/storage/storage_test.go +++ b/trpcservice/runtime/storage/storage_test.go @@ -1,9 +1,13 @@ package storage_test import ( + "crypto/sha256" + "encoding/hex" "errors" + "strings" "testing" + "github.com/XnLemon/trpc-agent-service/trpcservice/attachment" "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" ) @@ -47,3 +51,60 @@ func TestReplyTargetValidation(t *testing.T) { t.Fatalf("legacy zero target = %v", err) } } + +func TestNormalizeReplyOutboxMediaContract(t *testing.T) { + if got, err := storage.NormalizeReplyOutbox(storage.ReplyOutbox{Payload: "hello"}); err != nil || got.Kind != storage.ReplyKindText { + t.Fatalf("legacy text normalize = %+v, %v", got, err) + } + textWithMedia := storage.ReplyOutbox{Kind: storage.ReplyKindText, Fallback: "fallback", Attachment: mediaReference(t, attachment.KindImage, "image/png", "chart.png", []byte("png"))} + if _, err := storage.NormalizeReplyOutbox(textWithMedia); !errors.Is(err, storage.ErrInvalid) { + t.Fatalf("text with media = %v", err) + } + + for _, test := range []struct { + name string + kind storage.ReplyKind + ref attachment.Reference + }{ + {name: "image", kind: " Image ", ref: mediaReference(t, attachment.KindImage, "IMAGE/PNG", " chart.png ", []byte("png"))}, + {name: "video", kind: storage.ReplyKindVideo, ref: mediaReference(t, attachment.KindVideo, "video/mp4", "clip.mp4", []byte("mp4"))}, + {name: "audio", kind: storage.ReplyKindAudio, ref: mediaReference(t, attachment.KindAudio, "audio/mpeg", "voice.mp3", []byte("mp3"))}, + {name: "document", kind: storage.ReplyKindDocument, ref: mediaReference(t, attachment.KindDocument, "application/pdf", "brief.pdf", []byte("pdf"))}, + } { + t.Run(test.name, func(t *testing.T) { + got, err := storage.NormalizeReplyOutbox(storage.ReplyOutbox{Kind: test.kind, Attachment: test.ref, Fallback: "fallback"}) + if err != nil { + t.Fatalf("NormalizeReplyOutbox = %v", err) + } + if got.Kind != storage.ReplyKind(strings.ToLower(strings.TrimSpace(string(test.kind)))) || got.Attachment.MIMEType != strings.ToLower(strings.TrimSpace(test.ref.MIMEType)) || got.Fallback != "fallback" { + t.Fatalf("normalized media reply = %+v", got) + } + }) + } + + for _, test := range []struct { + name string + value storage.ReplyOutbox + }{ + {name: "invalid kind", value: storage.ReplyOutbox{Kind: "sticker"}}, + {name: "invalid reference", value: storage.ReplyOutbox{Kind: storage.ReplyKindImage, Attachment: attachment.Reference{ID: "bad"}, Fallback: "fallback"}}, + {name: "kind mismatch", value: storage.ReplyOutbox{Kind: storage.ReplyKindImage, Attachment: mediaReference(t, attachment.KindDocument, "application/pdf", "brief.pdf", []byte("pdf")), Fallback: "fallback"}}, + {name: "missing fallback", value: storage.ReplyOutbox{Kind: storage.ReplyKindImage, Attachment: mediaReference(t, attachment.KindImage, "image/png", "chart.png", []byte("png"))}}, + } { + t.Run(test.name, func(t *testing.T) { + if _, err := storage.NormalizeReplyOutbox(test.value); !errors.Is(err, storage.ErrInvalid) && !errors.Is(err, attachment.ErrInvalid) { + t.Fatalf("NormalizeReplyOutbox accepted invalid value: %+v err=%v", test.value, err) + } + }) + } +} + +func mediaReference(t *testing.T, kind attachment.Kind, contentType, name string, data []byte) attachment.Reference { + t.Helper() + digest := sha256.Sum256(data) + reference := attachment.Reference{ID: "attachment-" + string(kind), Kind: kind, MIMEType: contentType, Name: name, Size: int64(len(data)), SHA256: hex.EncodeToString(digest[:])} + if _, err := reference.Normalize(); err != nil { + t.Fatalf("reference = %v", err) + } + return reference +}