diff --git a/README.md b/README.md index bfb665e8..560c2031 100644 --- a/README.md +++ b/README.md @@ -48,7 +48,7 @@ IM / HTTP -> Channel Adapter -> Gateway -> Queue/Outbox -> Agent Worker - PostgreSQL 控制面、migration、显式 `init` 初始化和受认证的 Admin API; - Tenant/App/Revision 的草稿、发布、回滚、灰度候选和乐观锁; - OpenAI 模型 provider,以及不访问外部服务的 deterministic fake provider; -- InMemory 与 PostgreSQL runtime storage、Session/Event、Reply Outbox 和租约恢复; +- InMemory、PostgreSQL 与 tenant-scoped Redis runtime storage;Redis 当前覆盖 Session/Event、Memory、Reply Outbox 和租约恢复; - 普通及流式 HTTP Chat API,企业微信自建应用文本 webhook,Telegram 文本 long polling; - OpenTelemetry trace/metrics、Prometheus 导出路径、审计事件和脱敏错误; - Docker Compose 本地验证、Kubernetes Kustomize base,以及版本 tag 触发的 GHCR 镜像发布。 diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml index a7b32406..794ebf86 100644 --- a/deploy/docker-compose.yml +++ b/deploy/docker-compose.yml @@ -16,6 +16,18 @@ services: retries: 12 start_period: 10s + redis: + image: redis:7-alpine + command: ["redis-server", "--appendonly", "yes"] + volumes: + - redis-data:/data + healthcheck: + test: ["CMD", "redis-cli", "ping"] + interval: 5s + timeout: 3s + retries: 12 + start_period: 5s + service: build: context: .. @@ -25,6 +37,8 @@ services: depends_on: postgres: condition: service_healthy + redis: + condition: service_healthy environment: TRPC_POSTGRES_DSN: "${TRPC_POSTGRES_DSN:-postgres://${POSTGRES_USER:-trpc}:${POSTGRES_PASSWORD:-trpc-local-password}@postgres:5432/${POSTGRES_DB:-trpc_agent}?sslmode=disable}" TRPC_API_TOKEN: ${TRPC_API_TOKEN:-local-api-token} @@ -34,6 +48,11 @@ services: TRPC_ADMIN_TENANTS: ${TRPC_ADMIN_TENANTS:-*} TRPC_MODEL_API_KEY: ${TRPC_MODEL_API_KEY:-local-development-key} TRPC_SESSION_BACKEND: ${TRPC_SESSION_BACKEND:-postgres} + TRPC_REDIS_ADDR: ${TRPC_REDIS_ADDR:-redis:6379} + TRPC_REDIS_PASSWORD: ${TRPC_REDIS_PASSWORD:-} + TRPC_REDIS_DB: ${TRPC_REDIS_DB:-0} + TRPC_REDIS_KEY_PREFIX: ${TRPC_REDIS_KEY_PREFIX:-trpc:runtime:v1} + TRPC_REDIS_SECRET_REF: ${TRPC_REDIS_SECRET_REF:-env/trpc-redis-password} TRPC_MODEL_PROVIDER: ${TRPC_MODEL_PROVIDER:-openai} TRPC_MODEL_NAMES: ${TRPC_MODEL_NAMES:-gpt-4o-mini} TRPC_MODEL_ENDPOINT_HOSTS: ${TRPC_MODEL_ENDPOINT_HOSTS:-api.openai.com} @@ -54,3 +73,4 @@ services: volumes: postgres-data: + redis-data: diff --git a/deploy/example.env b/deploy/example.env index 708fca55..670c6aa6 100644 --- a/deploy/example.env +++ b/deploy/example.env @@ -14,6 +14,18 @@ TRPC_ADMIN_TOKEN=local-admin-token TRPC_ADMIN_TENANTS=* TRPC_MODEL_API_KEY=local-development-key TRPC_SESSION_BACKEND=postgres +# Redis is available in Compose for an explicit TRPC_SESSION_BACKEND=redis +# deployment. The local service has no password; external deployments should +# inject TRPC_REDIS_PASSWORD through a secret manager. +# TRPC_REDIS_ADDR=redis:6379 +# TRPC_REDIS_PASSWORD= +# TRPC_REDIS_DB=0 +# TRPC_REDIS_KEY_PREFIX=trpc:runtime:v1 +# TRPC_REDIS_SECRET_REF=env/trpc-redis-password +# TRPC_REDIS_DIAL_TIMEOUT=500ms +# TRPC_REDIS_READ_TIMEOUT=500ms +# TRPC_REDIS_WRITE_TIMEOUT=500ms +# TRPC_REDIS_POOL_SIZE=10 TRPC_MODEL_PROVIDER=openai TRPC_MODEL_NAMES=gpt-4o-mini TRPC_MODEL_ENDPOINT_HOSTS=api.openai.com diff --git a/docs/docs/deployment.md b/docs/docs/deployment.md index aa771a63..7b61c775 100644 --- a/docs/docs/deployment.md +++ b/docs/docs/deployment.md @@ -58,6 +58,9 @@ demo 仅支持 PostgreSQL、本地开发使用,遇到部分或多租户歧义 - 容器健康检查使用静态 `/app/trpc-healthcheck`,不会依赖 distroless 镜像中的 shell。 - `/healthz` 用于存活和 startup probe;`/readyz` 用于流量接入前的 readiness。 - `TRPC_SERVICE_IMAGE` 可在 CI 或本地覆盖服务镜像名;默认值是 `trpc-agent-service:local`。 +- Compose 同时启动 Redis 7(AOF、`redis:6379`)作为可选运行时后端;只有显式设置 + `TRPC_SESSION_BACKEND=redis` 时服务才会读写它。生产环境应替换为受管 Redis,并通过 Secret + Manager 注入 `TRPC_REDIS_PASSWORD`。 填充后的 `deploy/service.env` 已被 `.dockerignore` 排除,不会进入 Docker build context;该文件 仍可能被 Compose 读取,因此不要提交到 Git,也不要把它作为生产 Secret 管理方案。 @@ -185,7 +188,14 @@ curl --fail http://127.0.0.1:8080/readyz | `TRPC_MODEL_NAMES` | 否,`gpt-4o-mini` | 逗号分隔模型白名单 | | `TRPC_MODEL_ENDPOINT_HOSTS` | 否,`api.openai.com` | 逗号分隔 HTTPS endpoint host 白名单 | | `TRPC_MODEL_SECRET_REF` | 否,`env/trpc-model-api-key` | 运行时 Secret 引用,不是 Secret 值 | -| `TRPC_SESSION_BACKEND` | 必需,显式 `postgres`/`inmemory` | Compose/Kubernetes 示例使用 `postgres`;MySQL 控制面当前应使用 `inmemory` | +| `TRPC_SESSION_BACKEND` | 必需,显式 `postgres`/`redis`/`inmemory` | Compose/Kubernetes 示例使用 `postgres`;Redis 模式只提供 Session/Memory,MySQL 控制面当前应使用 `inmemory` | +| `TRPC_REDIS_ADDR` | Redis 模式必需 | Redis `host:port`;Compose 默认使用 `redis:6379` | +| `TRPC_REDIS_PASSWORD` | 否 | Redis 认证密码,使用 Secret Manager 注入,不进入日志或快照 | +| `TRPC_REDIS_SECRET_REF` | 否,`env/trpc-redis-password` | Redis Backend Profile 的可选 SecretRef | +| `TRPC_REDIS_DB` | 否,`0` | Redis logical database | +| `TRPC_REDIS_KEY_PREFIX` | 否,`trpc:runtime:v1` | tenant-scoped key 前缀 | +| `TRPC_REDIS_DIAL_TIMEOUT` / `TRPC_REDIS_READ_TIMEOUT` / `TRPC_REDIS_WRITE_TIMEOUT` | 否 | Go duration,限制 Redis 客户端 I/O | +| `TRPC_REDIS_POOL_SIZE` | 否 | 大于 `0` 时覆盖连接池大小 | | `TRPC_DEMO_MODE` | 否,`false` | 仅由 `quickstart.sh --demo` 显式启用;要求 `TRPC_MODEL_PROVIDER=fake`,不读取模型凭据 | 模型 API key 只在受信任的 Secret Resolver/Factory 路径中使用,不进入 Execution Plan、缓存、 diff --git a/docs/docs/ops.md b/docs/docs/ops.md index f56fc034..f8ce79a4 100644 --- a/docs/docs/ops.md +++ b/docs/docs/ops.md @@ -1,8 +1,9 @@ # 运维、可观测性与生产风险 > 本页把 [生产架构设计](architecture.md) 转成可执行的发布、监控、恢复和风险检查表。 -> 当前仓库只有控制面领域模型、快照和最小 Runner spine;Gateway、队列、真实 IM/Storage -> Adapter、Dashboard 和告警规则仍是后续平台实现,不应把本页当作已经部署的运行手册。 +> 本页按代码和自动化测试证据标注能力状态,不等同于一份已经完成生产部署的运行手册。 +> 控制面、Runner spine、SQL/InMemory RuntimeStore,以及可选的 Redis Session/Memory +> provider 已落地;真实 IM 验签、Dashboard、生产告警平台和其他外部存储适配仍需单独验收。 ## 运行边界与值班目标 @@ -209,8 +210,19 @@ backpressure。高峰保护使用租户级 token bucket、全局队列上限、 ## 当前实现状态与后续门禁 -本仓库目前可以验证 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 | 已实现 | 控制面模型、快照和策略测试 | +| InMemory RuntimeStore | 已实现 | Tenant-scoped Session/Event/Memory 契约测试 | +| PostgreSQL RuntimeStore | 部分实现 | 迁移、CAS、幂等、Outbox 和可选 live conformance;需要外部 DSN 才能验证重启恢复 | +| Redis RuntimeStore | 已实现(Issue #108) | `TRPC_SESSION_BACKEND=redis`,Redis Session/Memory、WATCH/MULTI CAS、租户隔离和 readiness PING;可选 live reconnect 测试 | +| Redis capability 范围 | 明确限制 | 仅 `session`、`memory`;`summary`、`knowledge`、`artifact`、`audit` 和独立向量库 provider 会被拒绝 | +| Redis/PostgreSQL 迁移、双写、shadow read、自动 cutover | 未实现 | 迁移方案仍需后续工具和演练,不能把切换当作 Redis provider 自带能力 | +| 对象存储(S3/OSS)和生产向量库(Qdrant/Milvus/pgvector) | 未实现 | 当前只有接口/能力边界,未提供真实外部 adapter 或检索闭环 | +| 真实 IM 验签、多媒体、Dashboard 和生产告警平台 | 未实现或部分实现 | WeCom/Telegram 的已交付范围以各自 adapter 文档为准;本页不宣称生产运营集成 | + +后续每落地一个 Adapter 或运维组件,都必须补充:双租户隔离测试、重复/乱序/验签失败测试、 +provider 一致性契约测试、故障注入、审计字段检查和 `mkdocs build --strict`。只有在提供外部 +依赖并实际运行 live suite 后,才能把相应的 `✅*` 证据升级为生产验收结论。 diff --git a/docs/docs/runtime-storage.md b/docs/docs/runtime-storage.md index 03503c1a..a0fddc3d 100644 --- a/docs/docs/runtime-storage.md +++ b/docs/docs/runtime-storage.md @@ -1,19 +1,18 @@ # Tenant 运行时持久化契约(Issue #48) -> 本页是 Issue #48 的先行设计与实现 ledger。它把 Session、入站事件和回复 -> Outbox 的租户边界、顺序和错误契约固定下来,再由后续代码阶段逐项落地。 -> 在 ledger 全部完成前,PR 使用 `Updates #48`,不会把未实现的能力描述成已交付。 +> 本页记录 Issue #48 的通用 RuntimeStore 契约,以及 Issue #108 的 Redis 实现边界。 +> 代码、测试和部署示例只把已经验证的能力标为已实现;未覆盖的外部后端仍属于后续工作。 ## 目标与非目标 -运行时持久化的事实源是 PostgreSQL。每个操作都必须显式带 `tenant_id`; +PostgreSQL 仍是控制面和默认运行时事实源;Redis 是可选的共享运行时后端。每个操作都必须显式带 `tenant_id`; Session/Runner 使用的命名空间只用于防碰撞,不能替代数据库授权。第一阶段覆盖: - Session 元数据、状态版本和生命周期; - `message_event` 入站幂等事实、事件序号和执行状态; - `reply_outbox` 分段回复、租约/fencing、重试和供应商回执。 -本 Issue 不实现 Redis、Memory/Knowledge/Artifact 生产适配、AuditEvent/usage/cost、 +Issue #48 不实现 Memory/Knowledge/Artifact 的其他生产适配、AuditEvent/usage/cost、 完整 IM webhook/media、分布式调度、KMS/Vault 或告警平台。API principal 继续由 Gateway HTTP 层的进程内幂等存储保护;跨进程 durable inbound claim 只在已验证 Channel principal 上启用,因为 `message_event.binding_id` 必须引用真实的控制面 Binding。 @@ -137,9 +136,53 @@ provider reconciliation,`accepted` 直接确认,`rejected` 重试,`unknown ## Bootstrap 与恢复 Bootstrap 必须显式选择 Session capability。`TRPC_SESSION_BACKEND=postgres` 时, -必须同时提供已迁移的 `TRPC_POSTGRES_DSN`;未知值、缺失 DSN 或 migration 验证失败 -均 fail-closed。`inmemory` 只用于开发和测试,并在 readiness/启动日志中明确显示 -非持久化。新进程连接同一 DSN 后应能读取已有 Session、事件和未发送 Outbox。 +必须同时提供已迁移的 `TRPC_POSTGRES_DSN`;`TRPC_SESSION_BACKEND=redis` 时,必须提供 +`TRPC_REDIS_ADDR` 并在启动和 readiness 阶段成功 PING。未知值、缺失地址、连接失败或 +migration 验证失败均 fail-closed,不会静默回退到 InMemory。`inmemory` 只用于开发和测试, +并在 readiness/启动日志中明确显示非持久化。新进程连接同一后端后应能读取已有 Session、 +事件、Memory 和未发送 Outbox。 + +### Redis 实现范围(Issue #108) + +Redis provider 通过 Backend Profile 的 `Provider: "redis"` 选择,只注册 `session` 和 +`memory` capability;`summary`、`knowledge`、`artifact`、`audit` 等 capability 在 Catalog +校验阶段拒绝 Redis。每条 profile binding 必须使用 `redis://` endpoint;若 +设置 `SecretRef`,它只能解析到当前 tenant 的 Redis 密码,密码不会进入 profile、快照、日志 +或错误文本。Provider 构造时再次校验 endpoint/secret scope,避免不同 tenant 复用错误配置。 + +每个 tenant 使用一个 Redis key: + +```text +: +``` + +默认前缀为 `trpc:runtime:v1`。key 的 tenant 部分使用 UTF-8 字节 hex 编码,避免简单拼接造成 +边界碰撞。value 是版本化 JSON 状态文档,当前 `version` 为 `1`,包含 Session、Event、 +event history、Reply Outbox、correlation、Memory 和 index handoff 集合。写入使用 +`WATCH/MULTI` CAS;event 序号、重复消息 claim、lease/fencing 和完整 reply batch 在一次原子 +状态更新中提交。 + +Redis key 没有隐式 TTL。Session、事件、历史、Memory 和 Outbox 不会因为连接池或重启自动过期; +保留、归档和删除必须由显式业务操作或后续运维工具完成。当前没有 Redis/PostgreSQL 迁移、 +双写、shadow read 或自动 cutover 工具;迁移方案仍按 Backend Profile 版本切换另行设计。 + +主要环境变量如下: + +| 变量 | 必需/默认 | 说明 | +| --- | --- | --- | +| `TRPC_REDIS_ADDR` | Redis 模式必需 | `host:port`,地址不写入错误或日志 | +| `TRPC_REDIS_PASSWORD` | 否 | 通过 Secret 注入的 Redis 密码;不写入配置快照 | +| `TRPC_REDIS_SECRET_REF` | 否,`env/trpc-redis-password` | Backend Profile 可使用的租户 SecretRef | +| `TRPC_REDIS_DB` | 否,`0` | Redis logical database,范围 `0..32768` | +| `TRPC_REDIS_KEY_PREFIX` | 否,`trpc:runtime:v1` | 共享实例的命名空间前缀 | +| `TRPC_REDIS_DIAL_TIMEOUT` | 否 | Go duration,例如 `500ms` | +| `TRPC_REDIS_READ_TIMEOUT` | 否 | Go duration,例如 `500ms` | +| `TRPC_REDIS_WRITE_TIMEOUT` | 否 | Go duration,例如 `500ms` | +| `TRPC_REDIS_POOL_SIZE` | 否 | 大于 `0` 时覆盖客户端连接池大小 | + +本地 Compose 已包含带 AOF 的 Redis 7 服务;生产/Kubernetes 仍应使用外部 Redis,并通过 Secret +Manager 注入密码。可选 live conformance/reconnect 测试读取 `REDIS_RUNTIME_TEST_ADDR`;未设置 +时显式 skip,不把本地 miniredis 测试冒充生产 Redis 证据。 真实验收测试使用可选的 `POSTGRES_RUNTIME_TEST_DSN`,并要求该 DSN 已有可写的 `POSTGRES_RUNTIME_TEST_TENANT_ID` 与 `POSTGRES_RUNTIME_TEST_BINDING_ID`;测试会执行 完整 RuntimeStore 操作、关闭连接、重新打开连接并验证 Session/Event/History/Outbox @@ -156,12 +199,13 @@ Bootstrap 必须显式选择 Session capability。`TRPC_SESSION_BACKEND=postgres | Bootstrap 显式 Session capability 与 fail-closed | 3 | 环境配置、RuntimeStore-backed session.Service、重启恢复测试 | ✅ | | durable Event payload/history 与完整 Event 状态生命周期 | 4 | `runtime_event_history`、fresh delegate replay、状态迁移测试 | ✅ | | Outbox worker/reconciliation/provider delivery | 5 | fenced worker、重试/死信/过期 lease 与 provider 测试 | ✅ | +| Redis RuntimeStore/MemoryStore 与 tenant-scoped bootstrap | Issue #108 | `runtime/storage/redis` miniredis conformance、配置/Catalog 边界、Compose 服务与可选 live reconnect 测试 | ✅* | | 真实 PostgreSQL/InMemory conformance 与 fresh-process restart | 6 | `POSTGRES_RUNTIME_TEST_DSN` 可选 live suite 与 reopen 证据 | ✅* | | verified Channel duplicate Runner suppression | 6 | RuntimeStore claim + 并发 Gateway Runner invocation-count 测试 | ✅ | | 租户越权、取消、脱敏和防御性返回 | 1–6 | 双租户 conformance 与错误边界测试 | ✅ | | `go test`、race、vet、build、MkDocs strict | 最终 | PR 验证记录与 CI | ✅ | -`✅*` 表示测试代码和重启路径已交付;live PostgreSQL 证据只有在 CI/本地实际 -提供上述 DSN 时才可勾选,未设置 DSN 的默认测试运行会 skip。 +`✅*` 表示测试代码和重启路径已交付;live PostgreSQL/Redis 证据只有在 CI/本地实际 +提供对应 DSN/地址时才可勾选,未设置变量的默认测试运行会显式 skip。 在代码阶段完成后,本表必须与 PR 描述同步;未完成项目保留为明确的后续阶段。 diff --git a/go.mod b/go.mod index 35c4edfd..4dc5b4a9 100644 --- a/go.mod +++ b/go.mod @@ -6,10 +6,12 @@ require trpc.group/trpc-go/trpc-agent-go v1.11.2 require ( github.com/DATA-DOG/go-sqlmock v1.5.2 + github.com/alicebob/miniredis/v2 v2.34.0 github.com/go-sql-driver/mysql v1.8.1 github.com/go-telegram/bot v1.23.0 github.com/google/uuid v1.6.0 github.com/jackc/pgx/v5 v5.6.0 + github.com/redis/go-redis/v9 v9.6.1 go.opentelemetry.io/otel v1.29.0 go.opentelemetry.io/otel/exporters/otlp/otlpmetric/otlpmetrichttp v1.29.0 go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.29.0 @@ -22,21 +24,20 @@ require ( require ( filippo.io/edwards25519 v1.1.0 // indirect + github.com/alicebob/gopher-json v0.0.0-20230218143504-906a9b012302 // indirect github.com/bmatcuk/doublestar/v4 v4.9.1 // indirect github.com/cenkalti/backoff/v4 v4.3.0 // indirect + github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/creack/pty v1.1.24 // indirect + github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect github.com/go-logr/logr v1.4.3 // indirect github.com/go-logr/stdr v1.2.2 // indirect github.com/grpc-ecosystem/grpc-gateway/v2 v2.22.0 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a // indirect github.com/jackc/puddle/v2 v2.2.1 // indirect - github.com/openai/openai-go v1.12.0 // indirect github.com/panjf2000/ants/v2 v2.10.0 // indirect - github.com/tidwall/gjson v1.14.4 // indirect - github.com/tidwall/match v1.1.1 // indirect - github.com/tidwall/pretty v1.2.1 // indirect - github.com/tidwall/sjson v1.2.5 // indirect + github.com/yuin/gopher-lua v1.1.1 // indirect go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.29.0 // indirect go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.29.0 // indirect go.opentelemetry.io/proto/otlp v1.3.1 // indirect diff --git a/go.sum b/go.sum index 7a07bb1f..5af7d4a6 100644 --- a/go.sum +++ b/go.sum @@ -2,17 +2,29 @@ filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= github.com/DATA-DOG/go-sqlmock v1.5.2 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU= github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU= +github.com/alicebob/gopher-json v0.0.0-20230218143504-906a9b012302 h1:uvdUDbHQHO85qeSydJtItA4T55Pw6BtAejd0APRJOCE= +github.com/alicebob/gopher-json v0.0.0-20230218143504-906a9b012302/go.mod h1:SGnFV6hVsYE877CKEZ6tDNTjaSXYUk6QqoIK6PrAtcc= +github.com/alicebob/miniredis/v2 v2.34.0 h1:mBFWMaJSNL9RwdGRyEDoAAv8OQc5UlEhLDQggTglU/0= +github.com/alicebob/miniredis/v2 v2.34.0/go.mod h1:kWShP4b58T1CW0Y5dViCd5ztzrDqRWqM3nksiyXk5s8= github.com/bmatcuk/doublestar/v4 v4.9.1 h1:X8jg9rRZmJd4yRy7ZeNDRnM+T3ZfHv15JiBJ/avrEXE= github.com/bmatcuk/doublestar/v4 v4.9.1/go.mod h1:xBQ8jztBU6kakFMg+8WGxn0c6z1fTSPVIjEY1Wr7jzc= +github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs= +github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c= +github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA= +github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0= github.com/bufbuild/protocompile v0.14.1 h1:iA73zAf/fyljNjQKwYzUHD6AD4R8KMasmwa/FBatYVw= github.com/bufbuild/protocompile v0.14.1/go.mod h1:ppVdAIhbr2H8asPk6k4pY7t9zB1OU5DoEw9xY/FUi1c= github.com/cenkalti/backoff/v4 v4.3.0 h1:MyRJ/UdXutAwSAT+s3wNd7MfTIcy71VQueUuFK343L8= github.com/cenkalti/backoff/v4 v4.3.0/go.mod h1:Y3VNntkOUPxTVeUxJ/G5vcM//AlwfmyYozVcomhLiZE= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/creack/pty v1.1.24 h1:bJrF4RRfyJnbTJqzRLHzcGaZK1NeM5kTC9jGgovnR1s= github.com/creack/pty v1.1.24/go.mod h1:08sCNb52WyoAwi2QDyzUCTgcvVFhUzewun7wtTfvcwE= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78= +github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc= github.com/go-ego/gse v1.0.0 h1:GNbtH1WP7Yd1VvCZ85fIK6eVEe7RctmgmnwliEPUMNA= github.com/go-ego/gse v1.0.0/go.mod h1:Gt3A9Ry1Eso2Kza4MRaiZ7f2DTAvActmETY46Lxg0gU= github.com/go-logr/logr v1.2.2/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A= @@ -51,6 +63,8 @@ github.com/panjf2000/ants/v2 v2.10.0 h1:zhRg1pQUtkyRiOFo2Sbqwjp0GfBNo9cUY2/Grpx1 github.com/panjf2000/ants/v2 v2.10.0/go.mod h1:7ZxyxsqE4vvW0M7LSD8aI3cKwgFhBHbxnlN8mDqHa1I= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/redis/go-redis/v9 v9.6.1 h1:HHDteefn6ZkTtY5fGUE8tj8uy85AHk6zP7CpzIAM0y4= +github.com/redis/go-redis/v9 v9.6.1/go.mod h1:0C0c6ycQsdpVNQpxb1njEQIqkx5UcsM8FJCQLgE9+RA= github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= @@ -63,12 +77,10 @@ github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO github.com/stretchr/testify v1.8.2/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= -github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= github.com/tidwall/gjson v1.14.4 h1:uo0p8EbA09J7RQaflQ1aBRffTR7xedD2bcIVSYxLnkM= github.com/tidwall/gjson v1.14.4/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA= github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= -github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4= github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY= @@ -77,6 +89,8 @@ github.com/vcaesar/cedar v0.20.2 h1:TDx7AdZhilKcfE1WvdToTJf5VrC/FXcUOW+KY1upLZ4= github.com/vcaesar/cedar v0.20.2/go.mod h1:lyuGvALuZZDPNXwpzv/9LyxW+8Y6faN7zauFezNsnik= github.com/yuin/goldmark v1.4.13 h1:fVcFKWvrslecOb/tg+Cc05dkeYx540o0FuFt3nUVDoE= github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY= +github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M= +github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw= go.opentelemetry.io/otel v1.29.0 h1:PdomN/Al4q/lN6iBJEN3AwPvUiHPMlt93c8bqTG5Llw= go.opentelemetry.io/otel v1.29.0/go.mod h1:N/WtXPs1CNCUEx+Agz5uouwCba+i+bJGFicT8SR4NP8= go.opentelemetry.io/otel/exporters/otlp/otlpmetric/otlpmetrichttp v1.29.0 h1:xvhQxJ/C9+RTnAj5DpTg7LSM1vbbMTiXt7e9hsfqHNw= diff --git a/trpcservice/bootstrap/environment.go b/trpcservice/bootstrap/environment.go index 6c7a423f..4af52a80 100644 --- a/trpcservice/bootstrap/environment.go +++ b/trpcservice/bootstrap/environment.go @@ -32,6 +32,7 @@ import ( runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" runtimestorageinmemory "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/inmemory" runtimestoragepostgres "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/postgres" + runtimestorageredis "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/redis" "github.com/XnLemon/trpc-agent-service/trpcservice/storage/mysql" "github.com/XnLemon/trpc-agent-service/trpcservice/storage/postgres" "github.com/XnLemon/trpc-agent-service/trpcservice/tenant" @@ -66,7 +67,18 @@ const ( // #nosec G101 -- environment variable name, not a secret. envModelSecretRef = "TRPC_MODEL_SECRET_REF" envSessionBackend = "TRPC_SESSION_BACKEND" - envDemoMode = "TRPC_DEMO_MODE" + envRedisAddr = "TRPC_REDIS_ADDR" + // #nosec G101 -- environment variable name, not a credential. + envRedisPassword = "TRPC_REDIS_PASSWORD" + envRedisDB = "TRPC_REDIS_DB" + envRedisKeyPrefix = "TRPC_REDIS_KEY_PREFIX" + // #nosec G101 -- environment variable name, not a secret. + envRedisSecretRef = "TRPC_REDIS_SECRET_REF" + envRedisDialTimeout = "TRPC_REDIS_DIAL_TIMEOUT" + envRedisReadTimeout = "TRPC_REDIS_READ_TIMEOUT" + envRedisWriteTimeout = "TRPC_REDIS_WRITE_TIMEOUT" + envRedisPoolSize = "TRPC_REDIS_POOL_SIZE" + envDemoMode = "TRPC_DEMO_MODE" // #nosec G101 -- environment variable name, not a secret. envWeComCallbackToken = "WECOM_CALLBACK_TOKEN" envWeComEncodingAESKey = "WECOM_ENCODING_AES_KEY" @@ -87,6 +99,7 @@ const ( // #nosec G101 -- symbolic secret reference, not secret material. defaultModelSecretRef = "env/trpc-model-api-key" defaultSubjectID = "service" + maxRedisDB = 1 << 15 ) var ( @@ -97,6 +110,8 @@ var ( verifyEnvironmentMigrations = migrations.Verify verifyMySQLEnvironmentMigrations = migrations.VerifyMySQL newEnvironmentRuntimeStore = environmentRuntimeStore + newEnvironmentRedisRuntimeStore = environmentRedisRuntimeStore + newEnvironmentInMemoryFallback = func() runtimestorage.RuntimeStore { return runtimestorageinmemory.New() } environmentWeComOwnerFunc = environmentWeComOwner newEnvironmentWeComWorker = outbox.New ) @@ -122,6 +137,9 @@ type environmentConfig struct { endpointHosts []string secretRef string runtimeStorage string + redis runtimestorageredis.Config + redisEndpoint string + redisSecretRef string demoMode bool wecom *environmentWeComConfig telemetry observability.Provider @@ -135,6 +153,25 @@ type environmentWeComConfig struct { secretRef string } +// environmentRuntimeStores owns process-scoped runtime stores. The primary +// store serves ingress and outbox processing; provider stores serve Backend +// Profile capability materialization. +type environmentRuntimeStores struct { + primary runtimestorage.RuntimeStore + providers map[string]runtimestorage.RuntimeStore + owned []runtimestorage.RuntimeStore +} + +func (stores environmentRuntimeStores) Close() error { + var errs []error + for _, store := range stores.owned { + if store != nil { + errs = append(errs, store.Close()) + } + } + return errors.Join(errs...) +} + // NewFromEnvironment assembles the production bootstrap graph from explicit // process configuration. It fails before binding an HTTP server when the // durable control plane or required credentials are not configured. @@ -176,16 +213,23 @@ func NewFromEnvironment(ctx context.Context) (*Runtime, error) { return nil, err } delegateSessions := inmemory.NewSessionService() - runtimeStore, err := newEnvironmentRuntimeStore(config.runtimeStorage, db) + runtimeStores, err := newEnvironmentRuntimeStoresForConfig(ctx, config, db) if err != nil { _ = delegateSessions.Close() _ = db.Close() + if ctx.Err() != nil { + return nil, ctx.Err() + } + if config.runtimeStorage == "redis" { + return nil, fmt.Errorf("%w: Redis runtime storage is unavailable", ErrInvalidConfig) + } return nil, err } + runtimeStore := runtimeStores.primary tenantRepo, appRepo, channelRepo, auditWriter, err := environmentRepositories(config, db) if err != nil { _ = delegateSessions.Close() - _ = runtimeStore.Close() + _ = runtimeStores.Close() _ = db.Close() return nil, fmt.Errorf("%w: environment repositories: %v", ErrInvalidConfig, err) } @@ -193,21 +237,21 @@ func NewFromEnvironment(ctx context.Context) (*Runtime, error) { wecomFactory, wecomWorker, err := environmentWeComComponents(config, channelRepo, tenantRepo, appRepo, runtimeStore, auditWriter) if err != nil { _ = delegateSessions.Close() - _ = runtimeStore.Close() + _ = runtimeStores.Close() _ = db.Close() return nil, fmt.Errorf("%w: wecom components: %v", ErrInvalidConfig, err) } - secretRegistry, modelRegistry, backendRegistry, err := environmentRegistries(config, delegateSessions, runtimeStore) + secretRegistry, modelRegistry, backendRegistry, err := environmentRegistriesForStores(config, delegateSessions, runtimeStores) if err != nil { _ = delegateSessions.Close() - _ = runtimeStore.Close() + _ = runtimeStores.Close() _ = db.Close() return nil, fmt.Errorf("%w: environment registries: %v", ErrInvalidConfig, err) } storageFactory, err := backend.NewRegistryStorageFactory(backendRegistry, secretRegistry) if err != nil { _ = delegateSessions.Close() - _ = runtimeStore.Close() + _ = runtimeStores.Close() _ = db.Close() return nil, fmt.Errorf("%w: storage factory: %v", ErrInvalidConfig, err) } @@ -233,20 +277,17 @@ func NewFromEnvironment(ctx context.Context) (*Runtime, error) { OutboxPollInterval: time.Second, AuditWriter: auditWriter, Ping: func(pingContext context.Context) error { - if config.driver == ControlPlaneDriverMySQL { - return mysql.Ping(pingContext, db) - } - return postgres.Ping(pingContext, db) + return environmentPing(pingContext, config.driver, db, runtimeStore) }, Migrate: applyMigrations, VerifyMigrations: verifyMigrations, CloseDependencies: func() error { - return errors.Join(delegateSessions.Close(), runtimeStore.Close()) + return errors.Join(delegateSessions.Close(), runtimeStores.Close()) }, }) if err != nil { _ = delegateSessions.Close() - _ = runtimeStore.Close() + _ = runtimeStores.Close() _ = db.Close() return nil, err } @@ -345,19 +386,34 @@ func environmentWeComComponents(config environmentConfig, channelsRepo channels. } func environmentRegistries(config environmentConfig, delegateSessions session.Service, runtimeStore runtimestorage.RuntimeStore) (*modelprofile.SecretRegistry, *modelprofile.ModelProviderRegistry, *backend.ProviderRegistry, error) { + providerName := environmentRuntimeProviderName(config.runtimeStorage) + return environmentRegistriesForStores(config, delegateSessions, environmentRuntimeStores{ + primary: runtimeStore, + providers: map[string]runtimestorage.RuntimeStore{providerName: runtimeStore}, + }) +} + +type environmentRuntimeProviderSpec struct { + name string + capabilities []backend.Capability + store runtimestorage.RuntimeStore +} + +func environmentRegistriesForStores(config environmentConfig, delegateSessions session.Service, runtimeStores environmentRuntimeStores) (*modelprofile.SecretRegistry, *modelprofile.ModelProviderRegistry, *backend.ProviderRegistry, error) { secretRegistry := modelprofile.NewSecretRegistry() modelRegistry := modelprofile.NewModelProviderRegistry() backendRegistry := backend.NewProviderRegistry() + runtimeProviders, err := environmentRuntimeProviders(config, runtimeStores) + if err != nil { + return nil, nil, nil, err + } for _, identity := range config.apiIdentities { if config.demoMode { if err := modelRegistry.Register(identity.TenantID, demoModelProvider, environmentModelFactory{}); err != nil { return nil, nil, nil, err } - for _, capability := range []backend.Capability{backend.CapabilitySession, backend.CapabilityMemory, backend.CapabilitySummary, backend.CapabilityKnowledge, backend.CapabilityArtifact, backend.CapabilityAudit} { - provider := environmentRuntimeCapabilityProvider{capability: capability, delegate: delegateSessions, store: runtimeStore, telemetry: config.telemetry, backend: config.runtimeStorage} - if err := backendRegistry.Register(identity.TenantID, capability, "inmemory", provider); err != nil { - return nil, nil, nil, err - } + if err := registerEnvironmentRuntimeProviders(backendRegistry, identity.TenantID, delegateSessions, config, runtimeProviders); err != nil { + return nil, nil, nil, err } continue } @@ -374,16 +430,66 @@ func environmentRegistries(config environmentConfig, delegateSessions session.Se if err := modelRegistry.Register(identity.TenantID, config.modelProvider, environmentModelFactory{}); err != nil { return nil, nil, nil, err } - for _, capability := range []backend.Capability{backend.CapabilitySession, backend.CapabilityMemory, backend.CapabilitySummary, backend.CapabilityKnowledge, backend.CapabilityArtifact, backend.CapabilityAudit} { - provider := environmentRuntimeCapabilityProvider{capability: capability, delegate: delegateSessions, store: runtimeStore, telemetry: config.telemetry, backend: config.runtimeStorage} - if err := backendRegistry.Register(identity.TenantID, capability, "inmemory", provider); err != nil { + if config.runtimeStorage == "redis" && config.redis.Password != "" { + if err := secretRegistry.RegisterValue(modelprofile.SecretScope{TenantID: identity.TenantID, SecretRef: config.redisSecretRef}, config.redis.Password); err != nil { return nil, nil, nil, err } } + if err := registerEnvironmentRuntimeProviders(backendRegistry, identity.TenantID, delegateSessions, config, runtimeProviders); err != nil { + return nil, nil, nil, err + } } return secretRegistry, modelRegistry, backendRegistry, nil } +func environmentRuntimeProviders(config environmentConfig, stores environmentRuntimeStores) ([]environmentRuntimeProviderSpec, error) { + providerName := environmentRuntimeProviderName(config.runtimeStorage) + primary := stores.providers[providerName] + if primary == nil { + return nil, fmt.Errorf("%w: primary runtime provider is unavailable", ErrInvalidConfig) + } + providers := []environmentRuntimeProviderSpec{{name: providerName, capabilities: environmentRuntimeCapabilities(config.runtimeStorage), store: primary}} + if config.runtimeStorage != "redis" { + return providers, nil + } + fallback := stores.providers["inmemory"] + if fallback == nil { + return nil, fmt.Errorf("%w: in-memory runtime provider is unavailable", ErrInvalidConfig) + } + return append(providers, environmentRuntimeProviderSpec{name: "inmemory", capabilities: environmentRuntimeCapabilities("inmemory"), store: fallback}), nil +} + +func registerEnvironmentRuntimeProviders(registry *backend.ProviderRegistry, tenantID string, delegateSessions session.Service, config environmentConfig, runtimeProviders []environmentRuntimeProviderSpec) error { + for _, runtimeProvider := range runtimeProviders { + for _, capability := range runtimeProvider.capabilities { + provider := environmentRuntimeCapabilityProvider{capability: capability, delegate: delegateSessions, store: runtimeProvider.store, telemetry: config.telemetry, backend: runtimeProvider.name} + if runtimeProvider.name == "redis" { + provider.redisEndpoint = config.redisEndpoint + provider.redisSecretRef = config.redisSecretRef + provider.redisPasswordRequired = config.redis.Password != "" + } + if err := registry.Register(tenantID, capability, runtimeProvider.name, provider); err != nil { + return err + } + } + } + return nil +} + +func environmentRuntimeProviderName(runtimeStorage string) string { + if runtimeStorage == "redis" { + return "redis" + } + return "inmemory" +} + +func environmentRuntimeCapabilities(runtimeStorage string) []backend.Capability { + if runtimeStorage == "redis" { + return []backend.Capability{backend.CapabilitySession, backend.CapabilityMemory} + } + return []backend.Capability{backend.CapabilitySession, backend.CapabilityMemory, backend.CapabilitySummary, backend.CapabilityKnowledge, backend.CapabilityArtifact, backend.CapabilityAudit} +} + func loadEnvironment() (environmentConfig, error) { demoMode, err := environmentBool(envDemoMode) if err != nil { @@ -598,8 +704,17 @@ func parseEnvironmentModelAPIKeys(value string) (map[string]string, error) { func (config *environmentConfig) loadRuntime() error { config.subjectID = strings.TrimSpace(config.subjectID) - if config.runtimeStorage != "postgres" && config.runtimeStorage != "inmemory" { - return fmt.Errorf("%w: %s must be explicitly set to postgres or inmemory", ErrInvalidConfig, envSessionBackend) + switch config.runtimeStorage { + case "postgres", "inmemory": + case "redis": + if config.demoMode { + return fmt.Errorf("%w: %s cannot use redis in demo mode", ErrInvalidConfig, envSessionBackend) + } + if err := config.loadRedis(); err != nil { + return err + } + default: + return fmt.Errorf("%w: %s must be explicitly set to postgres, redis or inmemory", ErrInvalidConfig, envSessionBackend) } if config.demoMode && (config.driver != ControlPlaneDriverPostgres || config.runtimeStorage != "inmemory") { return fmt.Errorf("%w: %s requires PostgreSQL control plane and inmemory session backend", ErrInvalidConfig, envDemoMode) @@ -610,6 +725,76 @@ func (config *environmentConfig) loadRuntime() error { return nil } +func (config *environmentConfig) loadRedis() error { + addr, err := requiredEnvironment(envRedisAddr) + if err != nil { + return err + } + if strings.ContainsAny(addr, "\r\n") { + return fmt.Errorf("%w: %s is invalid", ErrInvalidConfig, envRedisAddr) + } + db, err := environmentInteger(envRedisDB, environmentOrDefault(envRedisDB, "0"), 0, maxRedisDB) + if err != nil { + return err + } + dialTimeout, err := environmentDuration(envRedisDialTimeout) + if err != nil { + return err + } + readTimeout, err := environmentDuration(envRedisReadTimeout) + if err != nil { + return err + } + writeTimeout, err := environmentDuration(envRedisWriteTimeout) + if err != nil { + return err + } + poolSize, err := environmentInteger(envRedisPoolSize, environmentOrDefault(envRedisPoolSize, "0"), 0, 0) + if err != nil { + return err + } + keyPrefix := environmentOrDefault(envRedisKeyPrefix, "trpc:runtime:v1") + if strings.ContainsAny(keyPrefix, "\r\n") || strings.TrimSpace(keyPrefix) == "" { + return fmt.Errorf("%w: %s is invalid", ErrInvalidConfig, envRedisKeyPrefix) + } + password := os.Getenv(envRedisPassword) + if strings.ContainsAny(password, "\r\n") { + return fmt.Errorf("%w: %s is invalid", ErrInvalidConfig, envRedisPassword) + } + config.redis = runtimestorageredis.Config{Addr: addr, Password: password, DB: db, KeyPrefix: keyPrefix, DialTimeout: dialTimeout, ReadTimeout: readTimeout, WriteTimeout: writeTimeout, PoolSize: poolSize} + config.redisEndpoint = redisEndpoint(addr) + config.redisSecretRef = environmentOrDefault(envRedisSecretRef, "env/trpc-redis-password") + if _, err := modelprofile.NewSecretValue(config.redis.Password); err != nil && config.redis.Password != "" { + return fmt.Errorf("%w: %s is invalid", ErrInvalidConfig, envRedisPassword) + } + for _, identity := range config.apiIdentities { + if err := (modelprofile.SecretScope{TenantID: identity.TenantID, SecretRef: config.redisSecretRef}).Validate(); err != nil { + return fmt.Errorf("%w: %s is invalid", ErrInvalidConfig, envRedisSecretRef) + } + } + return nil +} + +func environmentInteger(name, value string, min, max int) (int, error) { + parsed, err := strconv.Atoi(strings.TrimSpace(value)) + if err != nil || parsed < min || (max > 0 && parsed > max) { + return 0, fmt.Errorf("%w: %s is invalid", ErrInvalidConfig, name) + } + return parsed, nil +} + +func environmentDuration(name string) (time.Duration, error) { + value := strings.TrimSpace(os.Getenv(name)) + if value == "" { + return 0, nil + } + parsed, err := time.ParseDuration(value) + if err != nil || parsed <= 0 { + return 0, fmt.Errorf("%w: %s is invalid", ErrInvalidConfig, name) + } + return parsed, nil +} + func (config *environmentConfig) loadWeCom() error { values := []string{strings.TrimSpace(os.Getenv(envWeComCallbackToken)), strings.TrimSpace(os.Getenv(envWeComEncodingAESKey)), strings.TrimSpace(os.Getenv(envWeComAppSecret)), strings.TrimSpace(os.Getenv(envWeComSecretRef))} configured := 0 @@ -647,6 +832,65 @@ func environmentRuntimeStore(kind string, db *sql.DB) (runtimestorage.RuntimeSto } } +func newEnvironmentRuntimeStoreForConfig(ctx context.Context, config environmentConfig, db *sql.DB) (runtimestorage.RuntimeStore, error) { + if config.runtimeStorage == "redis" { + return newEnvironmentRedisRuntimeStore(ctx, config) + } + return newEnvironmentRuntimeStore(config.runtimeStorage, db) +} + +func newEnvironmentRuntimeStoresForConfig(ctx context.Context, config environmentConfig, db *sql.DB) (environmentRuntimeStores, error) { + primary, err := newEnvironmentRuntimeStoreForConfig(ctx, config, db) + if err != nil { + return environmentRuntimeStores{}, err + } + providerName := environmentRuntimeProviderName(config.runtimeStorage) + stores := environmentRuntimeStores{ + primary: primary, + providers: map[string]runtimestorage.RuntimeStore{providerName: primary}, + owned: []runtimestorage.RuntimeStore{primary}, + } + if config.runtimeStorage != "redis" { + return stores, nil + } + fallback := newEnvironmentInMemoryFallback() + stores.providers["inmemory"] = fallback + stores.owned = append(stores.owned, fallback) + return stores, nil +} + +func environmentRedisRuntimeStore(ctx context.Context, config environmentConfig) (runtimestorage.RuntimeStore, error) { + store, err := runtimestorageredis.NewFromConfig(ctx, config.redis) + if err != nil { + return nil, err + } + return store, nil +} + +func environmentPing(ctx context.Context, driver ControlPlaneDriver, db *sql.DB, runtimeStore runtimestorage.RuntimeStore) error { + if ctx == nil { + return ErrInvalidConfig + } + if err := ctx.Err(); err != nil { + return err + } + if driver == ControlPlaneDriverMySQL { + if err := mysql.Ping(ctx, db); err != nil { + return err + } + } else if err := postgres.Ping(ctx, db); err != nil { + return err + } + if pinger, ok := runtimeStore.(interface{ Ping(context.Context) error }); ok { + return pinger.Ping(ctx) + } + return nil +} + +func redisEndpoint(addr string) string { + return "redis://" + strings.TrimSpace(addr) +} + func environmentCatalogs(config environmentConfig) (*modelprofile.ProviderCatalog, *backend.ProviderCatalog, error) { if config.demoMode { if config.modelProvider != demoModelProvider { @@ -661,7 +905,7 @@ func environmentCatalogs(config environmentConfig) (*modelprofile.ProviderCatalo if err != nil { return nil, nil, fmt.Errorf("%w: demo model catalog is invalid", ErrInvalidConfig) } - backendCatalog, err := newEnvironmentBackendCatalog() + backendCatalog, err := newEnvironmentBackendCatalog(config.runtimeStorage) if err != nil { return nil, nil, err } @@ -681,21 +925,36 @@ func environmentCatalogs(config environmentConfig) (*modelprofile.ProviderCatalo if err != nil { return nil, nil, fmt.Errorf("%w: model catalog is invalid", ErrInvalidConfig) } - backendCatalog, err := newEnvironmentBackendCatalog() + backendCatalog, err := newEnvironmentBackendCatalog(config.runtimeStorage) if err != nil { return nil, nil, err } return modelCatalog, backendCatalog, nil } -func newEnvironmentBackendCatalog() (*backend.ProviderCatalog, error) { - backendCatalog, err := backend.NewProviderCatalog(backend.ProviderSpec{ +func newEnvironmentBackendCatalog(runtimeStorage string) (*backend.ProviderCatalog, error) { + inMemory := backend.ProviderSpec{ Provider: "inmemory", Capabilities: []backend.Capability{backend.CapabilitySession, backend.CapabilityMemory, backend.CapabilitySummary, backend.CapabilityKnowledge, backend.CapabilityArtifact, backend.CapabilityAudit}, EndpointPolicy: backend.FieldForbidden, SecretRefPolicy: backend.FieldForbidden, Options: map[string]backend.OptionSpec{}, - }) + } + if runtimeStorage == "redis" { + backendCatalog, err := backend.NewProviderCatalog(backend.ProviderSpec{ + Provider: "redis", + Capabilities: []backend.Capability{backend.CapabilitySession, backend.CapabilityMemory}, + EndpointPolicy: backend.FieldRequired, + EndpointSchemes: []string{"redis"}, + SecretRefPolicy: backend.FieldOptional, + Options: map[string]backend.OptionSpec{}, + }, inMemory) + if err != nil { + return nil, fmt.Errorf("%w: backend catalog is invalid", ErrInvalidConfig) + } + return backendCatalog, nil + } + backendCatalog, err := backend.NewProviderCatalog(inMemory) if err != nil { return nil, fmt.Errorf("%w: backend catalog is invalid", ErrInvalidConfig) } @@ -827,17 +1086,46 @@ type environmentSessionCapabilityProvider struct { } type environmentRuntimeCapabilityProvider struct { - capability backend.Capability - delegate session.Service - store runtimestorage.RuntimeStore - telemetry observability.Provider - backend string + capability backend.Capability + delegate session.Service + store runtimestorage.RuntimeStore + telemetry observability.Provider + backend string + redisEndpoint string + redisSecretRef string + redisPasswordRequired bool } -func (provider environmentRuntimeCapabilityProvider) New(ctx context.Context, input backend.StorageFactoryInput, _ backend.CapabilityBinding, _ modelprofile.SecretValue) (any, error) { +func (provider environmentRuntimeCapabilityProvider) New(ctx context.Context, input backend.StorageFactoryInput, binding backend.CapabilityBinding, secret modelprofile.SecretValue) (any, error) { if ctx == nil { return nil, context.Canceled } + if err := provider.validateRedisBinding(binding, secret); err != nil { + return nil, err + } + return provider.newCapability(ctx, input) +} + +func (provider environmentRuntimeCapabilityProvider) validateRedisBinding(binding backend.CapabilityBinding, secret modelprofile.SecretValue) error { + if provider.backend != "redis" { + return nil + } + if provider.capability != backend.CapabilitySession && provider.capability != backend.CapabilityMemory { + return backend.ErrStorageFactory + } + if provider.redisEndpoint != "" && binding.Endpoint != provider.redisEndpoint { + return backend.ErrStorageFactory + } + if provider.redisSecretRef != "" && binding.SecretRef != "" && binding.SecretRef != provider.redisSecretRef { + return backend.ErrStorageFactory + } + if provider.redisPasswordRequired && secret.Value() == "" { + return backend.ErrStorageFactory + } + return nil +} + +func (provider environmentRuntimeCapabilityProvider) newCapability(ctx context.Context, input backend.StorageFactoryInput) (any, error) { if provider.capability == backend.CapabilitySession { return runtimesessionpostgres.NewWithObservability(input.TenantID, provider.delegate, provider.store, provider.telemetry, provider.backend) } diff --git a/trpcservice/bootstrap/environment_multitenant_test.go b/trpcservice/bootstrap/environment_multitenant_test.go index 25eb6dda..e9e7553a 100644 --- a/trpcservice/bootstrap/environment_multitenant_test.go +++ b/trpcservice/bootstrap/environment_multitenant_test.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "errors" + "strings" "testing" "github.com/XnLemon/trpc-agent-service/trpcservice/backend" @@ -13,7 +14,9 @@ import ( runtimesessionpostgres "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/sessionpostgres" runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" runtimestorageinmemory "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/inmemory" + runtimestorageredis "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/redis" "github.com/XnLemon/trpc-agent-service/trpcservice/storage/postgres" + "github.com/alicebob/miniredis/v2" "trpc.group/trpc-go/trpc-agent-go/session/inmemory" ) @@ -458,6 +461,445 @@ func TestEnvironmentCatalogsAndIdentityListBoundaries(t *testing.T) { } } +func TestLoadEnvironmentRedisConfigurationIsExplicitAndBounded(t *testing.T) { + setRequiredEnvironment(t) + t.Setenv(envSessionBackend, "redis") + t.Setenv(envRedisAddr, "127.0.0.1:6379") + t.Setenv(envRedisPassword, "redis-secret") + t.Setenv(envRedisDB, "3") + t.Setenv(envRedisKeyPrefix, "service:runtime") + t.Setenv(envRedisSecretRef, "env/redis-password") + t.Setenv(envRedisDialTimeout, "150ms") + t.Setenv(envRedisReadTimeout, "250ms") + t.Setenv(envRedisWriteTimeout, "350ms") + t.Setenv(envRedisPoolSize, "12") + + config, err := loadEnvironment() + if err != nil { + t.Fatal(err) + } + if config.redis.Addr != "127.0.0.1:6379" || config.redis.Password != "redis-secret" || config.redis.DB != 3 || config.redis.KeyPrefix != "service:runtime" || config.redis.PoolSize != 12 { + t.Fatalf("redis config = %+v", config.redis) + } + if config.redisEndpoint != "redis://127.0.0.1:6379" || config.redisSecretRef != "env/redis-password" || config.redis.DialTimeout.String() != "150ms" || config.redis.ReadTimeout.String() != "250ms" || config.redis.WriteTimeout.String() != "350ms" { + t.Fatalf("redis endpoint/timeouts = %q/%q/%s/%s/%s", config.redisEndpoint, config.redisSecretRef, config.redis.DialTimeout, config.redis.ReadTimeout, config.redis.WriteTimeout) + } + + for _, test := range []struct { + name string + value string + }{ + {name: "missing address", value: ""}, + {name: "invalid db", value: "not-an-integer"}, + {name: "invalid timeout", value: "not-a-duration"}, + } { + t.Run(test.name, func(t *testing.T) { + setRequiredEnvironment(t) + t.Setenv(envSessionBackend, "redis") + t.Setenv(envRedisAddr, "127.0.0.1:6379") + switch test.name { + case "missing address": + t.Setenv(envRedisAddr, test.value) + case "invalid db": + t.Setenv(envRedisDB, test.value) + case "invalid timeout": + t.Setenv(envRedisDialTimeout, test.value) + } + if _, err := loadEnvironment(); !errors.Is(err, ErrInvalidConfig) { + t.Fatalf("redis configuration error = %v", err) + } + }) + } +} + +func TestEnvironmentRedisCatalogAndRegistryBoundaries(t *testing.T) { + const ( + tenantA = "t_00000000000000000000000000" + tenantB = "t_00000000000000000000000001" + ) + config := environmentConfig{runtimeStorage: "redis", redisEndpoint: "redis://127.0.0.1:6379", redisSecretRef: "env/redis-password", redis: runtimestorageredis.Config{Password: "redis-secret"}, modelProvider: defaultModelProvider, modelNames: []string{"gpt-4o-mini"}, endpointHosts: []string{"api.openai.com"}, secretRef: "env/model", apiIdentities: map[string]gateway.APIIdentity{ + "token-a": {TenantID: tenantA, AppID: "app-a", SubjectID: "subject-a"}, + "token-b": {TenantID: tenantB, AppID: "app-b", SubjectID: "subject-b"}, + }, modelAPIKeys: map[string]string{tenantA: "model-a", tenantB: "model-b"}} + _, backendCatalog, err := environmentCatalogs(config) + if err != nil { + t.Fatal(err) + } + if _, err := backend.NewProfile(backend.CreateInput{TenantID: tenantA, ProfileKey: "redis", DisplayName: "Redis", Bindings: []backend.CapabilityBinding{{Capability: backend.CapabilitySession, Provider: "redis", Endpoint: config.redisEndpoint, SecretRef: config.redisSecretRef}}}, backendCatalog); err != nil { + t.Fatalf("redis session binding = %v", err) + } + if _, err := backend.NewProfile(backend.CreateInput{TenantID: tenantA, ProfileKey: "redis-summary", DisplayName: "Redis Summary", Bindings: []backend.CapabilityBinding{{Capability: backend.CapabilitySummary, Provider: "redis", Endpoint: config.redisEndpoint}}}, backendCatalog); !errors.Is(err, backend.ErrInvalid) { + t.Fatalf("unsupported redis capability = %v", err) + } + if _, err := backend.NewProfile(backend.CreateInput{TenantID: tenantA, ProfileKey: "redis-endpoint", DisplayName: "Redis Endpoint", Bindings: []backend.CapabilityBinding{{Capability: backend.CapabilitySession, Provider: "redis", Endpoint: "redis://other:6379"}}}, backendCatalog); err != nil { + t.Fatalf("catalog should accept a valid redis endpoint before provider binding: %v", err) + } + + delegate := inmemory.NewSessionService() + redisStore := runtimestorageinmemory.New() + inMemoryStore := runtimestorageinmemory.New() + t.Cleanup(func() { _ = delegate.Close(); _ = redisStore.Close(); _ = inMemoryStore.Close() }) + secrets, _, providers, err := environmentRegistriesForStores(config, delegate, environmentRuntimeStores{ + primary: redisStore, + providers: map[string]runtimestorage.RuntimeStore{ + "redis": redisStore, + "inmemory": inMemoryStore, + }, + }) + if err != nil { + t.Fatal(err) + } + secret, err := secrets.Resolve(context.Background(), modelprofile.SecretScope{TenantID: tenantA, SecretRef: config.redisSecretRef}) + if err != nil || secret.Value() != "redis-secret" { + t.Fatalf("tenant redis secret = %q, %v", secret.Value(), err) + } + if _, err := secrets.Resolve(context.Background(), modelprofile.SecretScope{TenantID: tenantA, SecretRef: "env/other"}); err == nil { + t.Fatal("foreign redis secret reference was accepted") + } + provider, err := providers.Resolve(context.Background(), backend.StorageFactoryInput{TenantID: tenantA}, backend.CapabilityBinding{Capability: backend.CapabilitySession, Provider: "redis"}) + if err != nil || provider == nil { + t.Fatalf("tenant redis provider = %v", err) + } + if _, err := providers.Resolve(context.Background(), backend.StorageFactoryInput{TenantID: tenantA}, backend.CapabilityBinding{Capability: backend.CapabilitySummary, Provider: "redis"}); !errors.Is(err, backend.ErrProviderUnavailable) { + t.Fatalf("unsupported redis provider capability = %v", err) + } + if _, err := providers.Resolve(context.Background(), backend.StorageFactoryInput{TenantID: "t_00000000000000000000000002"}, backend.CapabilityBinding{Capability: backend.CapabilitySession, Provider: "redis"}); !errors.Is(err, backend.ErrProviderUnavailable) { + t.Fatalf("unregistered tenant redis provider = %v", err) + } + value, err := provider.New(context.Background(), backend.StorageFactoryInput{TenantID: tenantA}, backend.CapabilityBinding{Capability: backend.CapabilitySession, Provider: "redis", Endpoint: "redis://other:6379"}, secret) + if !errors.Is(err, backend.ErrStorageFactory) || value != nil { + t.Fatalf("mismatched redis endpoint = %T, %v", value, err) + } + if strings.Contains(err.Error(), "redis://other:6379") { + t.Fatal("redis endpoint leaked in provider error") + } + value, err = provider.New(context.Background(), backend.StorageFactoryInput{TenantID: tenantA}, backend.CapabilityBinding{Capability: backend.CapabilitySession, Provider: "redis", Endpoint: config.redisEndpoint, SecretRef: "env/other"}, secret) + if !errors.Is(err, backend.ErrStorageFactory) || value != nil { + t.Fatalf("mismatched redis secret reference = %T, %v", value, err) + } +} + +func TestEnvironmentRedisProfilesUseSeparateInMemoryProvider(t *testing.T) { + const ( + tenantA = "t_00000000000000000000000000" + tenantB = "t_00000000000000000000000001" + sessionID = "same-session" + memoryID = "same-memory" + ) + config := environmentConfig{runtimeStorage: "redis", redisEndpoint: "redis://127.0.0.1:6379", redisSecretRef: "env/redis-password", redis: runtimestorageredis.Config{Password: "redis-secret"}, modelProvider: defaultModelProvider, modelNames: []string{"gpt-4o-mini"}, endpointHosts: []string{"api.openai.com"}, secretRef: "env/model", apiIdentities: map[string]gateway.APIIdentity{ + "token-a": {TenantID: tenantA, AppID: "app-a", SubjectID: "subject-a"}, + "token-b": {TenantID: tenantB, AppID: "app-b", SubjectID: "subject-b"}, + }, modelAPIKeys: map[string]string{tenantA: "model-a", tenantB: "model-b"}} + _, catalog, err := environmentCatalogs(config) + if err != nil { + t.Fatal(err) + } + redisProfile, err := backend.NewProfile(backend.CreateInput{TenantID: tenantA, ProfileKey: "redis-runtime", DisplayName: "Redis Runtime", Bindings: []backend.CapabilityBinding{ + {Capability: backend.CapabilitySession, Provider: "redis", Endpoint: config.redisEndpoint, SecretRef: config.redisSecretRef}, + {Capability: backend.CapabilityMemory, Provider: "redis", Endpoint: config.redisEndpoint, SecretRef: config.redisSecretRef}, + }}, catalog) + if err != nil { + t.Fatalf("redis profile = %v", err) + } + inMemoryProfile, err := backend.NewProfile(backend.CreateInput{TenantID: tenantB, ProfileKey: "inmemory-runtime", DisplayName: "InMemory Runtime", Bindings: []backend.CapabilityBinding{ + {Capability: backend.CapabilitySession, Provider: "inmemory"}, + {Capability: backend.CapabilityMemory, Provider: "inmemory"}, + }}, catalog) + if err != nil { + t.Fatalf("in-memory profile = %v", err) + } + + delegate := inmemory.NewSessionService() + redisStore := runtimestorageinmemory.New() + inMemoryStore := runtimestorageinmemory.New() + t.Cleanup(func() { _ = delegate.Close(); _ = redisStore.Close(); _ = inMemoryStore.Close() }) + secrets, _, providers, err := environmentRegistriesForStores(config, delegate, environmentRuntimeStores{ + primary: redisStore, + providers: map[string]runtimestorage.RuntimeStore{ + "redis": redisStore, + "inmemory": inMemoryStore, + }, + }) + if err != nil { + t.Fatal(err) + } + factory, err := backend.NewRegistryStorageFactory(providers, secrets) + if err != nil { + t.Fatal(err) + } + redisCapabilities, err := factory.New(context.Background(), backend.StorageFactoryInput{TenantID: tenantA, Bindings: redisProfile.Bindings}) + if err != nil { + t.Fatalf("redis capabilities = %v", err) + } + t.Cleanup(func() { _ = redisCapabilities.Close() }) + inMemoryCapabilities, err := factory.New(context.Background(), backend.StorageFactoryInput{TenantID: tenantB, Bindings: inMemoryProfile.Bindings}) + if err != nil { + t.Fatalf("in-memory capabilities = %v", err) + } + t.Cleanup(func() { _ = inMemoryCapabilities.Close() }) + if _, err := redisCapabilities.Session(); err != nil { + t.Fatalf("redis session capability = %v", err) + } + if _, err := inMemoryCapabilities.Session(); err != nil { + t.Fatalf("in-memory session capability = %v", err) + } + redisMemory, err := redisCapabilities.Memory() + if err != nil { + t.Fatalf("redis memory capability = %v", err) + } + inMemoryMemory, err := inMemoryCapabilities.Memory() + if err != nil { + t.Fatalf("in-memory memory capability = %v", err) + } + ctx := context.Background() + if _, err := redisMemory.PutMemory(ctx, runtimestorage.MemoryInput{TenantID: tenantA, MemoryID: memoryID, UserID: "user", Content: "redis"}); err != nil { + t.Fatalf("redis memory write = %v", err) + } + if _, err := inMemoryMemory.PutMemory(ctx, runtimestorage.MemoryInput{TenantID: tenantB, MemoryID: memoryID, UserID: "user", Content: "in-memory"}); err != nil { + t.Fatalf("in-memory memory write = %v", err) + } + if _, err := redisStore.CreateSession(ctx, tenantA, sessionID, map[string]any{"provider": "redis"}); err != nil { + t.Fatalf("redis session write = %v", err) + } + if _, err := inMemoryStore.CreateSession(ctx, tenantB, sessionID, map[string]any{"provider": "in-memory"}); err != nil { + t.Fatalf("in-memory session write = %v", err) + } + assertEnvironmentRuntimeStoreIsolation(t, redisStore, inMemoryStore, tenantA, tenantB, sessionID, memoryID) +} + +func TestEnvironmentRedisRuntimeStoresOwnPrimaryAndFallback(t *testing.T) { + primary := &environmentRuntimeStoreSpy{} + fallback := &environmentRuntimeStoreSpy{} + previousRedis := newEnvironmentRedisRuntimeStore + previousFallback := newEnvironmentInMemoryFallback + newEnvironmentRedisRuntimeStore = func(context.Context, environmentConfig) (runtimestorage.RuntimeStore, error) { + return primary, nil + } + newEnvironmentInMemoryFallback = func() runtimestorage.RuntimeStore { return fallback } + t.Cleanup(func() { + newEnvironmentRedisRuntimeStore = previousRedis + newEnvironmentInMemoryFallback = previousFallback + }) + stores, err := newEnvironmentRuntimeStoresForConfig(context.Background(), environmentConfig{runtimeStorage: "redis"}, nil) + if err != nil { + t.Fatal(err) + } + if stores.primary != primary || stores.providers["redis"] != primary || stores.providers["inmemory"] != fallback || len(stores.owned) != 2 { + t.Fatalf("redis runtime stores = %#v", stores) + } + if err := stores.Close(); err != nil || primary.closed != 1 || fallback.closed != 1 { + t.Fatalf("runtime store close = %v, primary closes=%d fallback closes=%d", err, primary.closed, fallback.closed) + } + if _, err := environmentRuntimeProviders(environmentConfig{runtimeStorage: "redis"}, environmentRuntimeStores{providers: map[string]runtimestorage.RuntimeStore{"redis": primary}}); !errors.Is(err, ErrInvalidConfig) { + t.Fatalf("missing in-memory fallback = %v", err) + } + if _, err := environmentRuntimeProviders(environmentConfig{runtimeStorage: "redis"}, environmentRuntimeStores{}); !errors.Is(err, ErrInvalidConfig) { + t.Fatalf("missing redis primary = %v", err) + } +} + +func TestEnvironmentRedisRuntimeStoreFailsClosed(t *testing.T) { + server := miniredis.RunT(t) + config := environmentConfig{redis: runtimestorageredis.Config{Addr: server.Addr()}} + store, err := environmentRedisRuntimeStore(context.Background(), config) + if err != nil { + t.Fatal(err) + } + if err := store.Close(); err != nil { + t.Fatal(err) + } + addr := server.Addr() + server.Close() + if _, err := environmentRedisRuntimeStore(context.Background(), environmentConfig{redis: runtimestorageredis.Config{Addr: addr}}); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("unavailable Redis runtime store = %v", err) + } +} + +type environmentRuntimeStoreSpy struct { + runtimestorage.RuntimeStore + closed int +} + +func (store *environmentRuntimeStoreSpy) Close() error { + store.closed++ + return nil +} + +func assertEnvironmentRuntimeStoreIsolation(t *testing.T, redisStore, inMemoryStore runtimestorage.RuntimeStore, tenantA, tenantB, sessionID, memoryID string) { + t.Helper() + ctx := context.Background() + redisSession, err := redisStore.GetSession(ctx, tenantA, sessionID) + if err != nil || redisSession.State["provider"] != "redis" { + t.Fatalf("redis session = %#v, %v", redisSession, err) + } + inMemorySession, err := inMemoryStore.GetSession(ctx, tenantB, sessionID) + if err != nil || inMemorySession.State["provider"] != "in-memory" { + t.Fatalf("in-memory session = %#v, %v", inMemorySession, err) + } + redisMemory, err := redisStore.(runtimestorage.MemoryStore).GetMemory(ctx, tenantA, memoryID) + if err != nil || redisMemory.Content != "redis" { + t.Fatalf("redis memory = %#v, %v", redisMemory, err) + } + inMemoryMemory, err := inMemoryStore.(runtimestorage.MemoryStore).GetMemory(ctx, tenantB, memoryID) + if err != nil || inMemoryMemory.Content != "in-memory" { + t.Fatalf("in-memory memory = %#v, %v", inMemoryMemory, err) + } + if _, err := redisStore.GetSession(ctx, tenantB, sessionID); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("redis session leaked into in-memory tenant = %v", err) + } + if _, err := inMemoryStore.GetSession(ctx, tenantA, sessionID); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("in-memory session leaked into redis tenant = %v", err) + } + if _, err := redisStore.(runtimestorage.MemoryStore).GetMemory(ctx, tenantB, memoryID); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("redis memory leaked into in-memory tenant = %v", err) + } + if _, err := inMemoryStore.(runtimestorage.MemoryStore).GetMemory(ctx, tenantA, memoryID); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("in-memory memory leaked into redis tenant = %v", err) + } +} + +func TestNewFromEnvironmentRedisConnectionFailureIsRedacted(t *testing.T) { + setRequiredEnvironment(t) + t.Setenv(envSessionBackend, "redis") + t.Setenv(envRedisAddr, "redis.internal:6379") + t.Setenv(envRedisPassword, "redis-password") + registerBootstrapPingDriver.Do(func() { sql.Register("trpc-service-bootstrap-ping", bootstrapPingDriver{}) }) + db, err := sql.Open("trpc-service-bootstrap-ping", "") + if err != nil { + t.Fatal(err) + } + previousOpen := openEnvironmentDatabase + previousApply := applyEnvironmentMigrations + previousVerify := verifyEnvironmentMigrations + previousRedis := newEnvironmentRedisRuntimeStore + t.Cleanup(func() { + openEnvironmentDatabase = previousOpen + applyEnvironmentMigrations = previousApply + verifyEnvironmentMigrations = previousVerify + newEnvironmentRedisRuntimeStore = previousRedis + _ = db.Close() + }) + openEnvironmentDatabase = func(context.Context, string, postgres.Options) (*sql.DB, error) { return db, nil } + applyEnvironmentMigrations = func(context.Context, *sql.DB) error { return nil } + verifyEnvironmentMigrations = func(context.Context, *sql.DB) error { return nil } + newEnvironmentRedisRuntimeStore = func(context.Context, environmentConfig) (runtimestorage.RuntimeStore, error) { + return nil, errors.New("dial redis.internal:6379 with password redis-password failed") + } + _, err = NewFromEnvironment(context.Background()) + if !errors.Is(err, ErrInvalidConfig) { + t.Fatalf("redis connection error = %v", err) + } + if strings.Contains(err.Error(), "redis.internal:6379") || strings.Contains(err.Error(), "redis-password") { + t.Fatalf("redis connection details leaked: %v", err) + } +} + +func TestNewFromEnvironmentRuntimeStoreFailureBoundaries(t *testing.T) { + tests := []struct { + name string + runtimeStorage string + configureStore func(context.CancelFunc) + wantError error + }{ + { + name: "redis cancellation wins", + runtimeStorage: "redis", + configureStore: func(cancel context.CancelFunc) { + newEnvironmentRedisRuntimeStore = func(context.Context, environmentConfig) (runtimestorage.RuntimeStore, error) { + cancel() + return nil, errors.New("redis dial failure") + } + }, + wantError: context.Canceled, + }, + { + name: "non redis preserves source error", + runtimeStorage: "inmemory", + configureStore: func(context.CancelFunc) { + newEnvironmentRuntimeStore = func(string, *sql.DB) (runtimestorage.RuntimeStore, error) { + return nil, errEnvironmentRuntimeStore + } + }, + wantError: errEnvironmentRuntimeStore, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + setRequiredEnvironment(t) + t.Setenv(envSessionBackend, tt.runtimeStorage) + if tt.runtimeStorage == "redis" { + t.Setenv(envRedisAddr, "redis.internal:6379") + } + registerBootstrapPingDriver.Do(func() { sql.Register("trpc-service-bootstrap-ping", bootstrapPingDriver{}) }) + db, err := sql.Open("trpc-service-bootstrap-ping", "") + if err != nil { + t.Fatal(err) + } + previousOpen := openEnvironmentDatabase + previousRedis := newEnvironmentRedisRuntimeStore + previousStore := newEnvironmentRuntimeStore + t.Cleanup(func() { + openEnvironmentDatabase = previousOpen + newEnvironmentRedisRuntimeStore = previousRedis + newEnvironmentRuntimeStore = previousStore + _ = db.Close() + }) + openEnvironmentDatabase = func(context.Context, string, postgres.Options) (*sql.DB, error) { return db, nil } + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + tt.configureStore(cancel) + + _, err = NewFromEnvironment(ctx) + if !errors.Is(err, tt.wantError) { + t.Fatalf("runtime store failure = %v, want %v", err, tt.wantError) + } + if pingErr := db.Ping(); pingErr == nil { + t.Fatal("database remained open after runtime store failure") + } + }) + } +} + +func TestNewFromEnvironmentRedisClosesPrimaryAndFallbackAfterBootstrapFailure(t *testing.T) { + setRequiredEnvironment(t) + t.Setenv(envSessionBackend, "redis") + t.Setenv(envRedisAddr, "redis.internal:6379") + registerBootstrapPingDriver.Do(func() { sql.Register("trpc-service-bootstrap-ping", bootstrapPingDriver{}) }) + db, err := sql.Open("trpc-service-bootstrap-ping", "") + if err != nil { + t.Fatal(err) + } + primary := &environmentRuntimeStoreSpy{} + fallback := &environmentRuntimeStoreSpy{} + previousOpen := openEnvironmentDatabase + previousApply := applyEnvironmentMigrations + previousRedis := newEnvironmentRedisRuntimeStore + previousFallback := newEnvironmentInMemoryFallback + t.Cleanup(func() { + openEnvironmentDatabase = previousOpen + applyEnvironmentMigrations = previousApply + newEnvironmentRedisRuntimeStore = previousRedis + newEnvironmentInMemoryFallback = previousFallback + _ = db.Close() + }) + openEnvironmentDatabase = func(context.Context, string, postgres.Options) (*sql.DB, error) { return db, nil } + applyEnvironmentMigrations = func(context.Context, *sql.DB) error { return errors.New("migration failed") } + newEnvironmentRedisRuntimeStore = func(context.Context, environmentConfig) (runtimestorage.RuntimeStore, error) { return primary, nil } + newEnvironmentInMemoryFallback = func() runtimestorage.RuntimeStore { return fallback } + + if _, err := NewFromEnvironment(context.Background()); !errors.Is(err, ErrInvalidConfig) { + t.Fatalf("bootstrap failure = %v", err) + } + if primary.closed != 1 || fallback.closed != 1 { + t.Fatalf("runtime store closes: primary=%d fallback=%d", primary.closed, fallback.closed) + } + if pingErr := db.Ping(); pingErr == nil { + t.Fatal("database remained open after bootstrap failure") + } +} + +var errEnvironmentRuntimeStore = errors.New("runtime store initialization failed") + func TestDemoEnvironmentConfigurationBranches(t *testing.T) { setRequiredEnvironment(t) t.Setenv(envDemoMode, "not-bool") diff --git a/trpcservice/runtime/storage/redis/integration_test.go b/trpcservice/runtime/storage/redis/integration_test.go new file mode 100644 index 00000000..b2e1c4fa --- /dev/null +++ b/trpcservice/runtime/storage/redis/integration_test.go @@ -0,0 +1,78 @@ +package redis_test + +import ( + "context" + "os" + "strconv" + "strings" + "testing" + "time" + + runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" + redisstore "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/redis" +) + +// TestRedisRuntimeLiveReconnect is opt-in so the default suite remains +// hermetic. It proves committed runtime state survives client recreation +// against an externally managed Redis service. +func TestRedisRuntimeLiveReconnect(t *testing.T) { + addr := strings.TrimSpace(os.Getenv("REDIS_RUNTIME_TEST_ADDR")) + if addr == "" { + t.Skip("REDIS_RUNTIME_TEST_ADDR is not configured") + } + db := 0 + if raw := strings.TrimSpace(os.Getenv("REDIS_RUNTIME_TEST_DB")); raw != "" { + parsed, err := strconv.Atoi(raw) + if err != nil || parsed < 0 { + t.Fatalf("invalid REDIS_RUNTIME_TEST_DB") + } + db = parsed + } + config := redisstore.Config{Addr: addr, Password: os.Getenv("REDIS_RUNTIME_TEST_PASSWORD"), DB: db, KeyPrefix: "trpc:test:redis-live"} + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + store, err := redisstore.NewFromConfig(ctx, config) + if err != nil { + t.Fatalf("open live redis: %v", err) + } + tenantID := "t_redis_live_" + strconv.FormatInt(time.Now().UnixNano(), 10) + sessionID := "session" + eventID := "event" + if _, err := store.CreateSession(ctx, tenantID, sessionID, map[string]any{"live": true}); err != nil { + _ = store.Close() + t.Fatalf("create live session: %v", err) + } + if _, _, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: tenantID, SessionID: sessionID, BindingID: "binding", ExternalMessageID: "external", EventID: eventID}); err != nil { + _ = store.Close() + t.Fatalf("record live event: %v", err) + } + if _, err := store.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: tenantID, ReplyID: "reply", EventID: eventID, SegmentIndex: 0, SegmentCount: 1, Payload: "live"}); err != nil { + _ = store.Close() + t.Fatalf("write live reply: %v", err) + } + if _, err := store.PutMemory(ctx, runtimestorage.MemoryInput{TenantID: tenantID, MemoryID: "memory", UserID: "user", Content: "live"}); err != nil { + _ = store.Close() + t.Fatalf("write live memory: %v", err) + } + if err := store.Close(); err != nil { + t.Fatalf("close live redis: %v", err) + } + + reopened, err := redisstore.NewFromConfig(ctx, config) + if err != nil { + t.Fatalf("reopen live redis: %v", err) + } + defer func() { _ = reopened.Close() }() + if _, err := reopened.GetSession(ctx, tenantID, sessionID); err != nil { + t.Fatalf("reopened session: %v", err) + } + if _, err := reopened.GetMessage(ctx, tenantID, eventID); err != nil { + t.Fatalf("reopened event: %v", err) + } + if _, err := reopened.GetMemory(ctx, tenantID, "memory"); err != nil { + t.Fatalf("reopened memory: %v", err) + } + if _, err := reopened.GetReply(ctx, tenantID, "reply", 0); err != nil { + t.Fatalf("reopened reply: %v", err) + } +} diff --git a/trpcservice/runtime/storage/redis/redis.go b/trpcservice/runtime/storage/redis/redis.go new file mode 100644 index 00000000..d2a0931d --- /dev/null +++ b/trpcservice/runtime/storage/redis/redis.go @@ -0,0 +1,1139 @@ +// Package redis provides a tenant-scoped Redis runtime and memory store. +// +// Each tenant is represented by one versioned JSON document. Mutations use +// WATCH/MULTI so event sequencing, idempotency claims, leases and memory CAS +// remain atomic when multiple workers share the same Redis server. +package redis + +import ( + "context" + "encoding/hex" + "encoding/json" + "errors" + "reflect" + "sort" + "strconv" + "strings" + "sync" + "time" + + "github.com/XnLemon/trpc-agent-service/trpcservice/observability" + runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" + "github.com/google/uuid" + redisclient "github.com/redis/go-redis/v9" +) + +const ( + stateVersion = 1 + maxWatchRetries = 8 +) + +// Config configures a Redis-backed runtime store. Password is transient +// connection input and must come from a trusted secret boundary. +type Config struct { + Addr string + Password string + DB int + KeyPrefix string + DialTimeout time.Duration + ReadTimeout time.Duration + WriteTimeout time.Duration + PoolSize int +} + +// Store implements RuntimeStore and MemoryStore over Redis. +type Store struct { + client redisclient.UniversalClient + keyPrefix string + closeOnce sync.Once + closeErr error + owned bool +} + +type state struct { + Version int `json:"version"` + Sessions map[string]runtimestorage.Session `json:"sessions,omitempty"` + Events map[string]runtimestorage.MessageEvent `json:"events,omitempty"` + Messages map[string]string `json:"messages,omitempty"` + Histories map[string][]runtimestorage.EventPayload `json:"histories,omitempty"` + Replies map[string]runtimestorage.ReplyOutbox `json:"replies,omitempty"` + Correlations map[string]runtimestorage.ReplyCorrelation `json:"correlations,omitempty"` + Memories map[string]runtimestorage.MemoryRecord `json:"memories,omitempty"` + MemoryIndexHandoffs map[string]int64 `json:"memory_index_handoffs,omitempty"` +} + +// New creates a store using a caller-owned Redis client. +func New(client redisclient.UniversalClient, keyPrefix string) (*Store, error) { + if client == nil { + return nil, runtimestorage.ErrInvalid + } + keyPrefix = strings.TrimSpace(keyPrefix) + if keyPrefix == "" { + keyPrefix = "trpc:runtime:v1" + } + return &Store{client: client, keyPrefix: keyPrefix}, nil +} + +// NewFromURL creates and pings a Redis client from a redis:// URL. The URL is +// never included in returned errors or logs. +func NewFromURL(ctx context.Context, rawURL string) (*Store, error) { + if ctx == nil { + return nil, runtimestorage.ErrInvalid + } + if err := ctx.Err(); err != nil { + return nil, err + } + options, err := redisclient.ParseURL(strings.TrimSpace(rawURL)) + if err != nil || options == nil || options.Addr == "" { + return nil, runtimestorage.ErrInvalid + } + client := redisclient.NewClient(options) + if err := client.Ping(ctx).Err(); err != nil { + _ = client.Close() + return nil, mapRedisError(ctx, err) + } + store, err := New(client, "trpc:runtime:v1") + if err != nil { + _ = client.Close() + return nil, err + } + store.owned = true + return store, nil +} + +// NewFromConfig creates and pings an owned client from explicit connection +// settings. Password is accepted only as resolved runtime input. +func NewFromConfig(ctx context.Context, config Config) (*Store, error) { + if ctx == nil || strings.TrimSpace(config.Addr) == "" || config.DB < 0 { + return nil, runtimestorage.ErrInvalid + } + if err := ctx.Err(); err != nil { + return nil, err + } + options := &redisclient.Options{Addr: strings.TrimSpace(config.Addr), Password: config.Password, DB: config.DB} + if config.DialTimeout > 0 { + options.DialTimeout = config.DialTimeout + } + if config.ReadTimeout > 0 { + options.ReadTimeout = config.ReadTimeout + } + if config.WriteTimeout > 0 { + options.WriteTimeout = config.WriteTimeout + } + if config.PoolSize > 0 { + options.PoolSize = config.PoolSize + } + client := redisclient.NewClient(options) + if err := client.Ping(ctx).Err(); err != nil { + _ = client.Close() + return nil, mapRedisError(ctx, err) + } + store, err := New(client, config.KeyPrefix) + if err != nil { + _ = client.Close() + return nil, err + } + store.owned = true + return store, nil +} + +// Ping checks readiness without exposing provider errors. +func (s *Store) Ping(ctx context.Context) error { + if err := s.check(ctx); err != nil { + return err + } + return mapRedisError(ctx, s.client.Ping(ctx).Err()) +} + +func (s *Store) check(ctx context.Context) error { + if ctx == nil { + return runtimestorage.ErrInvalid + } + if err := ctx.Err(); err != nil { + return err + } + if s == nil || s.client == nil { + return runtimestorage.ErrStorage + } + return nil +} + +func (s *Store) key(tenantID string) string { + return s.keyPrefix + ":" + hex.EncodeToString([]byte(tenantID)) +} + +func emptyState() state { + return state{Version: stateVersion, Sessions: map[string]runtimestorage.Session{}, Events: map[string]runtimestorage.MessageEvent{}, Messages: map[string]string{}, Histories: map[string][]runtimestorage.EventPayload{}, Replies: map[string]runtimestorage.ReplyOutbox{}, Correlations: map[string]runtimestorage.ReplyCorrelation{}, Memories: map[string]runtimestorage.MemoryRecord{}, MemoryIndexHandoffs: map[string]int64{}} +} + +func normalizeState(value state) state { + if value.Version == 0 { + value.Version = stateVersion + } + if value.Sessions == nil { + value.Sessions = map[string]runtimestorage.Session{} + } + if value.Events == nil { + value.Events = map[string]runtimestorage.MessageEvent{} + } + if value.Messages == nil { + value.Messages = map[string]string{} + } + if value.Histories == nil { + value.Histories = map[string][]runtimestorage.EventPayload{} + } + if value.Replies == nil { + value.Replies = map[string]runtimestorage.ReplyOutbox{} + } + if value.Correlations == nil { + value.Correlations = map[string]runtimestorage.ReplyCorrelation{} + } + if value.Memories == nil { + value.Memories = map[string]runtimestorage.MemoryRecord{} + } + if value.MemoryIndexHandoffs == nil { + value.MemoryIndexHandoffs = map[string]int64{} + } + return value +} + +func (s *Store) load(ctx context.Context, tenantID string) (state, error) { + result, err := s.client.Get(ctx, s.key(tenantID)).Bytes() + if errors.Is(err, redisclient.Nil) { + return emptyState(), nil + } + if err != nil { + return state{}, mapRedisError(ctx, err) + } + var value state + if err := json.Unmarshal(result, &value); err != nil || value.Version != stateVersion { + return state{}, runtimestorage.ErrStorage + } + return normalizeState(value), nil +} + +func (s *Store) mutate(ctx context.Context, tenantID string, fn func(*state) error) error { + if err := s.check(ctx); err != nil { + return err + } + if err := runtimestorage.ValidateTenant(tenantID); err != nil { + return err + } + key := s.key(tenantID) + for attempt := 0; attempt < maxWatchRetries; attempt++ { + if err := ctx.Err(); err != nil { + return err + } + err := s.client.Watch(ctx, func(tx *redisclient.Tx) error { + value, err := s.loadFromTx(ctx, tx, key) + if err != nil { + return err + } + if err := fn(&value); err != nil { + return err + } + encoded, err := json.Marshal(value) + if err != nil { + return runtimestorage.ErrStorage + } + _, err = tx.TxPipelined(ctx, func(pipe redisclient.Pipeliner) error { + pipe.Set(ctx, key, encoded, 0) + return nil + }) + return err + }, key) + if err == nil { + return nil + } + if errors.Is(err, redisclient.TxFailedErr) { + continue + } + return mapRedisError(ctx, err) + } + return runtimestorage.ErrConflict +} + +func (s *Store) loadFromTx(ctx context.Context, tx *redisclient.Tx, key string) (state, error) { + result, err := tx.Get(ctx, key).Bytes() + if errors.Is(err, redisclient.Nil) { + return emptyState(), nil + } + if err != nil { + return state{}, mapRedisError(ctx, err) + } + var value state + if err := json.Unmarshal(result, &value); err != nil || value.Version != stateVersion { + return state{}, runtimestorage.ErrStorage + } + return normalizeState(value), nil +} + +func mapRedisError(ctx context.Context, err error) error { + if err == nil { + return nil + } + if ctx != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + } + if errors.Is(err, context.Canceled) { + return context.Canceled + } + if errors.Is(err, context.DeadlineExceeded) { + return context.DeadlineExceeded + } + for _, sentinel := range []error{runtimestorage.ErrNotFound, runtimestorage.ErrDuplicate, runtimestorage.ErrConflict, runtimestorage.ErrInvalid, runtimestorage.ErrIllegalTransition, runtimestorage.ErrStorage} { + if errors.Is(err, sentinel) { + return err + } + } + return runtimestorage.ErrStorage +} + +func scopedKey(parts ...string) string { + var b strings.Builder + for _, part := range parts { + b.WriteString(strconv.Itoa(len(part))) + b.WriteByte(':') + b.WriteString(part) + } + return b.String() +} + +func replyKey(replyID string, segment int) string { return scopedKey(replyID, strconv.Itoa(segment)) } +func messageKey(bindingID, externalID string) string { return scopedKey(bindingID, externalID) } + +func cloneMap(input map[string]any) map[string]any { + if input == nil { + return nil + } + data, err := json.Marshal(input) + if err != nil { + return nil + } + var output map[string]any + if json.Unmarshal(data, &output) != nil { + return nil + } + return output +} + +func cloneSession(v runtimestorage.Session) runtimestorage.Session { + v.State = cloneMap(v.State) + return v +} +func cloneEvent(v runtimestorage.MessageEvent) runtimestorage.MessageEvent { + if v.LeaseExpiresAt != nil { + x := *v.LeaseExpiresAt + v.LeaseExpiresAt = &x + } + return v +} +func clonePayload(v runtimestorage.EventPayload) runtimestorage.EventPayload { + v.Payload = append([]byte(nil), v.Payload...) + return v +} +func cloneReply(v runtimestorage.ReplyOutbox) runtimestorage.ReplyOutbox { + if v.LeaseExpiresAt != nil { + x := *v.LeaseExpiresAt + v.LeaseExpiresAt = &x + } + return v +} +func cloneMemory(v runtimestorage.MemoryRecord) runtimestorage.MemoryRecord { + v.Topics = append([]string(nil), v.Topics...) + v.Metadata = cloneMap(v.Metadata) + v.Embedding = append([]float64(nil), v.Embedding...) + if v.DeletedAt != nil { + x := *v.DeletedAt + v.DeletedAt = &x + } + return v +} + +func jsonEqual(left, right []byte) bool { + var a, b any + return json.Unmarshal(left, &a) == nil && json.Unmarshal(right, &b) == nil && reflect.DeepEqual(a, b) +} + +// GetReplyCorrelation returns the tenant-scoped request correlation for an event. +func (s *Store) GetReplyCorrelation(ctx context.Context, tenantID, eventID string) (runtimestorage.ReplyCorrelation, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.ReplyCorrelation{}, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || eventID == "" { + return runtimestorage.ReplyCorrelation{}, runtimestorage.ErrInvalid + } + value, err := s.load(ctx, tenantID) + if err != nil { + return runtimestorage.ReplyCorrelation{}, err + } + result, ok := value.Correlations[eventID] + if !ok { + return runtimestorage.ReplyCorrelation{}, runtimestorage.ErrNotFound + } + result.TraceParent = observability.NormalizeTraceParent(result.TraceParent) + return result, nil +} + +// GetSession returns one tenant-scoped session. +func (s *Store) GetSession(ctx context.Context, tenantID, sessionID string) (runtimestorage.Session, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.Session{}, err + } + if runtimestorage.ValidateSession(tenantID, sessionID) != nil { + return runtimestorage.Session{}, runtimestorage.ErrInvalid + } + value, err := s.load(ctx, tenantID) + if err != nil { + return runtimestorage.Session{}, err + } + result, ok := value.Sessions[sessionID] + if !ok { + return runtimestorage.Session{}, runtimestorage.ErrNotFound + } + return cloneSession(result), nil +} + +// CreateSession creates an active tenant-scoped session. +func (s *Store) CreateSession(ctx context.Context, tenantID, sessionID string, stateValue map[string]any) (runtimestorage.Session, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.Session{}, err + } + if runtimestorage.ValidateSession(tenantID, sessionID) != nil || (stateValue != nil && cloneMap(stateValue) == nil) { + return runtimestorage.Session{}, runtimestorage.ErrInvalid + } + result := runtimestorage.Session{TenantID: tenantID, SessionID: sessionID, Status: runtimestorage.SessionActive, Version: 1, State: cloneMap(stateValue), CreatedAt: time.Now().UTC()} + result.UpdatedAt = result.CreatedAt + err := s.mutate(ctx, tenantID, func(value *state) error { + if _, ok := value.Sessions[sessionID]; ok { + return runtimestorage.ErrDuplicate + } + value.Sessions[sessionID] = result + return nil + }) + if err != nil { + return runtimestorage.Session{}, err + } + return cloneSession(result), nil +} + +// UpdateSessionState applies an optimistic-concurrency session update. +func (s *Store) UpdateSessionState(ctx context.Context, tenantID, sessionID string, expectedVersion int64, stateValue map[string]any) (runtimestorage.Session, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.Session{}, err + } + if runtimestorage.ValidateSession(tenantID, sessionID) != nil || (stateValue != nil && cloneMap(stateValue) == nil) { + return runtimestorage.Session{}, runtimestorage.ErrInvalid + } + var result runtimestorage.Session + err := s.mutate(ctx, tenantID, func(value *state) error { + current, ok := value.Sessions[sessionID] + if !ok { + return runtimestorage.ErrNotFound + } + if current.Version != expectedVersion { + return runtimestorage.ErrConflict + } + current.Version++ + current.State = cloneMap(stateValue) + current.UpdatedAt = time.Now().UTC() + value.Sessions[sessionID] = current + result = cloneSession(current) + return nil + }) + return result, err +} + +// DeleteSession removes a session and its dependent runtime records. +func (s *Store) DeleteSession(ctx context.Context, tenantID, sessionID string) error { + if err := s.check(ctx); err != nil { + return err + } + if runtimestorage.ValidateSession(tenantID, sessionID) != nil { + return runtimestorage.ErrInvalid + } + return s.mutate(ctx, tenantID, func(value *state) error { + if _, ok := value.Sessions[sessionID]; !ok { + return runtimestorage.ErrNotFound + } + delete(value.Sessions, sessionID) + delete(value.Histories, sessionID) + for id, event := range value.Events { + if event.SessionID != sessionID { + continue + } + delete(value.Events, id) + delete(value.Messages, messageKey(event.BindingID, event.ExternalMessageID)) + for key, reply := range value.Replies { + if reply.EventID == event.EventID { + delete(value.Replies, key) + } + } + } + return nil + }) +} + +// RecordMessage records or returns an idempotent inbound message event. +func (s *Store) RecordMessage(ctx context.Context, input runtimestorage.MessageEventInput) (runtimestorage.MessageEvent, bool, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.MessageEvent{}, false, err + } + if runtimestorage.ValidateSession(input.TenantID, input.SessionID) != nil || input.BindingID == "" || input.ExternalMessageID == "" || input.EventID == "" || runtimestorage.ValidateReplyTarget(input.ReplyTarget) != nil || (input.ReplyTarget != (runtimestorage.ReplyTarget{}) && input.ReplyTarget.BindingID != input.BindingID) { + return runtimestorage.MessageEvent{}, false, runtimestorage.ErrInvalid + } + var result runtimestorage.MessageEvent + duplicate := false + err := s.mutate(ctx, input.TenantID, func(value *state) error { + if id, ok := value.Messages[messageKey(input.BindingID, input.ExternalMessageID)]; ok { + result = cloneEvent(value.Events[id]) + duplicate = true + return nil + } + if _, ok := value.Events[input.EventID]; ok { + return runtimestorage.ErrDuplicate + } + sess, ok := value.Sessions[input.SessionID] + if !ok { + return runtimestorage.ErrNotFound + } + sess.Version++ + sess.UpdatedAt = time.Now().UTC() + value.Sessions[input.SessionID] = sess + now := time.Now().UTC() + result = runtimestorage.MessageEvent{TenantID: input.TenantID, EventID: input.EventID, SessionID: input.SessionID, BindingID: input.BindingID, ExternalMessageID: input.ExternalMessageID, IdempotencyKey: input.IdempotencyKey, EventSeq: sess.Version, Status: runtimestorage.EventReceived, ReplyTarget: input.ReplyTarget, CreatedAt: now, UpdatedAt: now} + value.Events[input.EventID] = result + value.Messages[messageKey(input.BindingID, input.ExternalMessageID)] = input.EventID + return nil + }) + return result, duplicate, err +} + +func validateMessageTransition(t runtimestorage.MessageTransition) error { + if runtimestorage.ValidateTenant(t.TenantID) != nil || t.EventID == "" || t.Owner == "" { + return runtimestorage.ErrInvalid + } + if !runtimestorage.ValidateMessageTransition(t.From, t.To) { + return runtimestorage.ErrIllegalTransition + } + if t.To == runtimestorage.EventRunning && t.LeaseDuration <= 0 { + return runtimestorage.ErrInvalid + } + return nil +} + +// GetMessage returns one tenant-scoped inbound message event. +func (s *Store) GetMessage(ctx context.Context, tenantID, eventID string) (runtimestorage.MessageEvent, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.MessageEvent{}, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || eventID == "" { + return runtimestorage.MessageEvent{}, runtimestorage.ErrInvalid + } + value, err := s.load(ctx, tenantID) + if err != nil { + return runtimestorage.MessageEvent{}, err + } + result, ok := value.Events[eventID] + if !ok { + return runtimestorage.MessageEvent{}, runtimestorage.ErrNotFound + } + return cloneEvent(result), nil +} + +// TransitionMessage advances an inbound event through its fenced lifecycle. +func (s *Store) TransitionMessage(ctx context.Context, transition runtimestorage.MessageTransition) (runtimestorage.MessageEvent, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.MessageEvent{}, err + } + if err := validateMessageTransition(transition); err != nil { + return runtimestorage.MessageEvent{}, err + } + var result runtimestorage.MessageEvent + err := s.mutate(ctx, transition.TenantID, func(value *state) error { + current, ok := value.Events[transition.EventID] + if !ok { + return runtimestorage.ErrNotFound + } + if current.Status != transition.From { + return runtimestorage.ErrConflict + } + now := time.Now().UTC() + if transition.From == runtimestorage.EventRunning { + if transition.To == runtimestorage.EventExecutionReconciling { + if current.LeaseExpiresAt == nil || current.LeaseExpiresAt.After(now) { + return runtimestorage.ErrConflict + } + } else if current.LeaseOwner != transition.Owner || transition.FencingToken == 0 || current.FencingToken != transition.FencingToken || current.LeaseExpiresAt == nil || !current.LeaseExpiresAt.After(now) { + return runtimestorage.ErrConflict + } + } + if transition.To == runtimestorage.EventRunning { + deadline := now.Add(transition.LeaseDuration) + current.LeaseOwner = transition.Owner + current.LeaseExpiresAt = &deadline + } else { + current.LeaseOwner = "" + current.LeaseExpiresAt = nil + } + current.Status = transition.To + if transition.ReplyID != "" { + current.ReplyID = transition.ReplyID + } + if transition.SegmentCount > 0 { + current.SegmentCount = transition.SegmentCount + } + current.FencingToken++ + current.UpdatedAt = now + value.Events[transition.EventID] = current + result = cloneEvent(current) + return nil + }) + return result, err +} + +func validatePayload(value runtimestorage.EventPayload) error { + if runtimestorage.ValidateSession(value.TenantID, value.SessionID) != nil || value.EventID == "" || len(value.Payload) == 0 || !json.Valid(value.Payload) { + return runtimestorage.ErrInvalid + } + return nil +} + +// AppendEventPayload appends or idempotently replays one session event payload. +func (s *Store) AppendEventPayload(ctx context.Context, payload runtimestorage.EventPayload) (runtimestorage.EventPayload, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.EventPayload{}, err + } + if err := validatePayload(payload); err != nil { + return runtimestorage.EventPayload{}, err + } + var result runtimestorage.EventPayload + err := s.mutate(ctx, payload.TenantID, func(value *state) error { + if _, ok := value.Sessions[payload.SessionID]; !ok { + return runtimestorage.ErrNotFound + } + entries := value.Histories[payload.SessionID] + for _, existing := range entries { + if existing.EventID != payload.EventID { + continue + } + if !jsonEqual(existing.Payload, payload.Payload) { + return runtimestorage.ErrConflict + } + result = clonePayload(existing) + return nil + } + payload.HistorySeq = int64(len(entries) + 1) + payload.CreatedAt = time.Now().UTC() + value.Histories[payload.SessionID] = append(entries, clonePayload(payload)) + result = clonePayload(payload) + return nil + }) + return result, err +} + +// ListEventPayloads returns a session's event history in sequence order. +func (s *Store) ListEventPayloads(ctx context.Context, tenantID, sessionID string) ([]runtimestorage.EventPayload, error) { + if err := s.check(ctx); err != nil { + return nil, err + } + if runtimestorage.ValidateSession(tenantID, sessionID) != nil { + return nil, runtimestorage.ErrInvalid + } + value, err := s.load(ctx, tenantID) + if err != nil { + return nil, err + } + if _, ok := value.Sessions[sessionID]; !ok { + return nil, runtimestorage.ErrNotFound + } + entries := value.Histories[sessionID] + result := make([]runtimestorage.EventPayload, len(entries)) + for i, item := range entries { + result[i] = clonePayload(item) + } + return result, 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 prepareReply(value runtimestorage.ReplyOutbox) (runtimestorage.ReplyOutbox, error) { + if value.Status == "" { + value.Status = runtimestorage.ReplyPending + } + if err := validateReplySegment(value); err != nil { + return runtimestorage.ReplyOutbox{}, err + } + if value.Status != runtimestorage.ReplyPending { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + return value, nil +} +func sameReply(a, b runtimestorage.ReplyOutbox) bool { + return a.EventID == b.EventID && a.SegmentCount == b.SegmentCount && a.Payload == b.Payload && a.ReplyTarget == b.ReplyTarget +} + +// EnqueueReply materializes one pending reply segment. +func (s *Store) EnqueueReply(ctx context.Context, input runtimestorage.ReplyOutbox) (runtimestorage.ReplyOutbox, error) { + value, err := prepareReply(input) + if err != nil { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + var result runtimestorage.ReplyOutbox + err = s.mutate(ctx, value.TenantID, func(current *state) error { + event, ok := current.Events[value.EventID] + if !ok { + return runtimestorage.ErrNotFound + } + if event.ReplyTarget != value.ReplyTarget { + return runtimestorage.ErrConflict + } + key := replyKey(value.ReplyID, value.SegmentIndex) + if existing, ok := current.Replies[key]; ok { + if !sameReply(existing, value) { + return runtimestorage.ErrConflict + } + result = cloneReply(existing) + return nil + } + now := time.Now().UTC() + value.CreatedAt, value.UpdatedAt = now, now + current.Replies[key] = value + result = cloneReply(value) + return nil + }) + return result, err +} + +func validateReplyBatch(values []runtimestorage.ReplyOutbox) (runtimestorage.ReplyOutbox, error) { + if len(values) == 0 { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + first, err := prepareReply(values[0]) + if err != nil { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + seen := map[int]bool{} + for _, raw := range values { + value, err := prepareReply(raw) + if err != nil || value.TenantID != first.TenantID || value.ReplyID != first.ReplyID || value.EventID != first.EventID || value.SegmentCount != first.SegmentCount || value.ReplyTarget != first.ReplyTarget || seen[value.SegmentIndex] { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + seen[value.SegmentIndex] = true + } + if len(seen) != first.SegmentCount { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + for i := 0; i < first.SegmentCount; i++ { + if !seen[i] { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + } + return first, nil +} + +// EnqueueReplies materializes a complete reply batch atomically. +func (s *Store) EnqueueReplies(ctx context.Context, values []runtimestorage.ReplyOutbox) ([]runtimestorage.ReplyOutbox, error) { + return s.enqueueReplies(ctx, runtimestorage.ReplyCorrelation{}, values) +} + +// EnqueueRepliesWithCorrelation atomically stores correlation and reply segments. +func (s *Store) EnqueueRepliesWithCorrelation(ctx context.Context, correlation runtimestorage.ReplyCorrelation, values []runtimestorage.ReplyOutbox) ([]runtimestorage.ReplyOutbox, error) { + if correlation.TenantID == "" || correlation.EventID == "" || correlation.RequestID == "" { + return nil, runtimestorage.ErrInvalid + } + correlation.TraceParent = observability.NormalizeTraceParent(correlation.TraceParent) + return s.enqueueReplies(ctx, correlation, values) +} +func (s *Store) enqueueReplies(ctx context.Context, correlation runtimestorage.ReplyCorrelation, values []runtimestorage.ReplyOutbox) ([]runtimestorage.ReplyOutbox, error) { + if err := s.check(ctx); err != nil { + return nil, err + } + first, err := validateReplyBatch(values) + if err != nil { + return nil, err + } + normalized := make([]runtimestorage.ReplyOutbox, len(values)) + for i, raw := range values { + var prepareErr error + normalized[i], prepareErr = prepareReply(raw) + if prepareErr != nil { + return nil, prepareErr + } + } + var result []runtimestorage.ReplyOutbox + err = s.mutate(ctx, first.TenantID, func(current *state) error { + event, ok := current.Events[first.EventID] + if !ok { + return runtimestorage.ErrNotFound + } + if event.ReplyTarget != first.ReplyTarget { + return runtimestorage.ErrConflict + } + if correlation.RequestID != "" { + if correlation.TenantID != first.TenantID || correlation.EventID != first.EventID { + return runtimestorage.ErrInvalid + } + if existing, ok := current.Correlations[first.EventID]; ok && existing != correlation { + return runtimestorage.ErrConflict + } + } + for _, value := range normalized { + if existing, ok := current.Replies[replyKey(value.ReplyID, value.SegmentIndex)]; ok && !sameReply(existing, value) { + return runtimestorage.ErrConflict + } + } + now := time.Now().UTC() + result = make([]runtimestorage.ReplyOutbox, 0, len(normalized)) + for _, value := range normalized { + key := replyKey(value.ReplyID, value.SegmentIndex) + if existing, ok := current.Replies[key]; ok { + result = append(result, cloneReply(existing)) + continue + } + value.CreatedAt, value.UpdatedAt = now, now + current.Replies[key] = value + result = append(result, cloneReply(value)) + } + if correlation.RequestID != "" { + current.Correlations[first.EventID] = correlation + } + return nil + }) + return result, err +} + +// GetReply returns one tenant-scoped reply segment. +func (s *Store) GetReply(ctx context.Context, tenantID, replyID string, segment int) (runtimestorage.ReplyOutbox, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.ReplyOutbox{}, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || replyID == "" || segment < 0 { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + value, err := s.load(ctx, tenantID) + if err != nil { + return runtimestorage.ReplyOutbox{}, err + } + result, ok := value.Replies[replyKey(replyID, segment)] + if !ok { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrNotFound + } + return cloneReply(result), nil +} + +// ListReplyCandidates returns pending, retryable, and expired sending segments. +func (s *Store) ListReplyCandidates(ctx context.Context, tenantID string) ([]runtimestorage.ReplyOutbox, error) { + if err := s.check(ctx); err != nil { + return nil, err + } + if runtimestorage.ValidateTenant(tenantID) != nil { + return nil, runtimestorage.ErrInvalid + } + value, err := s.load(ctx, tenantID) + if err != nil { + return nil, err + } + result := make([]runtimestorage.ReplyOutbox, 0, len(value.Replies)) + for _, item := range value.Replies { + result = append(result, cloneReply(item)) + } + sort.Slice(result, func(i, j int) bool { + if result[i].ReplyID == result[j].ReplyID { + return result[i].SegmentIndex < result[j].SegmentIndex + } + return result[i].ReplyID < result[j].ReplyID + }) + return result, nil +} + +// ClaimReply claims one reply segment with an expiring fencing lease. +func (s *Store) ClaimReply(ctx context.Context, tenantID, replyID string, segment int, owner string, duration time.Duration) (runtimestorage.ReplyOutbox, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.ReplyOutbox{}, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || replyID == "" || segment < 0 || owner == "" || duration <= 0 { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + var result runtimestorage.ReplyOutbox + err := s.mutate(ctx, tenantID, func(value *state) error { + current, ok := value.Replies[replyKey(replyID, segment)] + if !ok { + return runtimestorage.ErrNotFound + } + now := time.Now().UTC() + expired := current.Status == runtimestorage.ReplySending && current.LeaseExpiresAt != nil && !current.LeaseExpiresAt.After(now) + if current.Status != runtimestorage.ReplyPending && current.Status != runtimestorage.ReplyRetryable && !expired { + return runtimestorage.ErrConflict + } + deadline := now.Add(duration) + current.Status, current.Attempts, current.FencingToken, current.LeaseOwner, current.LeaseExpiresAt, current.UpdatedAt = runtimestorage.ReplySending, current.Attempts+1, current.FencingToken+1, owner, &deadline, now + value.Replies[replyKey(replyID, segment)] = current + result = cloneReply(current) + return nil + }) + return result, err +} + +// TransitionReply advances a reply segment through its fenced lifecycle. +func (s *Store) TransitionReply(ctx context.Context, transition runtimestorage.ReplyTransition) (runtimestorage.ReplyOutbox, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.ReplyOutbox{}, err + } + if runtimestorage.ValidateTenant(transition.TenantID) != nil || transition.ReplyID == "" || transition.SegmentIndex < 0 || transition.Owner == "" { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrInvalid + } + if !runtimestorage.ValidateTransition(transition.From, transition.To) { + return runtimestorage.ReplyOutbox{}, runtimestorage.ErrIllegalTransition + } + var result runtimestorage.ReplyOutbox + err := s.mutate(ctx, transition.TenantID, func(value *state) error { + key := replyKey(transition.ReplyID, transition.SegmentIndex) + current, ok := value.Replies[key] + if !ok { + return runtimestorage.ErrNotFound + } + now := time.Now().UTC() + if current.Status != transition.From || (current.LeaseOwner != "" && current.LeaseOwner != transition.Owner) || (transition.FencingToken != 0 && current.FencingToken != transition.FencingToken) || (current.Status == runtimestorage.ReplySending && (current.LeaseExpiresAt == nil || !current.LeaseExpiresAt.After(now))) { + return runtimestorage.ErrConflict + } + current.Status = transition.To + current.LeaseOwner = transition.Owner + current.FencingToken++ + if transition.To == runtimestorage.ReplySending { + current.Attempts++ + if transition.LeaseDuration > 0 { + deadline := now.Add(transition.LeaseDuration) + current.LeaseExpiresAt = &deadline + } + } else { + current.LeaseOwner = "" + current.LeaseExpiresAt = nil + } + current.ProviderMessageID, current.LastErrorClass, current.UpdatedAt = transition.ProviderID, transition.ErrorClass, now + value.Replies[key] = current + result = cloneReply(current) + return nil + }) + return result, err +} + +// PutMemory creates or updates one durable tenant-scoped memory record. +func (s *Store) PutMemory(ctx context.Context, input runtimestorage.MemoryInput) (runtimestorage.MemoryRecord, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.MemoryRecord{}, err + } + if runtimestorage.ValidateTenant(input.TenantID) != nil || !runtimestorage.ValidateText(input.UserID, 256, true) || !runtimestorage.ValidateText(input.Content, 0, true) || !runtimestorage.ValidateText(input.MemoryID, 256, false) || !runtimestorage.ValidateText(input.SessionID, 256, false) || !runtimestorage.ValidateEmbedding(input.Embedding) || (input.Metadata != nil && cloneMap(input.Metadata) == nil) { + return runtimestorage.MemoryRecord{}, runtimestorage.ErrInvalid + } + if input.MemoryID == "" { + input.MemoryID = "mem_" + uuid.NewString() + } + if input.Topics == nil { + input.Topics = []string{} + } + if input.Metadata == nil { + input.Metadata = map[string]any{} + } + if input.Embedding == nil { + input.Embedding = []float64{} + } + var result runtimestorage.MemoryRecord + err := s.mutate(ctx, input.TenantID, func(value *state) error { + now := time.Now().UTC() + current, ok := value.Memories[input.MemoryID] + if ok { + current.Content, current.Topics, current.Metadata, current.Embedding = input.Content, append([]string(nil), input.Topics...), cloneMap(input.Metadata), append([]float64(nil), input.Embedding...) + current.UserID, current.SessionID, current.Version, current.UpdatedAt, current.DeletedAt = input.UserID, input.SessionID, current.Version+1, now, nil + value.Memories[input.MemoryID] = current + result = cloneMemory(current) + return nil + } + current = runtimestorage.MemoryRecord{TenantID: input.TenantID, MemoryID: input.MemoryID, UserID: input.UserID, SessionID: input.SessionID, Content: input.Content, Topics: append([]string(nil), input.Topics...), Metadata: cloneMap(input.Metadata), Embedding: append([]float64(nil), input.Embedding...), Version: 1, CreatedAt: now, UpdatedAt: now} + value.Memories[input.MemoryID] = current + result = cloneMemory(current) + return nil + }) + return result, err +} + +// GetMemory returns one non-deleted tenant-scoped memory record. +func (s *Store) GetMemory(ctx context.Context, tenantID, memoryID string) (runtimestorage.MemoryRecord, error) { + if err := s.check(ctx); err != nil { + return runtimestorage.MemoryRecord{}, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || memoryID == "" { + return runtimestorage.MemoryRecord{}, runtimestorage.ErrInvalid + } + value, err := s.load(ctx, tenantID) + if err != nil { + return runtimestorage.MemoryRecord{}, err + } + result, ok := value.Memories[memoryID] + if !ok || result.DeletedAt != nil { + return runtimestorage.MemoryRecord{}, runtimestorage.ErrNotFound + } + return cloneMemory(result), nil +} + +// ListMemories returns a tenant user's non-deleted memories. +func (s *Store) ListMemories(ctx context.Context, tenantID, userID string, limit int) ([]runtimestorage.MemoryRecord, error) { + if err := s.check(ctx); err != nil { + return nil, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || strings.TrimSpace(userID) == "" || limit < 0 { + return nil, runtimestorage.ErrInvalid + } + value, err := s.load(ctx, tenantID) + if err != nil { + return nil, err + } + result := make([]runtimestorage.MemoryRecord, 0) + for _, item := range value.Memories { + if item.UserID == userID && item.DeletedAt == nil { + result = append(result, cloneMemory(item)) + } + } + sort.Slice(result, func(i, j int) bool { + if result[i].UpdatedAt.Equal(result[j].UpdatedAt) { + return result[i].MemoryID < result[j].MemoryID + } + return result[i].UpdatedAt.After(result[j].UpdatedAt) + }) + if limit > 0 && len(result) > limit { + result = result[:limit] + } + return result, nil +} + +// SearchMemories searches a tenant user's memory content by text terms. +func (s *Store) SearchMemories(ctx context.Context, tenantID, userID, query string, limit int) ([]runtimestorage.MemorySearchResult, error) { + if err := s.check(ctx); err != nil { + return nil, err + } + if runtimestorage.ValidateTenant(tenantID) != nil || strings.TrimSpace(userID) == "" || strings.TrimSpace(query) == "" || limit < 0 { + return nil, runtimestorage.ErrInvalid + } + terms := strings.Fields(strings.ToLower(query)) + value, err := s.load(ctx, tenantID) + if err != nil { + return nil, err + } + result := make([]runtimestorage.MemorySearchResult, 0) + for _, item := range value.Memories { + if item.UserID != userID || item.DeletedAt != nil { + continue + } + hits := 0 + text := strings.ToLower(item.Content) + for _, term := range terms { + if strings.Contains(text, term) { + hits++ + } + } + if hits > 0 { + result = append(result, runtimestorage.MemorySearchResult{Memory: cloneMemory(item), Score: float64(hits) / float64(len(terms))}) + } + } + sort.Slice(result, func(i, j int) bool { + if result[i].Score == result[j].Score { + return result[i].Memory.MemoryID < result[j].Memory.MemoryID + } + return result[i].Score > result[j].Score + }) + if limit > 0 && len(result) > limit { + result = result[:limit] + } + return result, nil +} + +// DeleteMemory tombstones one tenant-scoped memory record. +func (s *Store) DeleteMemory(ctx context.Context, tenantID, memoryID string) error { + if err := s.check(ctx); err != nil { + return err + } + if runtimestorage.ValidateTenant(tenantID) != nil || memoryID == "" { + return runtimestorage.ErrInvalid + } + return s.mutate(ctx, tenantID, func(value *state) error { + current, ok := value.Memories[memoryID] + if !ok || current.DeletedAt != nil { + return runtimestorage.ErrNotFound + } + now := time.Now().UTC() + current.DeletedAt, current.UpdatedAt, current.Version = &now, now, current.Version+1 + value.Memories[memoryID] = current + delete(value.MemoryIndexHandoffs, scopedKey(memoryID)) + return nil + }) +} + +// EnqueueMemoryIndex records a durable index handoff for a memory version. +func (s *Store) EnqueueMemoryIndex(ctx context.Context, value runtimestorage.MemoryRecord) error { + if err := s.check(ctx); err != nil { + return err + } + if runtimestorage.ValidateTenant(value.TenantID) != nil || value.MemoryID == "" || value.Version < 1 { + return runtimestorage.ErrInvalid + } + return s.mutate(ctx, value.TenantID, func(current *state) error { + stored, ok := current.Memories[value.MemoryID] + if !ok || stored.DeletedAt != nil { + return runtimestorage.ErrNotFound + } + if stored.Version != value.Version { + return runtimestorage.ErrConflict + } + current.MemoryIndexHandoffs[scopedKey(value.MemoryID)] = value.Version + return nil + }) +} + +// WaitForMemoryIndex verifies that a memory version has been handed off. +func (s *Store) WaitForMemoryIndex(ctx context.Context, tenantID, memoryID string, version int64) error { + if err := s.check(ctx); err != nil { + return err + } + if runtimestorage.ValidateTenant(tenantID) != nil || memoryID == "" || version < 1 { + return runtimestorage.ErrInvalid + } + value, err := s.load(ctx, tenantID) + if err != nil { + return err + } + current, ok := value.Memories[memoryID] + if !ok || current.DeletedAt != nil { + return runtimestorage.ErrNotFound + } + if current.Version < version || value.MemoryIndexHandoffs[scopedKey(memoryID)] < version { + return runtimestorage.ErrConflict + } + return nil +} + +// Close is idempotent. A caller-owned client remains open. +func (s *Store) Close() error { + if s == nil { + return nil + } + s.closeOnce.Do(func() { + if s.owned { + if closer, ok := s.client.(interface{ Close() error }); ok { + s.closeErr = mapRedisError(context.Background(), closer.Close()) + } + } + }) + return s.closeErr +} + +var _ runtimestorage.RuntimeStore = (*Store)(nil) +var _ runtimestorage.MemoryStore = (*Store)(nil) +var _ runtimestorage.ReplyBatchEnqueuer = (*Store)(nil) +var _ runtimestorage.ReplyBatchCorrelationEnqueuer = (*Store)(nil) +var _ runtimestorage.ReplyCorrelationStore = (*Store)(nil) diff --git a/trpcservice/runtime/storage/redis/redis_test.go b/trpcservice/runtime/storage/redis/redis_test.go new file mode 100644 index 00000000..0dfe36e7 --- /dev/null +++ b/trpcservice/runtime/storage/redis/redis_test.go @@ -0,0 +1,966 @@ +package redis_test + +import ( + "context" + "encoding/hex" + "errors" + "math" + "strings" + "sync" + "testing" + "time" + + runtimestorage "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage" + redisstore "github.com/XnLemon/trpc-agent-service/trpcservice/runtime/storage/redis" + "github.com/alicebob/miniredis/v2" + redisclient "github.com/redis/go-redis/v9" +) + +func newStore(t *testing.T, server *miniredis.Miniredis) *redisstore.Store { + t.Helper() + client := redisclient.NewClient(&redisclient.Options{Addr: server.Addr()}) + store, err := redisstore.New(client, "test:runtime:v1") + if err != nil { + _ = client.Close() + t.Fatal(err) + } + t.Cleanup(func() { + _ = store.Close() + _ = client.Close() + }) + return store +} + +func seedEvent(t *testing.T, store *redisstore.Store, tenantID, sessionID, eventID string) runtimestorage.MessageEvent { + t.Helper() + if _, err := store.CreateSession(context.Background(), tenantID, sessionID, nil); err != nil { + t.Fatal(err) + } + event, duplicate, err := store.RecordMessage(context.Background(), runtimestorage.MessageEventInput{ + TenantID: tenantID, SessionID: sessionID, BindingID: "binding-" + eventID, + ExternalMessageID: "external-" + eventID, EventID: eventID, + }) + if err != nil || duplicate { + t.Fatalf("seed event = %+v duplicate=%v err=%v", event, duplicate, err) + } + return event +} + +func TestRedisTenantIsolationAndReconnect(t *testing.T) { + server := miniredis.RunT(t) + first := newStore(t, server) + second := newStore(t, server) + ctx := context.Background() + for _, tenantID := range []string{"tenant-a", "tenant-b"} { + if _, err := first.CreateSession(ctx, tenantID, "same-session", map[string]any{"tenant": tenantID}); err != nil { + t.Fatal(err) + } + event, _, err := first.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: tenantID, SessionID: "same-session", BindingID: "same-binding", ExternalMessageID: "same-external", EventID: "same-event"}) + if err != nil { + t.Fatal(err) + } + if _, err := first.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: tenantID, ReplyID: "same-reply", EventID: event.EventID, SegmentIndex: 0, SegmentCount: 1, Payload: tenantID}); err != nil { + t.Fatal(err) + } + if _, err := first.PutMemory(ctx, runtimestorage.MemoryInput{TenantID: tenantID, MemoryID: "same-memory", UserID: "user", Content: tenantID}); err != nil { + t.Fatal(err) + } + } + if got, err := second.GetSession(ctx, "tenant-a", "same-session"); err != nil || got.State["tenant"] != "tenant-a" { + t.Fatalf("tenant-a session = %+v, %v", got, err) + } + if got, err := second.GetSession(ctx, "tenant-b", "same-session"); err != nil || got.State["tenant"] != "tenant-b" { + t.Fatalf("tenant-b session = %+v, %v", got, err) + } + if got, err := second.GetReply(ctx, "tenant-b", "same-reply", 0); err != nil || got.Payload != "tenant-b" { + t.Fatalf("tenant-b reply = %+v, %v", got, err) + } + if _, err := second.GetMemory(ctx, "tenant-a", "same-memory"); err != nil { + t.Fatal(err) + } + + // A fresh client sees committed state after the original store is closed. + if err := first.Close(); err != nil { + t.Fatal(err) + } + reopened := newStore(t, server) + if got, err := reopened.GetMessage(ctx, "tenant-a", "same-event"); err != nil || got.SessionID != "same-session" { + t.Fatalf("reopened event = %+v, %v", got, err) + } +} + +func TestRedisDuplicateDeliveryAndConcurrentSequence(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + if _, err := store.CreateSession(ctx, "tenant-a", "session", nil); err != nil { + t.Fatal(err) + } + input := runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session", BindingID: "binding", ExternalMessageID: "external", EventID: "event-1"} + first, duplicate, err := store.RecordMessage(ctx, input) + if err != nil || duplicate { + t.Fatalf("first = %+v duplicate=%v err=%v", first, duplicate, err) + } + input.EventID = "event-duplicate" + replayed, duplicate, err := store.RecordMessage(ctx, input) + if err != nil || !duplicate || replayed.EventID != first.EventID { + t.Fatalf("replayed = %+v duplicate=%v err=%v", replayed, duplicate, err) + } + + var wg sync.WaitGroup + results := make(chan int64, 2) + for _, eventID := range []string{"event-2", "event-3"} { + wg.Add(1) + go func(id string) { + defer wg.Done() + value, _, callErr := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session", BindingID: id, ExternalMessageID: id, EventID: id}) + if callErr == nil { + results <- value.EventSeq + } + }(eventID) + } + wg.Wait() + close(results) + var sequences []int64 + for value := range results { + sequences = append(sequences, value) + } + if len(sequences) != 2 || sequences[0] == sequences[1] { + t.Fatalf("concurrent sequences = %v", sequences) + } +} + +func TestRedisHistoryOrderingAndDefensiveCopies(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + seedEvent(t, store, "tenant-a", "session", "inbound") + payload := []byte(`{"id":1}`) + first, err := store.AppendEventPayload(context.Background(), runtimestorage.EventPayload{TenantID: "tenant-a", SessionID: "session", EventID: "runner-1", Payload: payload}) + if err != nil || first.HistorySeq != 1 { + t.Fatalf("first history = %+v, %v", first, err) + } + first.Payload[0] = 'x' + if _, err := store.AppendEventPayload(context.Background(), runtimestorage.EventPayload{TenantID: "tenant-a", SessionID: "session", EventID: "runner-1", Payload: []byte(`{"id":1}`)}); err != nil { + t.Fatal(err) + } + if _, err := store.AppendEventPayload(context.Background(), runtimestorage.EventPayload{TenantID: "tenant-a", SessionID: "session", EventID: "runner-1", Payload: []byte(`{"id":2}`)}); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("history conflict = %v", err) + } + second, err := store.AppendEventPayload(context.Background(), runtimestorage.EventPayload{TenantID: "tenant-a", SessionID: "session", EventID: "runner-2", Payload: []byte(`{"id":2}`)}) + if err != nil || second.HistorySeq != 2 { + t.Fatalf("second history = %+v, %v", second, err) + } + items, err := store.ListEventPayloads(context.Background(), "tenant-a", "session") + if err != nil || len(items) != 2 || items[0].HistorySeq != 1 || items[1].HistorySeq != 2 || string(items[0].Payload) != `{"id":1}` { + t.Fatalf("history = %+v, %v", items, err) + } +} + +func TestRedisMessageLeaseExpiryAndFencing(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + event := seedEvent(t, store, "tenant-a", "session", "event") + running, err := store.TransitionMessage(context.Background(), runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: event.EventID, From: runtimestorage.EventReceived, To: runtimestorage.EventRunning, Owner: "worker-a", LeaseDuration: 15 * time.Millisecond}) + if err != nil { + t.Fatal(err) + } + time.Sleep(30 * time.Millisecond) + reconciling, err := store.TransitionMessage(context.Background(), runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: event.EventID, From: runtimestorage.EventRunning, To: runtimestorage.EventExecutionReconciling, Owner: "worker-b"}) + if err != nil || reconciling.FencingToken != running.FencingToken+1 || reconciling.LeaseOwner != "" { + t.Fatalf("reconciling = %+v, %v", reconciling, err) + } + if _, err := store.TransitionMessage(context.Background(), runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: event.EventID, From: runtimestorage.EventExecutionReconciling, To: runtimestorage.EventRunning, Owner: "worker-b", LeaseDuration: time.Second}); err != nil { + t.Fatal(err) + } +} + +func TestRedisReplyBatchRetryAndDeadLetter(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + event := seedEvent(t, store, "tenant-a", "session", "event") + batch := []runtimestorage.ReplyOutbox{ + {TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentIndex: 0, SegmentCount: 2, Payload: "one"}, + {TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentIndex: 1, SegmentCount: 2, Payload: "two"}, + } + rows, err := store.EnqueueReplies(context.Background(), batch) + if err != nil || len(rows) != 2 || rows[0].Status != runtimestorage.ReplyPending { + t.Fatalf("batch = %+v, %v", rows, err) + } + if _, err := store.EnqueueReplies(context.Background(), []runtimestorage.ReplyOutbox{batch[0], {TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentIndex: 1, SegmentCount: 2, Payload: "changed"}}); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("batch conflict = %v", err) + } + claimed, err := store.ClaimReply(context.Background(), "tenant-a", "reply", 0, "worker-a", time.Second) + if err != nil { + t.Fatal(err) + } + if _, err := store.TransitionReply(context.Background(), runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply", SegmentIndex: 0, From: runtimestorage.ReplySending, To: runtimestorage.ReplyRetryable, Owner: "worker-b", FencingToken: claimed.FencingToken}); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("stale reply worker = %v", err) + } + if _, err := store.TransitionReply(context.Background(), runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply", SegmentIndex: 0, From: runtimestorage.ReplySending, To: runtimestorage.ReplyRetryable, Owner: "worker-a", FencingToken: claimed.FencingToken, ErrorClass: "rate_limited"}); err != nil { + t.Fatal(err) + } + retry, err := store.ClaimReply(context.Background(), "tenant-a", "reply", 0, "worker-b", time.Second) + if err != nil { + t.Fatal(err) + } + dead, err := store.TransitionReply(context.Background(), runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply", SegmentIndex: 0, From: runtimestorage.ReplySending, To: runtimestorage.ReplyDeadLetter, Owner: "worker-b", FencingToken: retry.FencingToken, ErrorClass: "permanent"}) + if err != nil || dead.Status != runtimestorage.ReplyDeadLetter || dead.LeaseOwner != "" || dead.LeaseExpiresAt != nil { + t.Fatalf("dead letter = %+v, %v", dead, err) + } +} + +func TestRedisMemoryDurabilityAndIndexHandoff(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + value, err := store.PutMemory(context.Background(), runtimestorage.MemoryInput{TenantID: "tenant-a", MemoryID: "memory", UserID: "user", Content: "likes coffee", Metadata: map[string]any{"kind": "fact"}, Embedding: []float64{1, 0}}) + if err != nil { + t.Fatal(err) + } + value.Metadata["kind"] = "changed" + if err := store.EnqueueMemoryIndex(context.Background(), value); err != nil { + t.Fatal(err) + } + if err := store.WaitForMemoryIndex(context.Background(), "tenant-a", "memory", value.Version); err != nil { + t.Fatalf("index handoff = %v", err) + } + value2, err := store.PutMemory(context.Background(), runtimestorage.MemoryInput{TenantID: "tenant-a", MemoryID: "memory", UserID: "user", Content: "likes tea"}) + if err != nil || value2.Version != 2 { + t.Fatalf("memory update = %+v, %v", value2, err) + } + if err := store.WaitForMemoryIndex(context.Background(), "tenant-a", "memory", value2.Version); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("missing new handoff = %v", err) + } + if err := store.EnqueueMemoryIndex(context.Background(), value2); err != nil { + t.Fatal(err) + } + if err := store.WaitForMemoryIndex(context.Background(), "tenant-a", "memory", value2.Version); err != nil { + t.Fatal(err) + } + reopened := newStore(t, server) + got, err := reopened.GetMemory(context.Background(), "tenant-a", "memory") + if err != nil || got.Content != "likes tea" || got.Metadata == nil { + t.Fatalf("reopened memory = %+v, %v", got, err) + } +} + +func TestRedisCancellationCloseOwnershipAndRedaction(t *testing.T) { + server := miniredis.RunT(t) + client := redisclient.NewClient(&redisclient.Options{Addr: server.Addr()}) + borrowed, err := redisstore.New(client, "test:runtime:v1") + if err != nil { + t.Fatal(err) + } + if err := borrowed.Close(); err != nil || client.Ping(context.Background()).Err() != nil { + t.Fatalf("borrowed close = %v", err) + } + canceled, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := borrowed.GetSession(canceled, "tenant-a", "session"); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled read = %v", err) + } + addr := server.Addr() + server.Close() + if _, err := redisstore.NewFromConfig(context.Background(), redisstore.Config{Addr: addr, Password: "super-secret"}); err == nil || strings.Contains(err.Error(), addr) || strings.Contains(err.Error(), "super-secret") { + t.Fatalf("redacted unavailable error = %v", err) + } + _ = client.Close() +} + +func TestRedisScopedKeyCannotCollide(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + event := seedEvent(t, store, "tenant-a", "session", "event") + for _, reply := range []runtimestorage.ReplyOutbox{ + {TenantID: "tenant-a", ReplyID: "ab", EventID: event.EventID, SegmentIndex: 0, SegmentCount: 1, Payload: "first"}, + {TenantID: "tenant-a", ReplyID: "a", EventID: event.EventID, SegmentIndex: 0, SegmentCount: 1, Payload: "second"}, + } { + if _, err := store.EnqueueReply(context.Background(), reply); err != nil { + t.Fatal(err) + } + } + for replyID, payload := range map[string]string{"ab": "first", "a": "second"} { + got, err := store.GetReply(context.Background(), "tenant-a", replyID, 0) + if err != nil || got.Payload != payload { + t.Fatalf("reply %q = %+v, %v", replyID, got, err) + } + } +} + +func TestRedisSessionCASAndDeleteCascade(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + created, err := store.CreateSession(ctx, "tenant-a", "session", map[string]any{"nested": map[string]any{"value": "one"}}) + if err != nil { + t.Fatal(err) + } + if _, err := store.CreateSession(ctx, "tenant-a", "session", nil); !errors.Is(err, runtimestorage.ErrDuplicate) { + t.Fatalf("duplicate session = %v", err) + } + if _, err := store.UpdateSessionState(ctx, "tenant-a", "session", created.Version+1, nil); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("stale session update = %v", err) + } + updated, err := store.UpdateSessionState(ctx, "tenant-a", "session", created.Version, map[string]any{"nested": map[string]any{"value": "two"}}) + if err != nil || updated.Version != created.Version+1 { + t.Fatalf("session update = %+v, %v", updated, err) + } + updated.State["nested"].(map[string]any)["value"] = "mutated" + got, err := store.GetSession(ctx, "tenant-a", "session") + if err != nil || got.State["nested"].(map[string]any)["value"] != "two" { + t.Fatalf("session copy = %+v, %v", got, err) + } + + event, duplicate, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session", BindingID: "binding", ExternalMessageID: "external", EventID: "event"}) + if err != nil || duplicate { + t.Fatalf("event = %+v duplicate=%v err=%v", event, duplicate, err) + } + if _, err := store.AppendEventPayload(ctx, runtimestorage.EventPayload{TenantID: "tenant-a", SessionID: "session", EventID: "runner", Payload: []byte(`{"kind":"event"}`)}); err != nil { + t.Fatal(err) + } + if _, err := store.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentCount: 1, Payload: "payload"}); err != nil { + t.Fatal(err) + } + if err := store.DeleteSession(ctx, "tenant-a", "session"); err != nil { + t.Fatal(err) + } + for name, check := range map[string]func() error{ + "session": func() error { _, err := store.GetSession(ctx, "tenant-a", "session"); return err }, + "event": func() error { _, err := store.GetMessage(ctx, "tenant-a", event.EventID); return err }, + "reply": func() error { _, err := store.GetReply(ctx, "tenant-a", "reply", 0); return err }, + "history": func() error { _, err := store.ListEventPayloads(ctx, "tenant-a", "session"); return err }, + } { + if err := check(); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("deleted %s = %v", name, err) + } + } + if _, err := store.CreateSession(ctx, "tenant-a", "session", nil); err != nil { + t.Fatal(err) + } + if _, duplicate, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session", BindingID: "binding", ExternalMessageID: "external", EventID: "event"}); err != nil || duplicate { + t.Fatalf("reused inbound identity duplicate=%v err=%v", duplicate, err) + } +} + +func TestRedisReplyCorrelationAndCandidateOrdering(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + const traceParent = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01" + target := runtimestorage.ReplyTarget{BindingID: "binding", ConversationKind: "direct", ReceiverID: "receiver"} + if _, err := store.CreateSession(ctx, "tenant-a", "session", nil); err != nil { + t.Fatal(err) + } + event, duplicate, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session", BindingID: "binding", ExternalMessageID: "external", EventID: "event", ReplyTarget: target}) + if err != nil || duplicate { + t.Fatalf("event = %+v duplicate=%v err=%v", event, duplicate, err) + } + correlation := runtimestorage.ReplyCorrelation{TenantID: "tenant-a", EventID: event.EventID, RequestID: "request", TraceID: "trace", TraceParent: traceParent} + batch := []runtimestorage.ReplyOutbox{ + {TenantID: "tenant-a", ReplyID: "reply-b", EventID: event.EventID, SegmentIndex: 1, SegmentCount: 2, Payload: "two", ReplyTarget: target}, + {TenantID: "tenant-a", ReplyID: "reply-b", EventID: event.EventID, SegmentIndex: 0, SegmentCount: 2, Payload: "one", ReplyTarget: target}, + } + if _, err := store.EnqueueRepliesWithCorrelation(ctx, correlation, batch); err != nil { + t.Fatal(err) + } + if _, err := store.EnqueueRepliesWithCorrelation(ctx, correlation, batch); err != nil { + t.Fatalf("idempotent correlated enqueue = %v", err) + } + if _, err := store.EnqueueRepliesWithCorrelation(ctx, runtimestorage.ReplyCorrelation{TenantID: "tenant-a", EventID: event.EventID, RequestID: "other"}, batch); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("correlation conflict = %v", err) + } + got, err := store.GetReplyCorrelation(ctx, "tenant-a", event.EventID) + if err != nil || got != correlation { + t.Fatalf("correlation = %+v, %v", got, err) + } + if _, err := store.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply-a", EventID: event.EventID, SegmentCount: 1, Payload: "first", ReplyTarget: target}); err != nil { + t.Fatal(err) + } + candidates, err := store.ListReplyCandidates(ctx, "tenant-a") + if err != nil || len(candidates) != 3 || candidates[0].ReplyID != "reply-a" || candidates[1].SegmentIndex != 0 || candidates[2].SegmentIndex != 1 { + t.Fatalf("sorted candidates = %+v, %v", candidates, err) + } + if _, err := store.GetReplyCorrelation(ctx, "tenant-a", "missing"); !errors.Is(err, runtimestorage.ErrNotFound) { + t.Fatalf("missing correlation = %v", err) + } +} + +func seedRedisMemories(t *testing.T, store *redisstore.Store) { + t.Helper() + for _, value := range []runtimestorage.MemoryInput{ + {TenantID: "tenant-a", MemoryID: "both", UserID: "user", Content: "coffee and tea", Topics: []string{"drink"}, Metadata: map[string]any{"source": "test"}, Embedding: []float64{1, 0}}, + {TenantID: "tenant-a", MemoryID: "coffee", UserID: "user", Content: "coffee", Embedding: []float64{0, 1}}, + {TenantID: "tenant-a", MemoryID: "other-user", UserID: "other", Content: "coffee and tea"}, + } { + if _, err := store.PutMemory(context.Background(), value); err != nil { + t.Fatal(err) + } + } +} + +func TestRedisMemoryQueriesAndDefensiveCopies(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + seedRedisMemories(t, store) + values, err := store.ListMemories(ctx, "tenant-a", "user", 1) + if err != nil || len(values) != 1 || values[0].UserID != "user" { + t.Fatalf("limited memories = %+v, %v", values, err) + } + results, err := store.SearchMemories(ctx, "tenant-a", "user", "coffee tea", 10) + if err != nil || len(results) != 2 || results[0].Memory.MemoryID != "both" || results[0].Score != 1 || results[1].Memory.MemoryID != "coffee" || results[1].Score != 0.5 { + t.Fatalf("memory search = %+v, %v", results, err) + } + record, err := store.GetMemory(ctx, "tenant-a", "both") + if err != nil { + t.Fatal(err) + } + record.Topics[0] = "changed" + record.Metadata["source"] = "changed" + record.Embedding[0] = 9 + persisted, err := store.GetMemory(ctx, "tenant-a", "both") + if err != nil || persisted.Topics[0] != "drink" || persisted.Metadata["source"] != "test" || persisted.Embedding[0] != 1 { + t.Fatalf("memory defensive copy = %+v, %v", persisted, err) + } +} + +func TestRedisMemoryTombstoneRemovesIndexHandoff(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + seedRedisMemories(t, store) + persisted, err := store.GetMemory(ctx, "tenant-a", "both") + if err != nil { + t.Fatal(err) + } + if err := store.EnqueueMemoryIndex(ctx, persisted); err != nil { + t.Fatal(err) + } + if err := store.DeleteMemory(ctx, "tenant-a", "both"); err != nil { + t.Fatal(err) + } + checks := []struct { + name string + call func() error + }{ + {name: "read", call: func() error { _, err := store.GetMemory(ctx, "tenant-a", "both"); return err }}, + {name: "index", call: func() error { return store.EnqueueMemoryIndex(ctx, persisted) }}, + {name: "wait", call: func() error { return store.WaitForMemoryIndex(ctx, "tenant-a", "both", persisted.Version) }}, + } + for _, check := range checks { + assertRedisError(t, check.name, check.call(), runtimestorage.ErrNotFound) + } + values, err := store.ListMemories(ctx, "tenant-a", "user", 10) + if err != nil || len(values) != 1 || values[0].MemoryID != "coffee" { + t.Fatalf("tombstoned memories = %+v, %v", values, err) + } +} + +func TestRedisOwnedConstructorsAndPing(t *testing.T) { + server := miniredis.RunT(t) + ctx := context.Background() + if _, err := redisstore.New(nil, ""); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("nil client = %v", err) + } + if _, err := redisstore.NewFromURL(ctx, "not a redis URL"); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("invalid url = %v", err) + } + fromURL, err := redisstore.NewFromURL(ctx, "redis://"+server.Addr()+"/2") + if err != nil { + t.Fatal(err) + } + if err := fromURL.Ping(ctx); err != nil { + t.Fatal(err) + } + if err := fromURL.Close(); err != nil { + t.Fatal(err) + } + if err := fromURL.Ping(ctx); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("closed url client ping = %v", err) + } + fromConfig, err := redisstore.NewFromConfig(ctx, redisstore.Config{Addr: server.Addr(), DB: 1, KeyPrefix: "custom", PoolSize: 2}) + if err != nil { + t.Fatal(err) + } + if err := fromConfig.Ping(ctx); err != nil { + t.Fatal(err) + } + if err := fromConfig.Close(); err != nil { + t.Fatal(err) + } +} + +func assertRedisError(t *testing.T, name string, err, want error) { + t.Helper() + if !errors.Is(err, want) { + t.Fatalf("%s = %v, want %v", name, err, want) + } +} + +func TestRedisConstructorAndContextValidation(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + var nilStore *redisstore.Store + assertRedisError(t, "nil store ping", nilStore.Ping(ctx), runtimestorage.ErrStorage) + assertRedisError(t, "nil context ping", store.Ping(nil), runtimestorage.ErrInvalid) + _, err := store.GetSession(nil, "tenant-a", "session") + assertRedisError(t, "nil context read", err, runtimestorage.ErrInvalid) + canceled, cancel := context.WithCancel(ctx) + cancel() + _, err = redisstore.NewFromURL(canceled, "redis://"+server.Addr()) + assertRedisError(t, "canceled URL constructor", err, context.Canceled) + _, err = redisstore.NewFromConfig(canceled, redisstore.Config{Addr: server.Addr()}) + assertRedisError(t, "canceled config constructor", err, context.Canceled) + checks := []struct { + name string + call func() error + }{ + {name: "empty config address", call: func() error { _, err := redisstore.NewFromConfig(ctx, redisstore.Config{}); return err }}, + {name: "negative config DB", call: func() error { + _, err := redisstore.NewFromConfig(ctx, redisstore.Config{Addr: server.Addr(), DB: -1}) + return err + }}, + {name: "nil URL context", call: func() error { _, err := redisstore.NewFromURL(nil, "redis://"+server.Addr()); return err }}, + {name: "empty URL", call: func() error { _, err := redisstore.NewFromURL(ctx, ""); return err }}, + } + for _, check := range checks { + assertRedisError(t, check.name, check.call(), runtimestorage.ErrInvalid) + } +} + +func TestRedisSessionAndEventValidation(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + checks := []struct { + name string + call func() error + }{ + {name: "invalid correlation read", call: func() error { _, err := store.GetReplyCorrelation(ctx, "", "event"); return err }}, + {name: "empty correlation event", call: func() error { _, err := store.GetReplyCorrelation(ctx, "tenant-a", ""); return err }}, + {name: "invalid session tenant", call: func() error { _, err := store.GetSession(ctx, "", "session"); return err }}, + {name: "empty session ID", call: func() error { _, err := store.CreateSession(ctx, "tenant-a", "", nil); return err }}, + {name: "unencodable session state", call: func() error { + _, err := store.CreateSession(ctx, "tenant-a", "bad-state", map[string]any{"channel": make(chan int)}) + return err + }}, + {name: "unencodable update state", call: func() error { + _, err := store.UpdateSessionState(ctx, "tenant-a", "missing", 1, map[string]any{"channel": make(chan int)}) + return err + }}, + {name: "invalid session delete", call: func() error { return store.DeleteSession(ctx, "", "session") }}, + {name: "incomplete inbound event", call: func() error { + _, _, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session"}) + return err + }}, + {name: "mismatched reply target", call: func() error { + _, _, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session", BindingID: "binding", ExternalMessageID: "external", EventID: "event", ReplyTarget: runtimestorage.ReplyTarget{BindingID: "other", ConversationKind: "direct", ReceiverID: "receiver"}}) + return err + }}, + {name: "invalid message read", call: func() error { _, err := store.GetMessage(ctx, "", "event"); return err }}, + {name: "invalid history payload", call: func() error { + _, err := store.AppendEventPayload(ctx, runtimestorage.EventPayload{TenantID: "tenant-a", SessionID: "session", EventID: "event", Payload: []byte("not-json")}) + return err + }}, + {name: "invalid history tenant", call: func() error { _, err := store.ListEventPayloads(ctx, "", "session"); return err }}, + } + for _, check := range checks { + assertRedisError(t, check.name, check.call(), runtimestorage.ErrInvalid) + } + transitions := []struct { + name string + value runtimestorage.MessageTransition + want error + }{ + {name: "missing owner", value: runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: "event", From: runtimestorage.EventReceived, To: runtimestorage.EventFailed}, want: runtimestorage.ErrInvalid}, + {name: "illegal transition", value: runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: "event", From: runtimestorage.EventReceived, To: runtimestorage.EventCompleted, Owner: "worker"}, want: runtimestorage.ErrIllegalTransition}, + {name: "running without lease", value: runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: "event", From: runtimestorage.EventReceived, To: runtimestorage.EventRunning, Owner: "worker"}, want: runtimestorage.ErrInvalid}, + } + for _, transition := range transitions { + _, err := store.TransitionMessage(ctx, transition.value) + assertRedisError(t, transition.name, err, transition.want) + } +} + +func TestRedisReplyValidation(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + checks := []struct { + name string + call func() error + want error + }{ + {name: "non-pending reply", call: func() error { + _, err := store.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: "event", SegmentCount: 1, Status: runtimestorage.ReplySent}) + return err + }, want: runtimestorage.ErrInvalid}, + {name: "out-of-range reply segment", call: func() error { + _, err := store.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: "event", SegmentIndex: 1, SegmentCount: 1}) + return err + }, want: runtimestorage.ErrInvalid}, + {name: "incomplete reply batch", call: func() error { + _, err := store.EnqueueReplies(ctx, []runtimestorage.ReplyOutbox{{TenantID: "tenant-a", ReplyID: "reply", EventID: "event", SegmentCount: 2}}) + return err + }, want: runtimestorage.ErrInvalid}, + {name: "empty correlation enqueue", call: func() error { + _, err := store.EnqueueRepliesWithCorrelation(ctx, runtimestorage.ReplyCorrelation{}, nil) + return err + }, want: runtimestorage.ErrInvalid}, + {name: "empty reply ID", call: func() error { _, err := store.GetReply(ctx, "tenant-a", "", 0); return err }, want: runtimestorage.ErrInvalid}, + {name: "invalid candidates tenant", call: func() error { _, err := store.ListReplyCandidates(ctx, ""); return err }, want: runtimestorage.ErrInvalid}, + {name: "empty reply owner", call: func() error { _, err := store.ClaimReply(ctx, "tenant-a", "reply", 0, "", time.Second); return err }, want: runtimestorage.ErrInvalid}, + {name: "illegal reply transition", call: func() error { + _, err := store.TransitionReply(ctx, runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply", From: runtimestorage.ReplySent, To: runtimestorage.ReplySending, Owner: "worker"}) + return err + }, want: runtimestorage.ErrIllegalTransition}, + } + for _, check := range checks { + assertRedisError(t, check.name, check.call(), check.want) + } +} + +func TestRedisMemoryValidation(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + checks := []struct { + name string + call func() error + }{ + {name: "empty memory user", call: func() error { + _, err := store.PutMemory(ctx, runtimestorage.MemoryInput{TenantID: "tenant-a", Content: "content"}) + return err + }}, + {name: "non-finite memory embedding", call: func() error { + _, err := store.PutMemory(ctx, runtimestorage.MemoryInput{TenantID: "tenant-a", UserID: "user", Content: "content", Embedding: []float64{math.NaN()}}) + return err + }}, + {name: "unencodable memory metadata", call: func() error { + _, err := store.PutMemory(ctx, runtimestorage.MemoryInput{TenantID: "tenant-a", UserID: "user", Content: "content", Metadata: map[string]any{"channel": make(chan int)}}) + return err + }}, + {name: "invalid memory read", call: func() error { _, err := store.GetMemory(ctx, "", "memory"); return err }}, + {name: "empty memory list user", call: func() error { _, err := store.ListMemories(ctx, "tenant-a", "", 0); return err }}, + {name: "negative memory limit", call: func() error { _, err := store.ListMemories(ctx, "tenant-a", "user", -1); return err }}, + {name: "empty memory query", call: func() error { _, err := store.SearchMemories(ctx, "tenant-a", "user", "", 0); return err }}, + {name: "empty memory delete", call: func() error { return store.DeleteMemory(ctx, "tenant-a", "") }}, + {name: "invalid memory index", call: func() error { + return store.EnqueueMemoryIndex(ctx, runtimestorage.MemoryRecord{TenantID: "tenant-a", MemoryID: "memory"}) + }}, + {name: "invalid memory handoff wait", call: func() error { return store.WaitForMemoryIndex(ctx, "tenant-a", "memory", 0) }}, + } + for _, check := range checks { + assertRedisError(t, check.name, check.call(), runtimestorage.ErrInvalid) + } +} + +func TestRedisOperationsHonorCanceledContext(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + reply := runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: "event", SegmentCount: 1} + calls := []struct { + name string + call func() error + }{ + {name: "ping", call: func() error { return store.Ping(ctx) }}, + {name: "correlation", call: func() error { _, err := store.GetReplyCorrelation(ctx, "tenant-a", "event"); return err }}, + {name: "get session", call: func() error { _, err := store.GetSession(ctx, "tenant-a", "session"); return err }}, + {name: "create session", call: func() error { _, err := store.CreateSession(ctx, "tenant-a", "session", nil); return err }}, + {name: "update session", call: func() error { _, err := store.UpdateSessionState(ctx, "tenant-a", "session", 1, nil); return err }}, + {name: "delete session", call: func() error { return store.DeleteSession(ctx, "tenant-a", "session") }}, + {name: "record message", call: func() error { + _, _, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session", BindingID: "binding", ExternalMessageID: "external", EventID: "event"}) + return err + }}, + {name: "get message", call: func() error { _, err := store.GetMessage(ctx, "tenant-a", "event"); return err }}, + {name: "transition message", call: func() error { + _, err := store.TransitionMessage(ctx, runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: "event", From: runtimestorage.EventReceived, To: runtimestorage.EventRunning, Owner: "worker", LeaseDuration: time.Second}) + return err + }}, + {name: "append payload", call: func() error { + _, err := store.AppendEventPayload(ctx, runtimestorage.EventPayload{TenantID: "tenant-a", SessionID: "session", EventID: "event", Payload: []byte(`{}`)}) + return err + }}, + {name: "list payloads", call: func() error { _, err := store.ListEventPayloads(ctx, "tenant-a", "session"); return err }}, + {name: "enqueue reply", call: func() error { _, err := store.EnqueueReply(ctx, reply); return err }}, + {name: "enqueue replies", call: func() error { _, err := store.EnqueueReplies(ctx, []runtimestorage.ReplyOutbox{reply}); return err }}, + {name: "enqueue correlated replies", call: func() error { + _, err := store.EnqueueRepliesWithCorrelation(ctx, runtimestorage.ReplyCorrelation{TenantID: "tenant-a", EventID: "event", RequestID: "request"}, []runtimestorage.ReplyOutbox{reply}) + return err + }}, + {name: "get reply", call: func() error { _, err := store.GetReply(ctx, "tenant-a", "reply", 0); return err }}, + {name: "list candidates", call: func() error { _, err := store.ListReplyCandidates(ctx, "tenant-a"); return err }}, + {name: "claim reply", call: func() error { + _, err := store.ClaimReply(ctx, "tenant-a", "reply", 0, "worker", time.Second) + return err + }}, + {name: "transition reply", call: func() error { + _, err := store.TransitionReply(ctx, runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply", From: runtimestorage.ReplyPending, To: runtimestorage.ReplySending, Owner: "worker"}) + return err + }}, + {name: "put memory", call: func() error { + _, err := store.PutMemory(ctx, runtimestorage.MemoryInput{TenantID: "tenant-a", UserID: "user", Content: "content"}) + return err + }}, + {name: "get memory", call: func() error { _, err := store.GetMemory(ctx, "tenant-a", "memory"); return err }}, + {name: "list memories", call: func() error { _, err := store.ListMemories(ctx, "tenant-a", "user", 0); return err }}, + {name: "search memories", call: func() error { _, err := store.SearchMemories(ctx, "tenant-a", "user", "content", 0); return err }}, + {name: "delete memory", call: func() error { return store.DeleteMemory(ctx, "tenant-a", "memory") }}, + {name: "enqueue memory index", call: func() error { + return store.EnqueueMemoryIndex(ctx, runtimestorage.MemoryRecord{TenantID: "tenant-a", MemoryID: "memory", Version: 1}) + }}, + {name: "wait memory index", call: func() error { return store.WaitForMemoryIndex(ctx, "tenant-a", "memory", 1) }}, + } + for _, value := range calls { + assertRedisError(t, value.name, value.call(), context.Canceled) + } +} + +func TestRedisMissingRecordsReturnNotFound(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + reply := runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: "event", SegmentCount: 1} + calls := []struct { + name string + call func() error + }{ + {name: "correlation", call: func() error { _, err := store.GetReplyCorrelation(ctx, "tenant-a", "event"); return err }}, + {name: "session", call: func() error { _, err := store.GetSession(ctx, "tenant-a", "session"); return err }}, + {name: "session update", call: func() error { _, err := store.UpdateSessionState(ctx, "tenant-a", "session", 1, nil); return err }}, + {name: "session delete", call: func() error { return store.DeleteSession(ctx, "tenant-a", "session") }}, + {name: "message record", call: func() error { + _, _, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session", BindingID: "binding", ExternalMessageID: "external", EventID: "event"}) + return err + }}, + {name: "message", call: func() error { _, err := store.GetMessage(ctx, "tenant-a", "event"); return err }}, + {name: "message transition", call: func() error { + _, err := store.TransitionMessage(ctx, runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: "event", From: runtimestorage.EventReceived, To: runtimestorage.EventRunning, Owner: "worker", LeaseDuration: time.Second}) + return err + }}, + {name: "payload append", call: func() error { + _, err := store.AppendEventPayload(ctx, runtimestorage.EventPayload{TenantID: "tenant-a", SessionID: "session", EventID: "event", Payload: []byte(`{}`)}) + return err + }}, + {name: "payload list", call: func() error { _, err := store.ListEventPayloads(ctx, "tenant-a", "session"); return err }}, + {name: "reply enqueue", call: func() error { _, err := store.EnqueueReply(ctx, reply); return err }}, + {name: "reply batch", call: func() error { _, err := store.EnqueueReplies(ctx, []runtimestorage.ReplyOutbox{reply}); return err }}, + {name: "reply", call: func() error { _, err := store.GetReply(ctx, "tenant-a", "reply", 0); return err }}, + {name: "reply claim", call: func() error { + _, err := store.ClaimReply(ctx, "tenant-a", "reply", 0, "worker", time.Second) + return err + }}, + {name: "reply transition", call: func() error { + _, err := store.TransitionReply(ctx, runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply", From: runtimestorage.ReplyPending, To: runtimestorage.ReplySending, Owner: "worker"}) + return err + }}, + {name: "memory", call: func() error { _, err := store.GetMemory(ctx, "tenant-a", "memory"); return err }}, + {name: "memory delete", call: func() error { return store.DeleteMemory(ctx, "tenant-a", "memory") }}, + {name: "memory index", call: func() error { + return store.EnqueueMemoryIndex(ctx, runtimestorage.MemoryRecord{TenantID: "tenant-a", MemoryID: "memory", Version: 1}) + }}, + {name: "memory index wait", call: func() error { return store.WaitForMemoryIndex(ctx, "tenant-a", "memory", 1) }}, + } + for _, value := range calls { + assertRedisError(t, value.name, value.call(), runtimestorage.ErrNotFound) + } + memories, err := store.ListMemories(ctx, "tenant-a", "user", 0) + if err != nil || len(memories) != 0 { + t.Fatalf("empty memories = %+v, %v", memories, err) + } + candidates, err := store.ListReplyCandidates(ctx, "tenant-a") + if err != nil || len(candidates) != 0 { + t.Fatalf("empty reply candidates = %+v, %v", candidates, err) + } +} + +func TestRedisMessageSuccessTransitions(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + event := seedEvent(t, store, "tenant-a", "session", "event") + running, err := store.TransitionMessage(ctx, runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: event.EventID, From: runtimestorage.EventReceived, To: runtimestorage.EventRunning, Owner: "worker", LeaseDuration: time.Second}) + if err != nil || running.Status != runtimestorage.EventRunning || running.FencingToken != 1 || running.LeaseOwner != "worker" { + t.Fatalf("running event = %+v, %v", running, err) + } + completed, err := store.TransitionMessage(ctx, runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: event.EventID, From: runtimestorage.EventRunning, To: runtimestorage.EventCompleted, Owner: "worker", FencingToken: running.FencingToken}) + if err != nil || completed.Status != runtimestorage.EventCompleted || completed.LeaseExpiresAt != nil { + t.Fatalf("completed event = %+v, %v", completed, err) + } + pending, err := store.TransitionMessage(ctx, runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: event.EventID, From: runtimestorage.EventCompleted, To: runtimestorage.EventReplyPending, Owner: "worker", ReplyID: "reply", SegmentCount: 1}) + if err != nil || pending.Status != runtimestorage.EventReplyPending || pending.ReplyID != "reply" || pending.SegmentCount != 1 { + t.Fatalf("pending reply event = %+v, %v", pending, err) + } + replied, err := store.TransitionMessage(ctx, runtimestorage.MessageTransition{TenantID: "tenant-a", EventID: event.EventID, From: runtimestorage.EventReplyPending, To: runtimestorage.EventReplied, Owner: "worker"}) + if err != nil || replied.Status != runtimestorage.EventReplied { + t.Fatalf("replied event = %+v, %v", replied, err) + } + if got, err := store.GetMessage(ctx, "tenant-a", event.EventID); err != nil || got.Status != runtimestorage.EventReplied { + t.Fatalf("persisted replied event = %+v, %v", got, err) + } +} + +func TestRedisReplyDeliverySuccessTransitions(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + event := seedEvent(t, store, "tenant-a", "session", "event") + reply, err := store.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentCount: 1, Payload: "response"}) + if err != nil || reply.Status != runtimestorage.ReplyPending { + t.Fatalf("pending reply = %+v, %v", reply, err) + } + claimed, err := store.ClaimReply(ctx, "tenant-a", "reply", 0, "worker", time.Second) + if err != nil || claimed.Status != runtimestorage.ReplySending || claimed.Attempts != 1 { + t.Fatalf("claimed reply = %+v, %v", claimed, err) + } + sent, err := store.TransitionReply(ctx, runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply", From: runtimestorage.ReplySending, To: runtimestorage.ReplySent, Owner: "worker", FencingToken: claimed.FencingToken, ProviderID: "provider-message"}) + if err != nil || sent.Status != runtimestorage.ReplySent || sent.ProviderMessageID != "provider-message" || sent.LeaseExpiresAt != nil { + t.Fatalf("sent reply = %+v, %v", sent, err) + } +} + +func TestRedisMemoryDefaultValues(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + value, err := store.PutMemory(context.Background(), runtimestorage.MemoryInput{TenantID: "tenant-a", UserID: "user", Content: "content"}) + if err != nil || !strings.HasPrefix(value.MemoryID, "mem_") || value.Version != 1 { + t.Fatalf("defaulted memory = %+v, %v", value, err) + } + if _, err := store.SearchMemories(context.Background(), "tenant-a", "user", "absent", 0); err != nil { + t.Fatalf("empty search = %v", err) + } +} + +func TestRedisConstructorsNormalizeDefaultsAndOptions(t *testing.T) { + server := miniredis.RunT(t) + ctx := context.Background() + client := redisclient.NewClient(&redisclient.Options{Addr: server.Addr()}) + borrowed, err := redisstore.New(client, "") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = borrowed.Close(); _ = client.Close() }) + if _, err := borrowed.CreateSession(ctx, "tenant-a", "session", nil); err != nil { + t.Fatal(err) + } + if !server.Exists("trpc:runtime:v1:" + hex.EncodeToString([]byte("tenant-a"))) { + t.Fatal("default Redis key prefix was not used") + } + configured, err := redisstore.NewFromConfig(ctx, redisstore.Config{ + Addr: server.Addr(), + DialTimeout: time.Second, + ReadTimeout: time.Second, + WriteTimeout: time.Second, + PoolSize: 2, + }) + if err != nil { + t.Fatal(err) + } + if err := configured.Close(); err != nil { + t.Fatal(err) + } + addr := server.Addr() + server.Close() + if _, err := redisstore.NewFromURL(ctx, "redis://"+addr); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("unavailable URL constructor = %v", err) + } +} + +func TestRedisReplyInsertionIsIdempotentAndFenced(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + ctx := context.Background() + target := runtimestorage.ReplyTarget{BindingID: "binding-target", ConversationKind: "direct", ReceiverID: "receiver"} + event := seedEvent(t, store, "tenant-a", "session", "event") + if _, _, err := store.RecordMessage(ctx, runtimestorage.MessageEventInput{TenantID: "tenant-a", SessionID: "session", BindingID: "binding-target", ExternalMessageID: "external-target", EventID: "event-target", ReplyTarget: target}); err != nil { + t.Fatal(err) + } + input := runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentCount: 1, Payload: "payload"} + first, err := store.EnqueueReply(ctx, input) + if err != nil || first.Status != runtimestorage.ReplyPending { + t.Fatalf("first reply = %+v, %v", first, err) + } + second, err := store.EnqueueReply(ctx, input) + if err != nil || second.CreatedAt != first.CreatedAt { + t.Fatalf("idempotent reply = %+v, %v", second, err) + } + if _, err := store.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentCount: 1, Payload: "changed"}); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("reply payload conflict = %v", err) + } + if _, err := store.EnqueueReply(ctx, runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "target", EventID: "event-target", SegmentCount: 1}); !errors.Is(err, runtimestorage.ErrConflict) { + t.Fatalf("reply target conflict = %v", err) + } +} + +func TestRedisReplyBatchRejectsInconsistentSegments(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + event := seedEvent(t, store, "tenant-a", "session", "event") + base := runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentCount: 2} + tests := []struct { + name string + batch []runtimestorage.ReplyOutbox + }{ + {name: "empty", batch: nil}, + {name: "invalid first", batch: []runtimestorage.ReplyOutbox{{TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentCount: 1, Status: runtimestorage.ReplySent}}}, + {name: "different tenant", batch: []runtimestorage.ReplyOutbox{base, {TenantID: "tenant-b", ReplyID: "reply", EventID: event.EventID, SegmentIndex: 1, SegmentCount: 2}}}, + {name: "different reply", batch: []runtimestorage.ReplyOutbox{base, {TenantID: "tenant-a", ReplyID: "other", EventID: event.EventID, SegmentIndex: 1, SegmentCount: 2}}}, + {name: "different event", batch: []runtimestorage.ReplyOutbox{base, {TenantID: "tenant-a", ReplyID: "reply", EventID: "other", SegmentIndex: 1, SegmentCount: 2}}}, + {name: "different count", batch: []runtimestorage.ReplyOutbox{base, {TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentIndex: 1, SegmentCount: 3}}}, + {name: "duplicate segment", batch: []runtimestorage.ReplyOutbox{base, base}}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + if _, err := store.EnqueueReplies(context.Background(), test.batch); !errors.Is(err, runtimestorage.ErrInvalid) { + t.Fatalf("batch error = %v", err) + } + }) + } +} + +func TestRedisReplyTransitionCanStartDelivery(t *testing.T) { + server := miniredis.RunT(t) + store := newStore(t, server) + event := seedEvent(t, store, "tenant-a", "session", "event") + if _, err := store.EnqueueReply(context.Background(), runtimestorage.ReplyOutbox{TenantID: "tenant-a", ReplyID: "reply", EventID: event.EventID, SegmentCount: 1}); err != nil { + t.Fatal(err) + } + sending, err := store.TransitionReply(context.Background(), runtimestorage.ReplyTransition{TenantID: "tenant-a", ReplyID: "reply", From: runtimestorage.ReplyPending, To: runtimestorage.ReplySending, Owner: "worker", LeaseDuration: time.Second}) + if err != nil || sending.Status != runtimestorage.ReplySending || sending.Attempts != 1 || sending.LeaseExpiresAt == nil { + t.Fatalf("sending reply = %+v, %v", sending, err) + } +} + +func TestRedisMalformedStateAndUnavailableCommandsFailClosed(t *testing.T) { + server := miniredis.RunT(t) + client := redisclient.NewClient(&redisclient.Options{Addr: server.Addr()}) + store, err := redisstore.New(client, "test:runtime:v1") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close(); _ = client.Close() }) + ctx := context.Background() + key := "test:runtime:v1:" + hex.EncodeToString([]byte("tenant-a")) + if err := client.Set(ctx, key, "{", 0).Err(); err != nil { + t.Fatal(err) + } + if _, err := store.GetSession(ctx, "tenant-a", "session"); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("malformed state read = %v", err) + } + if err := client.Set(ctx, key, `{"version":99}`, 0).Err(); err != nil { + t.Fatal(err) + } + if _, err := store.GetSession(ctx, "tenant-a", "session"); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("unsupported state version = %v", err) + } + server.Close() + if _, err := store.GetSession(ctx, "tenant-a", "session"); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("unavailable read = %v", err) + } + if _, err := store.CreateSession(ctx, "tenant-a", "session", nil); !errors.Is(err, runtimestorage.ErrStorage) { + t.Fatalf("unavailable mutation = %v", err) + } +}