From 8a6c31c14d233ef8d269bcc8c4f5f7129a8d53d6 Mon Sep 17 00:00:00 2001 From: andy Date: Sat, 5 Sep 2026 15:46:37 +0800 Subject: [PATCH] =?UTF-8?q?=E5=88=9D=E5=A7=8B=E5=8C=96=E7=AC=AC=E4=B8=80?= =?UTF-8?q?=E7=89=88?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .dockerignore | 35 + .env.example | 65 ++ .gitignore | 29 + AGENTS.md | 105 ++ CONTEXT.md | 99 ++ Dockerfile | 40 + PROJECT_STATE.md | 113 ++ cmd/postgis-probe/main.go | 64 ++ cmd/postgis-srid-migrate/main.go | 101 ++ cmd/server/main.go | 31 + cmd/superagent-probe/main.go | 118 +++ compose.yaml | 44 + deploy/nginx/fire-safety-ymd.conf.example | 133 +++ docs/architecture/chat-api-v1.md | 82 ++ docs/architecture/public-chat-entry-v1.md | 57 ++ docs/architecture/spatial-mcp-v1.md | 73 ++ docs/import/db-samples/README.md | 34 + ...ai-native-software-engineering-standard.md | 350 +++++++ .../ai-native-templates/ADR.template.md | 40 + .../ai-native-templates/AGENTS.template.md | 39 + .../AGENT_HANDOFF.template.md | 57 ++ .../ARCHITECTURE.template.md | 31 + .../CHANGE_REQUEST.template.md | 61 ++ .../ai-native-templates/CONTEXT.template.md | 50 + .../DOMAIN_OBJECT.template.md | 27 + .../PROJECT_STATE.template.md | 39 + .../reusable/ai-native-templates/README.md | 23 + .../ai-native-templates/SPEC.template.md | 75 ++ .../ai-native-templates/WORKFLOW.template.md | 37 + .../general-development-guidelines.md | 151 +++ docs/import/字段组.xlsx | Bin 0 -> 9179 bytes docs/import/数据表映射.xlsx | Bin 0 -> 3789 bytes docs/project/README.md | 50 + docs/project/ai-native-adoption.md | 85 ++ docs/project/ai-nses-project-overlay.md | 74 ++ .../project/backend-development-guidelines.md | 141 +++ .../superagent-mcp-client.example.json | 33 + .../integrations/superagent-mcp-spatial.md | 160 +++ .../integrations/superagent-openapi.md | 125 +++ .../operations/docker-test-deployment.md | 379 +++++++ docs/project/operations/nginx-public-entry.md | 175 ++++ docs/project/operations/postgis-srid-4326.md | 81 ++ .../security-access-control-boundary.md | 133 +++ docs/specs/fire-safety-ymd-chat-api-v1.md | 164 +++ ...safety-ymd-dashscope-compatible-chat-v1.md | 143 +++ ...-ymd-superagent-mcp-spatial-readonly-v1.md | 210 ++++ ...ety-ymd-superagent-openapi-connectivity.md | 159 +++ docs/workflows/exercise-plan-evidence.md | 59 ++ docs/workflows/user-chat.md | 83 ++ go.mod | 13 + go.sum | 26 + internal/app/app.go | 184 ++++ internal/app/app_test.go | 116 +++ internal/config/config.go | 631 ++++++++++++ internal/config/config_test.go | 490 +++++++++ internal/domain/doc.go | 2 + internal/domain/spatial.go | 112 ++ internal/handler/chat.go | 466 +++++++++ internal/handler/chat_test.go | 273 +++++ internal/handler/dashscope_chat.go | 429 ++++++++ internal/handler/dashscope_chat_test.go | 330 ++++++ internal/handler/health.go | 43 + internal/handler/health_test.go | 83 ++ internal/handler/mcp.go | 854 +++++++++++++++ internal/handler/mcp_test.go | 479 +++++++++ internal/integration/superagent/chat.go | 100 ++ internal/integration/superagent/chat_test.go | 115 +++ internal/integration/superagent/client.go | 661 ++++++++++++ .../integration/superagent/client_test.go | 451 ++++++++ internal/integration/superagent/sse.go | 512 +++++++++ internal/integration/superagent/sse_test.go | 169 +++ internal/integration/superagent/types.go | 143 +++ internal/migration/postgis_srid.go | 403 ++++++++ internal/migration/postgis_srid_test.go | 150 +++ internal/repository/doc.go | 2 + internal/repository/postgis.go | 969 ++++++++++++++++++ internal/repository/postgis_test.go | 231 +++++ internal/service/chat.go | 455 ++++++++ internal/service/chat_test.go | 278 +++++ internal/service/doc.go | 2 + internal/service/spatial.go | 334 ++++++ internal/service/spatial_test.go | 334 ++++++ pkg/README.md | 5 + 83 files changed, 14302 insertions(+) create mode 100644 .dockerignore create mode 100644 .env.example create mode 100644 .gitignore create mode 100644 AGENTS.md create mode 100644 CONTEXT.md create mode 100644 Dockerfile create mode 100644 PROJECT_STATE.md create mode 100644 cmd/postgis-probe/main.go create mode 100644 cmd/postgis-srid-migrate/main.go create mode 100644 cmd/server/main.go create mode 100644 cmd/superagent-probe/main.go create mode 100644 compose.yaml create mode 100644 deploy/nginx/fire-safety-ymd.conf.example create mode 100644 docs/architecture/chat-api-v1.md create mode 100644 docs/architecture/public-chat-entry-v1.md create mode 100644 docs/architecture/spatial-mcp-v1.md create mode 100644 docs/import/db-samples/README.md create mode 100644 docs/import/reusable/ai-native-software-engineering-standard.md create mode 100644 docs/import/reusable/ai-native-templates/ADR.template.md create mode 100644 docs/import/reusable/ai-native-templates/AGENTS.template.md create mode 100644 docs/import/reusable/ai-native-templates/AGENT_HANDOFF.template.md create mode 100644 docs/import/reusable/ai-native-templates/ARCHITECTURE.template.md create mode 100644 docs/import/reusable/ai-native-templates/CHANGE_REQUEST.template.md create mode 100644 docs/import/reusable/ai-native-templates/CONTEXT.template.md create mode 100644 docs/import/reusable/ai-native-templates/DOMAIN_OBJECT.template.md create mode 100644 docs/import/reusable/ai-native-templates/PROJECT_STATE.template.md create mode 100644 docs/import/reusable/ai-native-templates/README.md create mode 100644 docs/import/reusable/ai-native-templates/SPEC.template.md create mode 100644 docs/import/reusable/ai-native-templates/WORKFLOW.template.md create mode 100644 docs/import/reusable/general-development-guidelines.md create mode 100644 docs/import/字段组.xlsx create mode 100644 docs/import/数据表映射.xlsx create mode 100644 docs/project/README.md create mode 100644 docs/project/ai-native-adoption.md create mode 100644 docs/project/ai-nses-project-overlay.md create mode 100644 docs/project/backend-development-guidelines.md create mode 100644 docs/project/integrations/superagent-mcp-client.example.json create mode 100644 docs/project/integrations/superagent-mcp-spatial.md create mode 100644 docs/project/integrations/superagent-openapi.md create mode 100644 docs/project/operations/docker-test-deployment.md create mode 100644 docs/project/operations/nginx-public-entry.md create mode 100644 docs/project/operations/postgis-srid-4326.md create mode 100644 docs/project/security-access-control-boundary.md create mode 100644 docs/specs/fire-safety-ymd-chat-api-v1.md create mode 100644 docs/specs/fire-safety-ymd-dashscope-compatible-chat-v1.md create mode 100644 docs/specs/fire-safety-ymd-superagent-mcp-spatial-readonly-v1.md create mode 100644 docs/specs/fire-safety-ymd-superagent-openapi-connectivity.md create mode 100644 docs/workflows/exercise-plan-evidence.md create mode 100644 docs/workflows/user-chat.md create mode 100644 go.mod create mode 100644 go.sum create mode 100644 internal/app/app.go create mode 100644 internal/app/app_test.go create mode 100644 internal/config/config.go create mode 100644 internal/config/config_test.go create mode 100644 internal/domain/doc.go create mode 100644 internal/domain/spatial.go create mode 100644 internal/handler/chat.go create mode 100644 internal/handler/chat_test.go create mode 100644 internal/handler/dashscope_chat.go create mode 100644 internal/handler/dashscope_chat_test.go create mode 100644 internal/handler/health.go create mode 100644 internal/handler/health_test.go create mode 100644 internal/handler/mcp.go create mode 100644 internal/handler/mcp_test.go create mode 100644 internal/integration/superagent/chat.go create mode 100644 internal/integration/superagent/chat_test.go create mode 100644 internal/integration/superagent/client.go create mode 100644 internal/integration/superagent/client_test.go create mode 100644 internal/integration/superagent/sse.go create mode 100644 internal/integration/superagent/sse_test.go create mode 100644 internal/integration/superagent/types.go create mode 100644 internal/migration/postgis_srid.go create mode 100644 internal/migration/postgis_srid_test.go create mode 100644 internal/repository/doc.go create mode 100644 internal/repository/postgis.go create mode 100644 internal/repository/postgis_test.go create mode 100644 internal/service/chat.go create mode 100644 internal/service/chat_test.go create mode 100644 internal/service/doc.go create mode 100644 internal/service/spatial.go create mode 100644 internal/service/spatial_test.go create mode 100644 pkg/README.md diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..dad944b --- /dev/null +++ b/.dockerignore @@ -0,0 +1,35 @@ +# VCS and local editor state +.git +.git/** +.idea +.vscode +.DS_Store +*.iml +*.swp +*.swo + +# Secrets and local-only configuration +.env +.env.* +!.env.example +*.pem +*.key + +# Project documentation and data exports are not runtime image inputs. +docs/ +deploy/ +*.sql +*.xlsx + +# Go build/test outputs and local logs +/bin/ +/build/ +/dist/ +/tmp/ +*.test +*.out +coverage.* +*.log + +# Other local dependency trees, if present +node_modules/ diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..a4c3e38 --- /dev/null +++ b/.env.example @@ -0,0 +1,65 @@ +# HTTP server +FIRE_SAFETY_HTTP_ADDR=:8080 + +# SuperAgent Open API (disabled until a project-specific test credential is supplied) +FIRE_SAFETY_SUPERAGENT_ENABLED=false +# Required when enabled, for example: https://superagent.example.com +FIRE_SAFETY_SUPERAGENT_BASE_URL= +FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY= +FIRE_SAFETY_SUPERAGENT_CONNECT_TIMEOUT=15s +FIRE_SAFETY_SUPERAGENT_RECOVERY_MAX_ATTEMPTS=5 +FIRE_SAFETY_SUPERAGENT_RECOVERY_INITIAL_BACKOFF=250ms +FIRE_SAFETY_SUPERAGENT_MAX_MESSAGE_BYTES=65536 + +# Connectivity probe only; do not use a real user ID or production data. +FIRE_SAFETY_SUPERAGENT_PROBE_SUBJECT_ID=fire-safety-ymd-connectivity-probe +FIRE_SAFETY_SUPERAGENT_PROBE_TIMEOUT=10m + +# User-facing chat API (disabled by default; requires SuperAgent above) +FIRE_SAFETY_CHAT_ENABLED=false +# Test-stage static Bearer only. Generate a distinct high-entropy value of at least 32 printable ASCII characters. +# Never reuse FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY or FIRE_SAFETY_MCP_AUTH_TOKEN. +FIRE_SAFETY_CHAT_AUTH_TOKEN= +# Server-controlled test subject. This is not final end-user identity or authorization. +FIRE_SAFETY_CHAT_SUBJECT_ID=fire-safety-ymd-chat-test-subject +# Optional legacy-client compatibility route: +# /api/v1/apps/{this-value}/completion. This is a public identifier, not a secret. +# Leave empty to expose only the native /api/chat route. +FIRE_SAFETY_CHAT_COMPAT_APP_ID= +# Exact comma-separated browser origins, for example: http://localhost:5173,https://fire.example.com +# Leave empty to allow only clients that do not send Origin, such as curl or a backend service. +FIRE_SAFETY_CHAT_ALLOWED_ORIGINS= +FIRE_SAFETY_CHAT_MAX_BODY_BYTES=131072 +FIRE_SAFETY_CHAT_RUN_TIMEOUT=10m +# Conversations are process-local and disappear on restart. +FIRE_SAFETY_CHAT_SESSION_TTL=30m +FIRE_SAFETY_CHAT_MAX_SESSIONS=1000 + +# Inbound MCP for SuperAgent (disabled by default) +FIRE_SAFETY_MCP_ENABLED=false +# Required when enabled. Use a new per-environment high-entropy token (at least 32 printable ASCII characters). +# Never reuse FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY. +FIRE_SAFETY_MCP_AUTH_TOKEN= +# Trusted server-side data scope: town_allowlist (default) or all. +# Use all only when every record in the MCP's fixed query tables is authorized for this credential. +FIRE_SAFETY_MCP_SCOPE_MODE=town_allowlist +# Required for town_allowlist and must be empty for all. This scope is not accepted from MCP tool arguments. +FIRE_SAFETY_MCP_ALLOWED_TOWNS= +FIRE_SAFETY_MCP_MAX_BODY_BYTES=262144 +FIRE_SAFETY_MCP_TOOL_TIMEOUT=5s + +# PostgreSQL/PostGIS (disabled by default) +FIRE_SAFETY_POSTGIS_ENABLED=false +# Secret. Example for local development only: +# postgresql://fire_safety_readonly:password@127.0.0.1:5432/fire_safety?sslmode=disable +FIRE_SAFETY_POSTGIS_DSN= +# One-off operator credential for cmd/postgis-srid-migrate only. It must be able +# to UPDATE all 8 source tables and ALTER the 7 legacy 2D geometry typmods. +# Never use this credential for cmd/server; remove it after the migration. +# compose.yaml additionally overrides it to empty inside the server container. +FIRE_SAFETY_POSTGIS_MIGRATION_DSN= +# Project source CRS is confirmed as EPSG:4326; strict readiness rejects raw SRID 0 imports. +FIRE_SAFETY_POSTGIS_EXPECTED_SRID=4326 +FIRE_SAFETY_POSTGIS_CONNECT_TIMEOUT=5s +FIRE_SAFETY_POSTGIS_QUERY_TIMEOUT=3s +FIRE_SAFETY_POSTGIS_MAX_CONNS=4 diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..4a7a7fd --- /dev/null +++ b/.gitignore @@ -0,0 +1,29 @@ +# Go build and test outputs +/bin/ +/build/ +/dist/ +/tmp/ +*.exe +*.test +*.out +coverage.* + +# Local environment files (keep documented examples) +.env +.env.* +!.env.example + +# Local database exports may contain operational contacts or destructive DDL. +/docs/import/db-samples/*.sql +!/docs/import/db-samples/README.md + +# IDE and editor files +.idea/ +.vscode/ +*.iml +*.swp +*.swo + +# Operating-system metadata +.DS_Store +Thumbs.db diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 0000000..f74f970 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,105 @@ +# fire-safety-ymd 项目协作与开发规范 + +## 1. 必读入口 + +开始任务前按顺序阅读: + +1. `AGENTS.md` +2. `CONTEXT.md` +3. `PROJECT_STATE.md` +4. `docs/project/README.md` +5. 与任务相关的项目文档、Spec、ADR 或 Workflow + +通用标准位于: + +- `docs/import/reusable/ai-native-software-engineering-standard.md` +- `docs/import/reusable/general-development-guidelines.md` + +项目规则与通用标准冲突时,以用户当前明确要求和 `docs/project/ai-nses-project-overlay.md` 中更具体的项目约束为准。 + +## 2. 项目定位 + +本仓库建设一个 Go 后端,负责连接用户侧应用与既有 SuperAgent 平台,并通过受控 MCP 工具向 Agent 提供森林防火业务数据。PostgreSQL/PostGIS 是空间业务事实的预期权威来源,SuperAgent 负责理解问题、选择工具和组织回答。 + +当前已完成项目骨架、`GET /health`、默认关闭的原生用户对话 API、可选 DashScope 风格兼容对话入口、SuperAgent Open API 出站适配器,以及默认关闭的空间只读 MCP/PostGIS 基线。对话入口只有独立静态联调凭证和单进程内存会话,不是最终用户认证;兼容 SSE 只在严格成功后返回正文。仓库已有多阶段 Docker/Compose 测试部署基线和把 `agent.nianxx.com` 精确反代到本 Go 服务的 Nginx 示例;容器端口只发布到宿主机回环地址,目标服务器和公网链路尚未验证。MCP 具备 7 个固定工具(含地名候选搜索)、独立 Bearer、显式服务端数据范围(数据库全范围或镇街白名单)、readiness 校验和只读参数化查询。数据提供方确认源数据为 EPSG:4326 且无偏移;2026-09-05 已通过受控事务为 4,048 条非空几何补齐 SRID 4326 元数据,严格 readiness 和 7 个工具的本地真实数据库冒烟测试均已通过,但真实消防 Profile 对话、公网部署与 SuperAgent 回调联调尚未完成。地名候选必须由用户确认,不要把查询成功当成资源可用性、路线、实时队伍位置或生产鉴权已经验证。 + +## 3. 默认工作方式 + +- 先确认 checkpoint 的目标、边界、验收标准和允许的写操作。 +- 修改前检查 Git 状态,保留用户已有变更,不回滚、不覆盖、不顺手整理无关文件。 +- 一次只完成一个明确 checkpoint;新需求进入 Spec、Change Request 或下一个 checkpoint。 +- 先读现有接口和项目文档,再新增抽象或依赖。 +- 对未确认的外部接口、数据结构、坐标系、权限规则和部署条件明确标记 `待确认`,不得猜测为事实。 +- 完成后报告变更、验证命令、验证结果、未确认项和建议下一步。 +- 除非用户明确要求,不自动创建提交、推送、迁移生产数据或调用外部系统写接口。 + +## 4. Go 工程规则 + +- 使用 `gofmt` 作为唯一基础格式化标准,变更后至少执行 `go test ./...`。 +- `cmd/server` 只负责进程启动、信号和依赖装配,不放业务规则。 +- `internal/handler` 负责 HTTP/MCP 入站契约、输入校验和响应映射。 +- `internal/service` 编排用例;`internal/domain` 保存稳定领域概念与规则。 +- `internal/repository` 保存持久化适配器和确实稳定的仓储契约,不允许 Handler 直连数据库。 +- `internal/config` 集中读取、默认值处理和配置校验;业务代码不散落读取环境变量。 +- `pkg` 只放确实需要被其他 Go module 导入的稳定 API;没有明确复用方时优先放入 `internal`。 +- 接口应小且由使用方定义;依赖通过构造函数显式注入,避免全局可变状态和隐藏初始化。 +- 错误必须携带上下文并使用 `%w` 保留错误链;正常业务失败不使用 `panic`。 +- 先使用标准库;引入框架或第三方依赖前记录用途、维护状态和替代方案。 + +完整后端规则见 `docs/project/backend-development-guidelines.md`。 + +## 5. 数据、AI 与安全边界 + +- 不允许 Agent 执行任意 SQL;MCP 只暴露固定、参数化、可授权和可审计的领域工具。 +- 用户身份、租户/区域范围和角色来自可信服务端上下文,不接受模型自由填写这些授权参数。 +- 未确认 SRID 前不得上线空间距离或包含关系计算;不得把 WKB 坐标数值猜测直接升级为数据库事实。 +- **原始几何维度不得被静默降维:**用户现场导入已确认 `st_2_xianyoufanghuotongdao.geom` 同时存在二维和带 Z 维度的几何;原导出声明 `geometry(GEOMETRY)` 会以二维 typmod 拒绝 Z 记录(PostgreSQL 错误 `22023`)。该表的原始导入列必须保持为不限定 typmod 的 `geometry`,除非另有经过审查的数据迁移决定。 +- 未经用户或数据所有者明确授权,不得对原始表执行 `ST_Force2D`、重写 WKB 或以其他方式丢弃 Z。若 MCP 算法只接受二维数据,应在只读查询/派生层显式投影并记录语义,不得修改原始事实。几何维度与 SRID 是两件事;成功容纳 Z 不代表坐标系已确认为 EPSG:4326。 +- 实库已发现 35 条无效面几何。当前决策是不自动修复或覆盖原始几何:readiness 明确告警,所有 MCP 查询使用 `ST_IsValid` 排除,工具响应提示结果可能不完整。SRID、坐标范围和几何类型仍是启用 MCP 的硬门禁。 +- 不把资源记录存在等同于设备当前可用;返回值需要能表达来源、更新时间、不确定性和现场确认要求。 +- 联系人、电话、精确位置、访问令牌和数据库凭证按敏感信息处理,不进入普通日志或模型无权限上下文。 +- AI 生成内容是辅助信息,不能虚构现场状态,也不能替代报警、人员撤离和现场指挥。 + +完整边界见 `docs/project/security-access-control-boundary.md`。 + +## 6. 常用命令 + +```bash +gofmt -w ./cmd ./internal +go test ./... +go vet ./... +go run ./cmd/server +curl -i http://localhost:8080/health +# 仅在只读 PostGIS 配置已显式加载时: +go run ./cmd/postgis-probe +# 仅使用临时表所有者/迁移凭证,详见项目 operations 文档: +go run ./cmd/postgis-srid-migrate +# 仅在项目专属测试凭证已配置时: +go run ./cmd/superagent-probe +``` + +默认监听地址为 `:8080`,可通过 `FIRE_SAFETY_HTTP_ADDR` 覆盖。Secret 只允许放在环境变量或未提交的 `.env` 文件中。 + +`docs/import/db-samples/*.sql` 含受限联系人和破坏性导出 DDL,已被 Git 忽略;只能只读分析,不得执行、提交或复制到测试输出。 + +## 7. 文档完成条件 + +每个 Feature 或 checkpoint 完成前检查: + +- `PROJECT_STATE.md` 是否反映真实状态。 +- Domain、Architecture、Workflow、ADR、Spec 是否需要新增或更新。 +- 接口、安全、权限、审计、外部系统和数据契约是否同步。 +- 验证命令和已知限制是否有可追踪记录。 + +没有文档变化时,交付说明中明确写 `No documentation changes required.`。 + +## 8. 多 Agent 协作 + +使用多 Agent 工作流时: + +- 主线程拆分任务、分配互不重叠的文件所有权并负责最终整合。 +- 边界明确、可独立完成的任务优先交给 `luna_worker`;需要修改代码且判断要求较高的任务交给 `sol_worker`。 +- 同一个文件同一时间只能由一个写入 Agent 负责,所有 Agent 必须保留他人的现有修改。 +- 实现结束后由 `sol_reviewer` 只读验证,并明确返回 `PASS` 或 `FAIL`。 +- 最多并行启动 6 个子 Agent;运行时上限更低时按实际上限执行。 +- 主线程等待所有子 Agent 返回后再统一总结。 diff --git a/CONTEXT.md b/CONTEXT.md new file mode 100644 index 0000000..b43b228 --- /dev/null +++ b/CONTEXT.md @@ -0,0 +1,99 @@ +# fire-safety-ymd 项目上下文 + +## 1. 项目目标 + +`fire-safety-ymd` 是一个面向森林防火场景的 Go 后端项目。目标调用链是:用户在业务应用中提问,后端维护会话并调用既有 SuperAgent 平台;SuperAgent 在需要业务事实时调用本项目提供的 MCP 工具;本项目从受控的 PostgreSQL/PostGIS 数据源查询消防资源和风险区域,再将结构化结果返回给 Agent 组织回答或辅助方案。 + +项目不自行训练或实现通用 Agent,也不允许大模型直接访问数据库。 + +## 2. 当前系统组成 + +| 路径或系统 | 当前职责 | 当前状态 | +| --- | --- | --- | +| `cmd/server` | Go 服务进程入口 | 已建立 | +| `internal/app` | 应用装配、readiness 和 HTTP 生命周期 | 已建立;按开关装配原生/兼容 Chat、SuperAgent 和 MCP/PostGIS | +| `internal/config` | 环境配置入口 | 已包含 HTTP、SuperAgent、Chat 兼容 App ID、MCP 与 PostGIS 配置校验和凭证分离门禁 | +| `internal/handler` | HTTP/MCP 入站协议层 | `GET /health` 已启用;默认关闭的原生 `/api/chat`、可选 DashScope 风格 `completion` 和 `/mcp` 已实现 | +| `internal/service` | 业务用例编排 | 已实现单进程聊天会话/并发 Run 控制,以及地名候选、有界空间查询、可信数据库全范围/镇街白名单和结果语义 | +| `internal/domain` | 森林防火领域模型与规则 | 已包含点位、水源、候选设施、通道、队伍和风险区模型 | +| `internal/repository` | PostgreSQL/PostGIS 持久化适配 | 已实现 pgxpool、只读固定 SQL 与 schema/SRID readiness;实库 SRID 元数据、严格 readiness 和 7 个工具真实查询已验证 | +| `cmd/postgis-srid-migrate` | 显式 SRID 元数据迁移 | 已执行;默认只读预检,写入需独立迁移凭证和明确 CRS 确认 | +| `internal/integration/superagent` | SuperAgent Open API 出站适配 | 已实现并通过模拟 Provider 测试,默认关闭 | +| `cmd/superagent-probe` | 无业务数据的显式连通性探针 | 已实现;需要项目专属测试配置 | +| `cmd/postgis-probe` | 不读取业务行的 PostGIS readiness 探针 | 已实现;需要只读数据库配置 | +| `Dockerfile` / `compose.yaml` | 测试环境容器构建与单实例进程托管 | 已建立;只发布宿主机回环端口,目标服务器尚未验证 | +| `pkg` | 可被外部 module 复用的稳定 Go API | 当前为空 | +| `docs/import` | 字段/表映射、通用模板和本地数据库样例 | 样例 SQL 含受限数据并被 Git 忽略,不会执行 | +| SuperAgent | 对话理解、工具选择和答案组织 | 已按仓库内 2026-07-12 协议基线实现客户端;当前环境待联调 | +| PostgreSQL/PostGIS | 消防空间业务事实的预期权威来源 | 8 表共 4,055 条记录;4,048 条非空几何已标记 EPSG:4326,严格 readiness 与本地真实工具冒烟均通过;35 条无效几何按当前策略排除并告警 | + +## 3. 技术栈 + +### 已采用 + +- Go module:`fire-safety-ymd`(临时 module 名)。 +- 本地工具链:Go `1.26.6`。 +- HTTP:Go 标准库 `net/http`。 +- 测试:Go 标准库 `testing`、`httptest`。 +- 配置:环境变量;支持 HTTP、SuperAgent、Chat、MCP 与 PostGIS 配置,并默认关闭 Chat 及两个外部方向。 +- Chat API:标准库 HTTP/SSE;原生 `/api/chat` 使用静态联调 Bearer,可选 `completion` 兼容入口使用同一信任方向的 `xtoken`;两者共享精确 Origin、严格 JSON、总超时、有界单进程会话和同会话并发冲突。兼容入口只在严格成功后发送正文。 +- SuperAgent:标准库 HTTP/SSE 客户端,分离 Session 创建和消息发送,支持严格完成判定与既有 Run 断流恢复。 +- MCP:标准库 HTTP/JSON-RPC,协议基线 `2025-06-18`,同步 JSON 响应,独立 Bearer 和 7 个只读工具;地名工具只搜索现有业务记录并要求用户确认候选。 +- PostgreSQL:`github.com/jackc/pgx/v5 v5.10.0` 原生连接池;连接默认只读并设置 statement timeout。 +- PostGIS:`ST_Covers`、`ST_DWithin`、`ST_Distance` 和 `ST_ClosestPoint`;只在实库确认 EPSG:4326 后启用。 +- 部署:多阶段 Docker 镜像与单实例 Compose;Secret 通过未提交的 `.env` 在运行时注入,应用端口只发布到宿主机 `127.0.0.1:8080`,由宿主机 Nginx 终止 TLS。 + +### 计划但尚未接入或确认 + +- PostgreSQL/PostGIS 的适用索引和生产查询计划验证。 +- SuperAgent 到 `/mcp` 的真实网络、TLS、Header 与 Token 联调。 +- 任意地址/山名的外部地理编码、别名词典和大数据量地名索引。 +- 用户聊天的真实身份认证、动态授权、共享/持久会话、主动取消和限流策略;首版默认关闭的静态 Bearer + 内存会话 API 已实现。 +- 最终用户鉴权、动态角色/区域或租户隔离、持久审计与完整可观测性方案。 +- 正式前端身份接入;现有客户端仅通过受限 DashScope 风格协议适配。 + +## 4. 已知业务数据范围 + +当前资料描述 8 类森林防火空间资源: + +| 资源 | 预期 PostgreSQL 表 | 预期几何形态 | +| --- | --- | --- | +| 蓄水池 | `st_2_xianyouxushuichiguan` | 点 | +| 防火通道 | `st_2_xianyoufanghuotongdao` | 线;样例为 MultiLineString,原始数据混合二维与 Z 维度 | +| 水源地 | `st_2_mpslfh_t_slfh_syd` | 点 | +| 林区工矿企业 | `st_2_linqugongkuangqiye` | 面 | +| 防火网格 | `st_2_fanghuowangge` | 面 | +| 防火瞭望哨 | `st_2_fanghuoliaowangshao` | 点 | +| 防火检查站 | `st_2_fanghuojianchazhan` | 点 | +| 墓地坟区 | `st_2_mudifenqu_mian` | 面 | + +仓库中的两份 Excel 是字段与数据表映射资料,不是可查询的业务数据库。`docs/import/db-samples/*.sql` 是含敏感字段的本地参考输入,已被 Git 忽略且不得由开发 Agent 执行。用户在 2026-09-04 的现场导入中确认,防火通道原始数据同时包含二维和带 Z 维度的几何:`geometry(GEOMETRY)` 导入到第 256 条附近时报 `Geometry has Z dimension but column does not`,将该原始列改为不限定 typmod 的 `geometry` 后重导成功。2026-09-05 的初始只读 audit 覆盖 8 表共 4,055 条记录:4,048 条非空几何均标记 SRID 0,7 条为空,35 条面几何无效,仅防火网格表存在 GiST 几何索引;类型和经纬度数值范围符合预期。数据提供方随后确认 8 表源数据均为 EPSG:4326、无坐标偏移,并允许 MCP 二维计算忽略 Z,但原始 Z 仍须保留。项目使用独立 `admin` 迁移凭证在单个事务中补齐了 4,048 条几何的 SRID 4326 元数据:7 张二维表改为 `geometry(Geometry,4326)`,防火通道保持裸 `geometry` 且 2 条 Z 几何未变;迁移前后无 SRID WKB 指纹、行数、类型、有效性和维度一致,随后只读严格 readiness 通过。无效几何仍不自动修复,查询排除并明确告警。 + +## 5. 目标职责边界 + +- 用户侧应用:采集用户输入并展示结果;具体形态待确认。 +- 本 Go 服务:鉴权上下文、会话转发、MCP 工具、数据查询、权限、安全和审计。 +- SuperAgent:理解自然语言、决定是否调用工具、组织自然语言结果。 +- MCP 工具:提供固定的森林防火领域查询,不开放任意 SQL 或跨权限访问。 +- PostgreSQL/PostGIS:保存和计算可信空间事实;资源是否可用仍取决于明确状态和数据时效。 + +## 6. 当前开发方向与非目标 + +当前阶段已有可运行、可测试、文档自解释的 Go 基线、SuperAgent Open API Adapter、默认关闭的原生用户对话 API、可选 DashScope 风格兼容入口,以及空间只读 MCP/PostGIS 实现。对话入口使用独立静态联调凭证和单进程内存会话;仓库已有多阶段 Docker/Compose 基线以及精确路径、无 Secret 的 Nginx HTTPS 反向代理示例,但尚未在目标机验证。MCP 默认关闭,实库严格 readiness 和全部 7 个工具的本地真实查询已通过。下一阶段在测试服务器应用这些部署资产,并使用真实消防 Profile 联调公网对话与 MCP。 + +本阶段不实现: + +- 面向真实用户的认证、动态权限、持久会话和已验证生产能力;当前对话入口只用于受控联调。 +- PostgreSQL/PostGIS 数据导入或索引 DDL;SRID 元数据迁移是已经显式执行的一次受控运维操作,不是应用运行时行为。 +- 用户登录、权限模型、审计存储和生产部署。 +- 面向真实火情的自动决策或路径规划。 +- 队伍实时定位、集结点管理和资源调度。 + +## 7. 新 Agent 阅读顺序 + +1. `AGENTS.md` +2. `CONTEXT.md` +3. `PROJECT_STATE.md` +4. `docs/project/README.md` +5. `docs/project/ai-nses-project-overlay.md` +6. 与任务相关的后端或安全规范 diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..4b8ef2c --- /dev/null +++ b/Dockerfile @@ -0,0 +1,40 @@ +# syntax=docker/dockerfile:1 + +# Keep the builder and runtime versions explicit so a test deployment does not +# silently move to an unrelated Go or Alpine release. +FROM golang:1.26.8-alpine3.24 AS build + +WORKDIR /src + +RUN apk add --no-cache ca-certificates + +COPY go.mod go.sum ./ +RUN go mod download + +# Copy only Go sources into the build context. Documentation, SQL exports and +# local configuration are intentionally not part of the application image. +COPY cmd ./cmd +COPY internal ./internal +COPY pkg ./pkg + +RUN go test ./... \ + && CGO_ENABLED=0 GOOS=linux go build \ + -buildvcs=false \ + -trimpath \ + -ldflags="-s -w" \ + -o /out/fire-safety-server \ + ./cmd/server + +FROM alpine:3.24.1 + +RUN apk add --no-cache ca-certificates tzdata \ + && addgroup -S app \ + && adduser -S -D -H -G app app + +COPY --from=build --chown=app:app /out/fire-safety-server /usr/local/bin/fire-safety-server + +USER app +WORKDIR /app +EXPOSE 8080 + +ENTRYPOINT ["/usr/local/bin/fire-safety-server"] diff --git a/PROJECT_STATE.md b/PROJECT_STATE.md new file mode 100644 index 0000000..4fd79be --- /dev/null +++ b/PROJECT_STATE.md @@ -0,0 +1,113 @@ +# fire-safety-ymd 项目当前状态 + +| 项 | 内容 | +| --- | --- | +| 最近更新 | 2026-09-05 | +| 当前分支 | `main` | +| 当前阶段 | 对话、SuperAgent、空间 MCP 与测试环境容器部署基线已完成 | +| 当前重点 | 将代码提交并推送后,在目标服务器验证 Compose、Nginx、真实消防 Profile 与公网 MCP 链路 | + +## 1. 当前 Checkpoint + +- 名称:`fire-safety-ymd-container-test-deployment-bootstrap` +- 状态:Complete +- 目标:为 `/home/firee-safety-ymd` 测试服务器提供可复现的多阶段 Docker 镜像、单实例 Compose、回环端口边界、Nginx 公网入口以及不泄露 Secret 的部署、验证和回滚手册。 +- 非目标:替用户提交或推送 Git、直接修改远程服务器、创建数据库容器、迁移生产数据、签发证书、改变 DNS/安全组、实现真实用户认证、动态授权、会话持久化或生产审计。 + +当前进展: + +- `Dockerfile` 使用显式 Go/Alpine 版本的多阶段构建,在构建阶段执行全部 Go 测试,最终镜像只包含静态服务二进制、CA 和时区数据,并以非 root 用户运行。 +- `compose.yaml` 只运行一个 API 实例,从未提交的 `.env` 注入配置,强制清空一次性迁移 DSN,只把 8080 发布到宿主机 `127.0.0.1`,并设置健康检查、只读文件系统、权限收紧和日志轮转。 +- 现有 PostgreSQL/PostGIS 不进入 Compose;同宿主机数据库需要使用容器可达的宿主机地址,且仍需受 `listen_addresses`、`pg_hba.conf` 和防火墙约束。 +- Nginx 示例增加 HTTP 到 HTTPS 跳转和 HTTP 429 JSON 限流响应,仍不比较、保存或注入 Chat、MCP、SuperAgent 或数据库 Secret。 +- 运维手册记录 Git 前置条件、服务器目录、Secret 权限、Compose/Nginx 启停、Chat/MCP 冒烟、SuperAgent 回调、更新和回滚。 + +已实现验收项: + +- 镜像构建上下文排除 `.env`、密钥、证书、SQL、Excel、文档和本机构建产物;运行时镜像不包含 Go 工具链或源码。 +- Compose 配置不包含明文 Secret,且没有数据库容器、数据卷或迁移命令;现有业务数据库不会被部署动作重建。 +- Nginx 上游固定为宿主机回环地址,兼容 Chat SSE 禁用缓冲和自动重试,未列出路径固定 404。 +- 目标机执行步骤包含不渲染 `.env` 内容的 Compose 检查、`nginx -t` 前置门禁和可恢复的配置替换。 +- 目标服务器尚未被本 checkpoint 修改;Docker/Nginx 真实验证留给下一 checkpoint。 + +## 2. 当前优先级 + +1. 由用户审查当前变更后创建首个 Git 提交并推送 `origin/main`;当前无提交,服务器尚不能 clone。 +2. 在 `/home/firee-safety-ymd` 构建并启动 Compose,确认健康状态、数据库可达和 8080 只绑定回环地址。 +3. 在目标机迁移 Nginx 配置并执行 `nginx -t` 后 reload;只公开 HTTPS 兼容路径、`/mcp` 和可选 `/health`。 +4. 使用项目专属测试 Key 和已发布消防 Profile 做兼容 `completion` 首轮/多轮真实冒烟,核对最终回答、用量和会话复用。 +5. 为所有查询表补齐适用 GiST 索引并验证查询计划;当前小数据可做联调,但生产前必须完成索引与并发验证。 +6. 设计真实用户身份、动态角色/区域授权、共享会话、限流、Secret 轮换、指标和持久审计。 + +## 3. 已确认事实 + +- 仓库已初始化 Git,当前分支是 `main`,已有远程 `origin`;尚无提交。 +- `docs/import/字段组.xlsx` 和 `docs/import/数据表映射.xlsx` 是用户已有、已暂存的变更,本 checkpoint 未修改。 +- 用户已使用项目只读 probe 成功连接数据库 `fire_safety_ymd`;PostGIS 报告版本 `3.3 USE_GEOS=1 USE_PROJ=1 USE_STATS=1`,8 表合计 4,055 条记录。 +- 实库 4,048 条非空几何已在 2026-09-05 的受控事务中从 SRID 0 补齐为 SRID 4326;坐标 extent 约为经度 121.16 至 121.93、纬度 37.07 至 37.49,类型与预期一致且未发现 WGS84 数值越界。数据提供方确认 8 表源数据均为 EPSG:4326 且无坐标偏移,严格 readiness 已通过。 +- 实库包含 7 条空几何;防火网格 4 条、林区工矿企业 4 条、墓地坟区 27 条无效面几何。用户决定首版不修复原始几何,由 MCP 排除并明确告警。 +- 实库仅防火网格表报告存在 GiST 几何索引,其余 7 张表尚未发现 GiST 几何索引。 +- 防火通道首次重导在约第 256 条因 `Geometry has Z dimension but column does not` 中止,证明源数据混合二维与 Z 维度;用户将 `st_2_xianyoufanghuotongdao.geom` 改为不限定 typmod 的 `geometry` 后报告重导成功。原始列不得无审查改回二维 `geometry(GEOMETRY)`,也不得静默 `ST_Force2D` 丢弃 Z。 +- 每张表的 2 条本地样例足以确定首版字段映射;样例不能证明全库质量或实时状态。 +- 原始导出样例的 `geom` 列均声明为 `geometry(GEOMETRY)`,但该声明不兼容防火通道中的 Z 记录;防火通道实库导入列已按用户现场操作改为裸 `geometry`。数据提供方允许当前二维算法忽略 Z,但原始 Z 仍须保留,不能对原始表执行 `ST_Force2D`。 +- Go module 当前使用临时名称 `fire-safety-ymd`,本地工具链为 Go `1.26.6`。 +- 直接第三方依赖为 `github.com/jackc/pgx/v5 v5.10.0`;HTTP、JSON-RPC/MCP 和测试仍使用 Go 标准库。 +- 服务默认监听 `:8080`;`GET /health` 仍只是 liveness,不访问外部依赖。 +- 用户对话 API、SuperAgent Open API Adapter 和 MCP endpoint 都默认关闭;Chat Bearer、Open API Key 和 MCP Token 属于三个独立信任方向并禁止复用。 +- `/api/chat` 已实现单进程内存会话映射、同会话并发 Run 冲突和严格 SSE 最终回答;模拟 Provider 端到端测试通过,真实 SuperAgent 尚未通过该入口联调。 +- 可选兼容入口已实现截图所示路径、`xtoken`、`input.prompt/session_id` 和 `event: result` 外形;正文仍只在严格成功的 `stop` 事件中出现,不是 DashScope 全量 API。 +- `Dockerfile`、`.dockerignore` 和 `compose.yaml` 已建立测试部署基线;容器单实例运行,8080 仅发布到宿主机回环地址,一次性迁移 DSN 在服务容器中强制为空。 +- `deploy/nginx/fire-safety-ymd.conf.example` 已将公网调用指向 `127.0.0.1:8080`,不再保存或注入任何 Provider/Chat/MCP Secret;目标机 Nginx 尚未验证。 +- 用户确认真实数据包含大量镇街,环境变量不适合枚举全量值;MCP 现支持显式数据库全范围 `all` 和默认镇街白名单 `town_allowlist` 两种服务端范围。 +- 用户选择先实现简单地名能力、后续再优化;当前只查询既有森林防火记录,不调用外部地图服务,也不把候选代表点自动认定为演练点。 +- 本地真实 MCP 冒烟已完成:7 个工具均成功访问实库,响应和错误边界符合契约;该结果不等于公网、SuperAgent 或生产并发已验证。 +- 防火通道现有字段不能支持可靠路线规划;防火网格现有字段不能支持实时队伍位置或正式集结点。 + +## 4. Known Issues 与未确认项 + +- 目标 PostgreSQL 的只读连接和 PostGIS 3.3 已由 probe 验证;数据库版本、TLS/网络生产拓扑、凭证轮换和审计仍未确认。 +- 35 条无效面几何和 7 条空几何会被查询排除;尤其 4 条无效防火网格可能造成所属网格和责任中队结果缺口,27 条无效墓地面可能造成风险区域漏项。 +- 除防火网格外 7 张表缺少 GiST 几何索引;当前 geography 距离表达式的生产索引方案需根据实库查询计划确认。 +- 水源/设施 `syzt`、水源 `hc_datetime` 等字段的枚举、单位、时区和更新责任人尚未确认。 +- SuperAgent MCP 的公网/内网 URL、TLS、网络白名单、Header 行为、Token 注入和轮换尚未联调。 +- `/api/chat` 静态 Bearer 和兼容路径 `xtoken` 只适用于受控联调,浏览器用户可以看到它;真实用户身份、动态授权、生产速率限制和滥用防护尚未实现。 +- 目标公网机器的 Docker/Compose 和 Nginx 版本、配置 include 层级、证书、DNS、安全组及 PostgreSQL 网络拓扑尚未验证;仓库已有部署资产和操作手册,但没有远程部署变更。 +- Chat 会话只在单个 Go 进程内存中保存;重启或多实例切换会丢失上下文,且当前没有历史查询、持久审计或主动取消 Provider Run。 +- 当前 `all`/`town_allowlist` 都是服务账号静态范围,不是最终用户级授权;`all` 会授权当前数据库中 MCP 固定查询表内所有镇街和镇街字段为空的记录,身份提供方、角色、租户和精确位置权限尚未确定。 +- 地名搜索是无索引的有界包含匹配;真实数据量下的耗时、重名率和名称字段质量尚未验证,生产优化可能需要标准地名表、别名词典或 `pg_trgm` 索引。 +- 镇街或村庄名称可能匹配多条资源/网格;线面只返回只读计算的代表点,不能直接作为真实演练点。 +- 只实现最小结构化运行日志,没有持久审计、指标、限流网关或数据源版本。 +- 路线规划缺少路网拓扑、坡度、路面、宽度、车辆限制、封路、实时火场和天气数据。 +- 队伍集结缺少正式集结点、实时定位、战备状态、人员/车辆/装备和容量数据。 +- `fire-safety-ymd` 正式 module path 与 CI/部署 Go 版本仍待确认。 + +## 5. Next Checkpoint + +建议:`fire-safety-ymd-container-public-entry-live-smoke-test`。 + +完成条件: + +- 用户完成首个 Git 提交并推送后,目标机能在 `/home/firee-safety-ymd` clone 或 `git pull --ff-only` 到明确 revision。 +- `docker compose build --pull`、`up -d` 和容器健康检查通过,8080 只绑定 `127.0.0.1`,运行容器中没有迁移凭证。 +- 在目标机替换安全 App ID、执行 `nginx -t` 后 reload,并验证 TLS、HTTP 到 HTTPS 跳转、429 和未列出路径 404。 +- 使用无敏感信息的问题验证兼容首轮 `null -> stop`、后续 `session_id` 复用、错误 xtoken、断流和超时。 +- 配置 SuperAgent 对公网 `/mcp` 的独立 Bearer,验证真实工具调用、TLS 和 warning 保留。 +- 不在输出、命令历史、Nginx、镜像层或 Git 中记录任何 Secret。 + +## 6. 验证记录 + +- `gofmt -w ./cmd ./internal`:通过。 +- `GOCACHE=/private/tmp/fire-safety-go-cache go test -count=1 ./...`:通过;新增覆盖兼容 App ID、`xtoken`、CORS、严格请求子集、`null -> stop` SSE、会话复用映射、用量/模型名、错误脱敏和路由/应用装配;原有 Chat、SuperAgent、MCP/PostGIS 覆盖继续通过。 +- `GOCACHE=/private/tmp/fire-safety-go-cache go vet ./...`:通过。 +- `GOCACHE=/private/tmp/fire-safety-go-cache go test -race -count=1 ./...`:通过。 +- 模拟 SuperAgent Chat 端到端:原生与兼容入口均通过;应用创建 Provider Session、发送消息、解析严格完成事件,原生返回 `conversation/message/done`,兼容入口返回 `result` 且最终 `finish_reason=stop`。 +- Nginx:配置已完成静态检查且未包含真实 Secret;当前开发机未安装 Nginx,尚未执行目标环境 `nginx -t` 或 reload。 +- Docker/Compose:部署文件已通过 YAML/静态安全断言;当前开发机未安装 Docker,尚未执行真实镜像构建、`docker compose config --quiet`、健康检查或容器运行验证。 +- SRID 迁移预检:通过;使用临时 `admin` 连接确认 4,048 条候选、8 表 UPDATE 权限和 7 表 ALTER 权限,未输出 DSN 或业务记录。 +- SRID 数据迁移:通过;事务更新 4,048 条非空几何,7 张二维表改为 `geometry(Geometry,4326)`,防火通道保持裸 `geometry`,迁移前后几何载荷指纹与维度一致。 +- 真实 PostGIS 严格 audit:通过;运行时只读账号报告 8 表 SRID 均为 4326、类型与范围门禁通过。预期保留 7 条空几何、35 条无效面几何和 7 张缺 GiST 索引表 warning。 +- 本地真实 MCP 冒烟:通过;health、鉴权、Origin、方法限制、协议初始化、7 工具发现、7 tools/call、无结果、非法参数、响应一致性、字段脱敏、warning 和结构化日志均符合契约。第二轮实库查询约 0.35 至 0.99 秒。 +- 数据导入:用户报告 8 份 SQL 已导入,防火通道在使用裸 `geometry` 保留混合二维/Z 后重导成功;这是现场反馈,不替代项目只读 probe 的最终验证。 +- 真实 SuperAgent 对话与 MCP 联调:对话 API 仅完成模拟 Provider 端到端测试,真实消防 Profile 尚未测试;MCP readiness 和本地真实工具冒烟已通过,公网 HTTPS 回调未配置。 +- 样例 SQL:未执行;含受限数据的 `*.sql` 已被 Git 忽略。 +- Git 提交:未创建,符合用户要求。 diff --git a/cmd/postgis-probe/main.go b/cmd/postgis-probe/main.go new file mode 100644 index 0000000..4c6fe2e --- /dev/null +++ b/cmd/postgis-probe/main.go @@ -0,0 +1,64 @@ +// Command postgis-probe performs a read-only, non-record-level readiness audit. +package main + +import ( + "context" + "encoding/json" + "fmt" + "log" + "os" + "os/signal" + "syscall" + "time" + + "fire-safety-ymd/internal/config" + "fire-safety-ymd/internal/repository" +) + +const probeTimeout = 45 * time.Second + +func main() { + cfg, err := config.Load() + if err != nil { + log.Fatalf("load configuration: %v", err) + } + if !cfg.PostGIS.Enabled { + log.Fatalf("%s must be true to run the PostGIS readiness probe", config.PostGISEnabledEnv) + } + + signalCtx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + ctx, cancel := context.WithTimeout(signalCtx, probeTimeout) + defer cancel() + + postGIS, err := repository.OpenPostGIS(ctx, repository.PostGISOptions{ + DSN: cfg.PostGIS.DSN, + MaxConns: cfg.PostGIS.MaxConns, + ConnectTimeout: cfg.PostGIS.ConnectTimeout, + QueryTimeout: cfg.PostGIS.QueryTimeout, + }) + if err != nil { + log.Fatalf("initialize PostGIS: %v", err) + } + defer postGIS.Close() + + var report repository.ReadinessReport + if cfg.PostGIS.ExpectedSRID == 0 { + report, err = postGIS.Audit(ctx) + } else { + report, err = postGIS.ValidateForMCP(ctx, cfg.PostGIS.ExpectedSRID) + } + if len(report.Tables) > 0 { + if encodeErr := json.NewEncoder(os.Stdout).Encode(report); encodeErr != nil { + log.Fatalf("encode readiness report: %v", encodeErr) + } + } + if err != nil { + log.Fatalf("audit PostGIS: %v", err) + } + if cfg.PostGIS.ExpectedSRID == 0 { + fmt.Fprintln(os.Stderr, "readiness note: FIRE_SAFETY_POSTGIS_EXPECTED_SRID is not set; audit only, MCP spatial validation was not performed") + return + } + fmt.Fprintln(os.Stderr, "MCP spatial validation passed") +} diff --git a/cmd/postgis-srid-migrate/main.go b/cmd/postgis-srid-migrate/main.go new file mode 100644 index 0000000..172f93e --- /dev/null +++ b/cmd/postgis-srid-migrate/main.go @@ -0,0 +1,101 @@ +// Command postgis-srid-migrate performs the explicitly approved EPSG:4326 metadata migration. +package main + +import ( + "context" + "encoding/json" + "flag" + "fmt" + "log" + "os" + "os/signal" + "strings" + "syscall" + "time" + + "fire-safety-ymd/internal/migration" + + "github.com/jackc/pgx/v5" +) + +const ( + migrationDSNEnv = "FIRE_SAFETY_POSTGIS_MIGRATION_DSN" + runtimeDSNEnv = "FIRE_SAFETY_POSTGIS_DSN" + commandTimeout = 3 * time.Minute +) + +func main() { + apply := flag.Bool("apply", false, "commit the reviewed SRID metadata migration") + confirmation := flag.String("confirm-source-crs", "", "required with --apply; must equal EPSG:4326") + flag.Parse() + + dsn, err := selectDSN(*apply) + if err != nil { + log.Fatal(err) + } + connectionConfig, err := pgx.ParseConfig(dsn) + if err != nil { + log.Fatal("parse PostGIS migration configuration") + } + connectionConfig.RuntimeParams["application_name"] = "fire-safety-ymd-srid-migration" + + signalCtx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + ctx, cancel := context.WithTimeout(signalCtx, commandTimeout) + defer cancel() + + conn, err := pgx.ConnectConfig(ctx, connectionConfig) + if err != nil { + log.Fatal("connect to PostGIS for SRID migration") + } + defer func() { _ = conn.Close(context.Background()) }() + + if !*apply { + report, err := migration.Inspect(ctx, conn) + if err != nil { + log.Fatalf("inspect SRID migration state: %v", err) + } + writeJSON(report) + if err := migration.ValidateForApply(report); err != nil { + log.Fatalf("SRID migration preflight failed: %v", err) + } + fmt.Fprintln(os.Stderr, "preflight passed; no database changes were made") + return + } + + if *confirmation != "EPSG:4326" { + log.Fatal("--apply requires --confirm-source-crs=EPSG:4326") + } + result, err := migration.ApplySRID4326(ctx, conn) + if err != nil { + log.Fatalf("apply SRID migration: %v", err) + } + writeJSON(result) + fmt.Fprintln(os.Stderr, "SRID metadata migration committed") +} + +func selectDSN(apply bool) (string, error) { + migrationDSN := strings.TrimSpace(os.Getenv(migrationDSNEnv)) + if apply { + if migrationDSN == "" { + return "", fmt.Errorf("%s is required with --apply", migrationDSNEnv) + } + return migrationDSN, nil + } + if migrationDSN != "" { + return migrationDSN, nil + } + runtimeDSN := strings.TrimSpace(os.Getenv(runtimeDSNEnv)) + if runtimeDSN == "" { + return "", fmt.Errorf("%s or %s is required", migrationDSNEnv, runtimeDSNEnv) + } + return runtimeDSN, nil +} + +func writeJSON(value any) { + encoder := json.NewEncoder(os.Stdout) + encoder.SetIndent("", " ") + if err := encoder.Encode(value); err != nil { + log.Fatalf("encode migration report: %v", err) + } +} diff --git a/cmd/server/main.go b/cmd/server/main.go new file mode 100644 index 0000000..4517f03 --- /dev/null +++ b/cmd/server/main.go @@ -0,0 +1,31 @@ +package main + +import ( + "context" + "log" + "os" + "os/signal" + "syscall" + + "fire-safety-ymd/internal/app" + "fire-safety-ymd/internal/config" +) + +func main() { + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + + cfg, err := config.Load() + if err != nil { + log.Fatalf("load configuration: %v", err) + } + application, err := app.New(ctx, cfg) + if err != nil { + log.Fatalf("initialize application: %v", err) + } + + log.Printf("HTTP server listening on %s", cfg.HTTPAddress) + if err := application.Run(ctx); err != nil { + log.Fatalf("run application: %v", err) + } +} diff --git a/cmd/superagent-probe/main.go b/cmd/superagent-probe/main.go new file mode 100644 index 0000000..80b9b93 --- /dev/null +++ b/cmd/superagent-probe/main.go @@ -0,0 +1,118 @@ +// Command superagent-probe performs a safe, explicit Open API connectivity check. +package main + +import ( + "context" + "crypto/rand" + "encoding/base64" + "encoding/json" + "fmt" + "log" + "os" + "os/signal" + "syscall" + + "fire-safety-ymd/internal/config" + "fire-safety-ymd/internal/integration/superagent" +) + +const probeMessage = "请只回复“SuperAgent connectivity OK”,不要调用业务工具。" + +type probeResult struct { + Status string `json:"status"` + SessionID string `json:"session_id"` + RunID string `json:"run_id,omitempty"` + ProfileID string `json:"profile_id,omitempty"` + ProfileVersionID string `json:"profile_version_id,omitempty"` + ModelName string `json:"model_name,omitempty"` + Answer string `json:"answer"` + Usage superagent.TokenUsage `json:"usage"` + EventTypes []string `json:"event_types"` +} + +func main() { + log.SetFlags(0) + cfg, err := config.Load() + if err != nil { + log.Fatalf("load configuration: %v", err) + } + if !cfg.SuperAgent.Enabled { + log.Fatalf("connectivity probe is disabled; set %s=true with a project-specific test credential", config.SuperAgentEnabledEnv) + } + + client, err := superagent.NewHTTPClient(superagent.Config{ + Enabled: cfg.SuperAgent.Enabled, + BaseURL: cfg.SuperAgent.BaseURL, + APIKey: cfg.SuperAgent.OpenAPIKey, + ConnectTimeout: cfg.SuperAgent.ConnectTimeout, + RecoveryMaxAttempts: cfg.SuperAgent.RecoveryMaxAttempts, + RecoveryInitialBackoff: cfg.SuperAgent.RecoveryInitialBackoff, + MaxMessageBytes: cfg.SuperAgent.MaxMessageBytes, + }) + if err != nil { + log.Fatalf("create SuperAgent client: %v", err) + } + + suffix, err := randomSuffix() + if err != nil { + log.Fatal("create probe correlation identifiers") + } + correlationID := "fsymd-probe-" + suffix + + signalContext, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + ctx, cancel := context.WithTimeout(signalContext, cfg.SuperAgent.ProbeTimeout) + defer cancel() + + session, err := client.CreateSession(ctx, superagent.CreateSessionRequest{ + ExternalSubjectID: cfg.SuperAgent.ProbeSubjectID, + IdempotencyKey: correlationID + "-session", + RequestID: correlationID, + Metadata: map[string]any{ + "source": "fire-safety-ymd", + "purpose": "connectivity-probe", + }, + }) + if err != nil { + log.Fatalf("create SuperAgent probe session: %v", err) + } + + result, err := client.StreamMessage(ctx, superagent.StreamMessageRequest{ + SessionID: session.ID, + Message: probeMessage, + IdempotencyKey: correlationID + "-message", + RequestID: correlationID, + Metadata: map[string]any{ + "source": "fire-safety-ymd", + "purpose": "connectivity-probe", + }, + }, nil) + if err != nil { + log.Fatalf("stream SuperAgent probe message: %v", err) + } + + output := probeResult{ + Status: "ok", + SessionID: result.SessionID, + RunID: result.RunID, + ProfileID: result.ProfileID, + ProfileVersionID: result.ProfileVersionID, + ModelName: result.ModelName, + Answer: result.Answer, + Usage: result.Usage, + EventTypes: result.EventTypes, + } + encoder := json.NewEncoder(os.Stdout) + encoder.SetIndent("", " ") + if err := encoder.Encode(output); err != nil { + log.Fatalf("write probe result: %v", err) + } +} + +func randomSuffix() (string, error) { + value := make([]byte, 12) + if _, err := rand.Read(value); err != nil { + return "", fmt.Errorf("generate random suffix: %w", err) + } + return base64.RawURLEncoding.EncodeToString(value), nil +} diff --git a/compose.yaml b/compose.yaml new file mode 100644 index 0000000..13a14c1 --- /dev/null +++ b/compose.yaml @@ -0,0 +1,44 @@ +services: + api: + build: + context: . + dockerfile: Dockerfile + image: fire-safety-ymd:test + container_name: fire-safety-ymd + restart: unless-stopped + env_file: + - .env + environment: + # The host binding below keeps this port private; the process must bind + # all container interfaces for Docker's loopback publish to work. + FIRE_SAFETY_HTTP_ADDR: ":8080" + # A one-off owner/migration credential must never enter the long-lived + # application container, even if an operator left it in the local file. + FIRE_SAFETY_POSTGIS_MIGRATION_DSN: "" + TZ: Asia/Shanghai + ports: + - "127.0.0.1:8080:8080" + extra_hosts: + # Use host.docker.internal in the DSN only when PostgreSQL runs on this + # same host; a remote/private database hostname is preferred otherwise. + - "host.docker.internal:host-gateway" + healthcheck: + test: ["CMD", "wget", "-q", "-O", "/dev/null", "http://127.0.0.1:8080/health"] + interval: 10s + timeout: 3s + retries: 6 + start_period: 10s + read_only: true + tmpfs: + - /tmp:size=16m,mode=1777 + security_opt: + - no-new-privileges:true + cap_drop: + - ALL + pids_limit: 256 + stop_grace_period: 20s + logging: + driver: json-file + options: + max-size: "10m" + max-file: "3" diff --git a/deploy/nginx/fire-safety-ymd.conf.example b/deploy/nginx/fire-safety-ymd.conf.example new file mode 100644 index 0000000..f0760f7 --- /dev/null +++ b/deploy/nginx/fire-safety-ymd.conf.example @@ -0,0 +1,133 @@ +# fire-safety-ymd public HTTPS entrypoint (example) +# +# This file is intended to be included from nginx's http context (normally +# /etc/nginx/conf.d/*.conf). Replace the safe app-id in the exact chat +# location with the value configured in FIRE_SAFETY_CHAT_COMPAT_APP_ID. +# Never put FIRE_SAFETY_CHAT_AUTH_TOKEN, FIRE_SAFETY_MCP_AUTH_TOKEN, a +# SuperAgent key, or a database credential in this file. + +limit_req_zone $binary_remote_addr zone=fire_safety_chat:10m rate=5r/s; +limit_req_zone $binary_remote_addr zone=fire_safety_mcp:10m rate=20r/s; + +upstream fire_safety_ymd_backend { + server 127.0.0.1:8080; + keepalive 16; +} + +server { + listen 80; + listen [::]:80; + server_name agent.nianxx.com; + + return 301 https://agent.nianxx.com$request_uri; +} + +server { + listen 443 ssl http2; + listen [::]:443 ssl http2; + server_name agent.nianxx.com; + + ssl_certificate "/cert/agent.nianxx.com.pem"; + ssl_certificate_key "/cert/agent.nianxx.com.key"; + ssl_session_cache shared:SSL:1m; + ssl_session_timeout 10m; + ssl_protocols TLSv1.2 TLSv1.3; + limit_req_status 429; + + # DashScope-compatible user chat. Keep this exact location restricted to + # the one configured app ID; do not replace it with a catch-all regex. + # The incoming xtoken is forwarded unchanged and validated by Go. + location = /api/v1/apps/replace-with-fire-safety-app-id/completion { + limit_req zone=fire_safety_chat burst=20 nodelay; + client_max_body_size 128k; + + proxy_pass http://fire_safety_ymd_backend; + proxy_http_version 1.1; + proxy_set_header Connection ""; + proxy_set_header Host $host; + proxy_set_header X-Real-IP $remote_addr; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + proxy_set_header X-Forwarded-Proto $scheme; + proxy_set_header X-Forwarded-Host $host; + proxy_set_header X-Forwarded-Port $server_port; + proxy_set_header X-Forwarded-Server $host; + proxy_set_header X-Request-ID $request_id; + # X-Accel-Buffering is a response header; the Go handler also sets it. + add_header X-Accel-Buffering no always; + + proxy_buffering off; + proxy_request_buffering off; + proxy_cache off; + gzip off; + proxy_connect_timeout 5s; + proxy_read_timeout 660s; + proxy_send_timeout 660s; + # Do not retry a streaming POST upstream and risk a duplicate run. + proxy_next_upstream off; + proxy_intercept_errors off; + } + + # SuperAgent's inbound MCP callback. Go validates its independent + # Authorization: Bearer token; this proxy deliberately does not inject it. + location = /mcp { + limit_req zone=fire_safety_mcp burst=40 nodelay; + client_max_body_size 256k; + + proxy_pass http://fire_safety_ymd_backend; + proxy_http_version 1.1; + proxy_set_header Connection ""; + proxy_set_header Host $host; + proxy_set_header X-Real-IP $remote_addr; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + proxy_set_header X-Forwarded-Proto $scheme; + proxy_set_header X-Forwarded-Host $host; + proxy_set_header X-Forwarded-Port $server_port; + proxy_set_header X-Forwarded-Server $host; + proxy_set_header X-Request-ID $request_id; + + proxy_buffering on; + proxy_cache off; + proxy_connect_timeout 5s; + proxy_read_timeout 30s; + proxy_send_timeout 30s; + proxy_next_upstream off; + proxy_intercept_errors off; + } + + # Optional public liveness check. It is intentionally not a readiness or + # database check; omit this location if health must remain private. + location = /health { + client_max_body_size 1k; + + proxy_pass http://fire_safety_ymd_backend; + proxy_http_version 1.1; + proxy_set_header Connection ""; + proxy_set_header Host $host; + proxy_set_header X-Real-IP $remote_addr; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + proxy_set_header X-Forwarded-Proto $scheme; + proxy_set_header X-Forwarded-Host $host; + proxy_set_header X-Forwarded-Port $server_port; + proxy_set_header X-Forwarded-Server $host; + proxy_set_header X-Request-ID $request_id; + + proxy_connect_timeout 5s; + proxy_read_timeout 5s; + proxy_send_timeout 5s; + proxy_next_upstream off; + proxy_intercept_errors off; + } + + # Do not expose every Go route through the public hostname. + location / { + default_type application/json; + return 404 '{"error":{"code":"NOT_FOUND","message":"Not found."}}'; + } + + error_page 429 = @rate_limited; + location @rate_limited { + internal; + default_type application/json; + return 429 '{"error":{"code":"RATE_LIMITED","message":"Too many requests."}}'; + } +} diff --git a/docs/architecture/chat-api-v1.md b/docs/architecture/chat-api-v1.md new file mode 100644 index 0000000..394c0d0 --- /dev/null +++ b/docs/architecture/chat-api-v1.md @@ -0,0 +1,82 @@ +# 用户对话 API v1 架构 + +## 1. 组件边界 + +```mermaid +flowchart LR + UI[用户侧应用或 curl] -->|POST /api/chat\nChat Bearer + JSON| H[Chat Handler] + H -->|Prepare / Stream| S[Chat Service] + S -->|CreateSession / StreamMessage| A[SuperAgent Chat Adapter] + A -->|Open API Key + HTTPS/SSE| SA[SuperAgent] + SA -->|独立 MCP Bearer| M[POST /mcp] + M --> P[(PostgreSQL/PostGIS)] +``` + +三种凭证属于不同信任方向: + +- Chat Bearer:用户侧应用到本 Go 服务,仅用于首版受控联调。 +- SuperAgent Open API Key:本 Go 服务到 SuperAgent,只保存在服务端。 +- MCP Bearer:SuperAgent 到本 Go 服务,只保护 `/mcp`。 + +三者必须使用不同值。 + +## 2. 请求生命周期 + +1. Handler 校验请求方法、精确 Origin、Chat Bearer、媒体类型、Accept、请求体大小和严格 JSON。 +2. Service 校验消息和本地 `conversation_id`。 +3. 首轮请求由 Service 生成随机 `conversation_id`,并使用服务端固定测试主体创建 SuperAgent Session。 +4. Service 在内存中保存 `conversation_id -> provider session_id`;Provider Session ID 不返回客户端。 +5. Handler 开始 SSE,先返回 `conversation` 事件。 +6. Adapter 把消息发送到已准备的 Session,只将经过清洗的 `run.*` / `tool.*` 进度投影给 Handler。 +7. Adapter 严格确认最终内容、成功 `run.completed` 和顶层 `end` 后,Handler 才发送 `message` 与 `done`。 +8. 发生断流、超时、协议错误或失败 Run 时,不返回部分回答;当前映射失效,客户端下一次应创建新对话。 + +Handler 在没有可公开进度时每 15 秒发送不含数据的 SSE comment heartbeat,并把单请求写期限扩展到配置的总运行 deadline 之后 5 秒,以便返回终止错误。反向代理仍需显式关闭该路径的响应缓冲并设置相容的空闲超时。 + +## 3. 会话状态机 + +```text +不存在 + -> 创建 Provider Session + -> busy + -> 成功:idle + -> 上游结果不确定或失败:删除映射 + +idle + -> 新一轮 Prepare:busy + -> TTL 到期:惰性删除 + +busy + -> 同 conversation_id 的并发请求:409 +``` + +会话存储是单进程内存映射: + +- 进程重启后全部丢失。 +- 多实例之间不共享,未配置粘性路由时后续轮次可能返回 404。 +- 只保存随机本地 ID、Provider Session ID、占用状态和最后使用时间,不保存消息历史。 +- SuperAgent 负责其 Session 内的上下文;本服务不会把完整历史在每轮重新发送。 + +## 4. 安全设计 + +- 客户端 DTO 没有 `user_id`、`external_subject_id`、角色、镇街范围、Provider Session 或 metadata 字段;未知字段直接拒绝。 +- 固定测试主体由服务端配置,不能作为最终用户身份或权限依据。 +- CORS 使用精确 Origin 列表,不支持 `*` 或 credentials。 +- Handler 日志只记录 request ID、结果、是否复用和耗时,不记录问题、答案或对话 ID。 +- 进度只允许安全字符组成的 `run.*` / `tool.*` 事件,以及工具名和状态;消息 delta、Provider ID、工具参数和结果不向客户端透传。 +- 错误映射为稳定代码,不回显 Provider 响应正文、URL、Session、堆栈或 Secret。 +- API 默认关闭;启用必须同时具备有效 SuperAgent 配置和独立 Chat Bearer。 + +## 5. 伸缩与后续替换点 + +当前内存实现适合单实例受控联调。进入多实例或真实用户阶段前,需要把 `ChatService` 的本地状态替换或扩展为: + +- 经验证的最终用户身份与动态数据授权上下文。 +- 持久化或共享的会话映射,并明确并发租约、过期和恢复语义。 +- 网关限流、滥用防护、请求配额和指标。 +- 对话、工具调用和安全事件的脱敏持久审计。 +- 主动取消 Provider Run,以及客户端断开后的终止策略。 + +这些扩展不能通过简单放宽当前静态 Token 或把客户端用户字段原样传给 Agent 来实现。 + +既有客户端所需的 DashScope 风格协议通过独立 Handler 适配并复用本文 Chat Service,不改变原生契约。公网入口和两种协议映射见 [`public-chat-entry-v1.md`](public-chat-entry-v1.md)。 diff --git a/docs/architecture/public-chat-entry-v1.md b/docs/architecture/public-chat-entry-v1.md new file mode 100644 index 0000000..232efb1 --- /dev/null +++ b/docs/architecture/public-chat-entry-v1.md @@ -0,0 +1,57 @@ +# 公网兼容对话入口 v1 架构 + +## 1. 调用链 + +```mermaid +flowchart LR + C[既有用户客户端] -->|HTTPS + xtoken\nDashScope 风格 JSON/SSE| N[Nginx] + N -->|HTTP 127.0.0.1:8080| D[兼容 Chat Handler] + D -->|ChatRequest / ChatTurn| S[共享 Chat Service] + S -->|Provider-neutral port| A[SuperAgent Adapter] + A -->|Open API Key + HTTPS/SSE| SA[SuperAgent] + SA -->|独立 MCP Bearer| M[同域 /mcp] + M --> P[(PostgreSQL/PostGIS)] +``` + +Nginx 不再直接调用 DashScope。它只负责 TLS、精确公开路径、基础限流和 SSE 传输设置;所有应用鉴权、输入校验、会话映射和上游调用都在 Go 服务内完成。 + +## 2. 双入站协议、单一用例 + +| 入站接口 | 面向对象 | 鉴权 Header | 请求字段 | 成功事件 | +| --- | --- | --- | --- | --- | +| `/api/chat` | 本项目原生客户端 | `Authorization: Bearer` | `message`、`conversation_id` | `conversation`、`progress`、`message`、`done` | +| `/api/v1/apps/{app_id}/completion` | 既有 DashScope 风格客户端 | `xtoken` | `input.prompt`、`input.session_id`、受限 `parameters` | `result`,最终 `finish_reason=stop` | + +两个 Handler 都调用同一个 `ChatService.Prepare` 和 `ChatTurn.Stream`。兼容层只负责协议转换,不复制会话或 SuperAgent 业务逻辑。 + +## 3. 标识映射 + +```text +兼容 output.session_id + = 本地 conversation_id + -> Chat Service 内存映射 + -> Provider session_id(永不返回客户端) +``` + +兼容 URL 中的 `app_id` 是配置的公开路由标识。它不能选择任意 Profile,也不参与授权;SuperAgent 目标仍由服务端 Base URL、Key 和已发布 Profile 决定。 + +## 4. 严格结果边界 + +兼容 Handler 会尽早发送一个无正文的 `result` 事件,让客户端取得会话 ID并建立 SSE。正文不会跟随上游 `message.delta` 实时透传。只有既有 Adapter 验证最终内容、成功 Run 和流终止标记后,Handler 才发送 `finish_reason=stop` 与最终文本。 + +这样保留了当前应急辅助场景的失败语义:断流或 Provider 协议不完整时,客户端不会把半截模型文本误当成已完成方案。代价是首版没有逐字动画;如后续必须实时输出,需要独立安全决策、取消/失败 UX 和新 Spec。 + +## 5. 网络与伸缩边界 + +- Go 推荐只监听 `127.0.0.1:8080`,公网只开放 Nginx 443;PostgreSQL 5432 不对公网开放。 +- Nginx 对兼容路径关闭缓冲、缓存、gzip 和重试,超时必须覆盖 Go Chat Run Timeout。 +- `/mcp` 使用独立 Bearer;Chat `xtoken`、MCP Bearer 与 SuperAgent Open API Key 三者不得复用。 +- 当前会话在单进程内存中。多实例部署必须先实现共享会话映射或粘性路由;否则后续轮次可能落到另一实例并返回 404。 +- 浏览器可读取静态 `xtoken`,所以该入口仍只适用于受控联调;生产最终用户入口需要真实身份认证和动态授权。 + +## 6. 相关文档 + +- [`../specs/fire-safety-ymd-dashscope-compatible-chat-v1.md`](../specs/fire-safety-ymd-dashscope-compatible-chat-v1.md) +- [`chat-api-v1.md`](chat-api-v1.md) +- [`../workflows/user-chat.md`](../workflows/user-chat.md) +- [`../project/operations/nginx-public-entry.md`](../project/operations/nginx-public-entry.md) diff --git a/docs/architecture/spatial-mcp-v1.md b/docs/architecture/spatial-mcp-v1.md new file mode 100644 index 0000000..3efd2ff --- /dev/null +++ b/docs/architecture/spatial-mcp-v1.md @@ -0,0 +1,73 @@ +# 空间只读 MCP v1 架构 + +## 决策摘要 + +首版 MCP 内嵌在现有 Go 服务中,使用标准库实现 HTTP/JSON-RPC 协议层,使用 `pgx/v5` 原生连接池访问 PostgreSQL/PostGIS。没有引入 MCP SDK、Web 框架或 ORM。 + +采用这一边界是因为当前只需要 SuperAgent 已验证的 `2025-06-18` 四个方法和同步 JSON 响应;业务复杂度位于固定空间查询、权限范围和安全语义,而不是协议框架。 + +## 模块关系 + +```mermaid +flowchart TD + APP["internal/app\n依赖装配与 readiness"] + H["internal/handler\nBearer、JSON-RPC、schema、限流边界"] + S["internal/service\n查询边界、超时、结果语义"] + D["internal/domain\n稳定空间领域对象"] + R["internal/repository\npgxpool、固定参数化 PostGIS SQL"] + DB["8 张既有 PostGIS 表"] + + APP --> H + APP --> S + APP --> R + H --> S + S --> D + S --> R + R --> D + R --> DB +``` + +Handler 不知道物理表名,Repository 不组织自然语言回答,Domain 不依赖 MCP 或 pgx。可信数据范围在应用装配时由配置传入 Service,并由 Repository 在 SQL 中应用:默认 `town_allowlist` 使用参数化镇街数组过滤;显式 `all` 使用服务端布尔参数放开镇街过滤。范围选择不进入工具 schema,模型无法扩大权限。 + +## 数据映射 + +| 领域结果 | 物理表 | +| --- | --- | +| 地名候选 | 以下 8 张表的名称、镇街、村庄和几何字段 | +| 网格上下文、责任中队 | `st_2_fanghuowangge` | +| 水源候选 | `st_2_mpslfh_t_slfh_syd`、`st_2_xianyouxushuichiguan` | +| 指挥部候选设施 | `st_2_fanghuojianchazhan`、`st_2_fanghuoliaowangshao` | +| 通道候选 | `st_2_xianyoufanghuotongdao` | +| 风险区域 | `st_2_mudifenqu_mian`、`st_2_linqugongkuangqiye` | + +姓名、联系电话和值班人员字段不进入领域结果,也没有出现在查询 SELECT 列表中。 + +## 空间约束 + +- MCP 入参固定为 WGS84 经纬度。 +- 地名搜索是坐标查询前的候选发现:记录点直接返回二维坐标;线使用首个组成线的起点,面使用 `ST_PointOnSurface` 产生代表点,并以 `location_kind` 明确区分。代表点不得自动升级为用户确认的演练点。 +- 防火通道原始 `geom` 已确认混合二维与 Z 维度,导入层必须使用不限定 typmod 的 `geometry` 保存原始事实。距离/最近点等 MCP v1 运算按二维地表语义解释;若底层函数需要降维,只能在只读查询或派生层显式处理,不得回写或静默修改原始几何。 +- 启用前要求实库所有非空几何 SRID 为 4326、类型符合表用途且坐标位于 WGS84 合法范围;这些属于硬门禁。 +- 点到点、点到线、点到面距离使用 PostGIS `geography` 米制计算。 +- 面覆盖使用 `ST_Covers`,边界上的点也视为位于网格/风险区内。 +- SQL 同时检查 SRID、类型和有效性。无效或空几何不自动修复、不改写原始事实,而是从全部 MCP 查询中排除;readiness 和工具结果明确告警结果可能不完整。readiness 发现异常类型、非 4326 SRID 或越界坐标时仍拒绝启用。 +- 2026-09-05 实库 audit 发现防火网格 4 条、林区工矿企业 4 条、墓地坟区 27 条无效面几何。用户选择首版排除这 35 条记录,后续如需修复必须在派生副本中审查,不覆盖原始 `geom`。 +- 距离结果只是地理邻近,不包含地形、路网、火势、天气和实时通行信息。 + +## 依赖决策 + +`github.com/jackc/pgx/v5 v5.10.0` 是当前唯一新增的直接依赖。项目只面向 PostgreSQL,并需要明确的连接池、context 和 PostgreSQL 参数行为,因此使用原生 `pgxpool`;空间计算仍全部由参数化 SQL/PostGIS 完成。 + +## 后续演进边界 + +- 动态用户/组织授权到位后,用可信身份解析器替换当前每 Token 的静态数据库全范围/镇街白名单,工具 schema 不增加可伪造的授权字段。 +- 路线规划必须新增专门的图网络/地形数据和高风险 Spec,不能扩写当前通道候选工具的描述来冒充路线能力。 +- 联系人字段如确需开放,必须有字段级权限、脱敏、审计和单独工具,不直接扩展当前结果。 +- 写入、派遣或状态变更工具需要人工确认、幂等、审计和独立安全评审。 + +## 协议与依赖依据 + +- [MCP 2025-06-18 Lifecycle](https://modelcontextprotocol.io/specification/2025-06-18/basic/lifecycle) +- [MCP 2025-06-18 Streamable HTTP Transport](https://modelcontextprotocol.io/specification/2025-06-18/basic/transports) +- [MCP 2025-06-18 Tools](https://modelcontextprotocol.io/specification/2025-06-18/server/tools) +- [pgx 官方仓库与版本策略](https://github.com/jackc/pgx) diff --git a/docs/import/db-samples/README.md b/docs/import/db-samples/README.md new file mode 100644 index 0000000..5d3c4d8 --- /dev/null +++ b/docs/import/db-samples/README.md @@ -0,0 +1,34 @@ +# PostgreSQL/PostGIS 本地样例说明 + +本目录的 `*.sql` 是用户提供的本地结构与少量数据样例,仅用于核对字段映射和设计查询。 + +- SQL 导出可能包含 `DROP TABLE`、真实联系人、电话和精确坐标。 +- 应用、测试和开发 Agent 不得执行这些 SQL。 +- `*.sql` 已由仓库根 `.gitignore` 忽略,避免误提交敏感数据。 +- 如需长期纳入版本控制,应另行生成脱敏 fixture:删除姓名/电话、替换精确坐标并移除所有 DDL,只保留最小测试字段。 +- 实库结构、SRID、几何类型、有效性、索引和数据时效必须通过只读 readiness probe 核验,不能由两条样例推定。 + +当前每表两条正常样例已足够确定 MCP v1 字段映射,不需要继续增加同类真实记录。后续若制作可提交的脱敏测试 fixture,优先补充:空几何、边界点、多网格重叠、无结果、不同状态枚举、空容量/时间、LineString 与 MultiLineString、Polygon 与 MultiPolygon,而不是更多真实联系人或生产坐标。 + +## 已确认的防火通道导入例外 + +2026-09-04 用户现场重导 `st_2_xianyoufanghuotongdao` 时,原导出 DDL 的 `geom geometry(GEOMETRY)` 在约第 256 条遇到带 Z 维度的记录并报 PostgreSQL/PostGIS `22023: Geometry has Z dimension but column does not`。此前可插入的二维记录与该错误共同确认:源文件包含混合二维/Z 几何。用户将原始导入列改为不限定 typmod 的 `geometry` 后报告导入成功。 + +后续 Agent 必须遵守: + +- 原始导入层保留 `geom geometry`,不要无审查改回二维 `geometry(GEOMETRY)`。 +- 不得为“导入成功”而对原始数据执行 `ST_Force2D`、改写 WKB 或丢弃 Z;二维算法需要时只在只读查询或派生层显式转换。 +- Z 维度与 SRID 无关。该次导入没有证明 EPSG:4326,也没有代替最终行数、维度分布和数据质量核验。 + +获得实库只读连接后,可先用以下非敏感聚合确认分布: + +```sql +SELECT + ST_NDims(geom) AS dimensions, + GeometryType(geom) AS geometry_type, + COUNT(*) AS row_count +FROM public.st_2_xianyoufanghuotongdao +WHERE geom IS NOT NULL +GROUP BY ST_NDims(geom), GeometryType(geom) +ORDER BY dimensions, geometry_type; +``` diff --git a/docs/import/reusable/ai-native-software-engineering-standard.md b/docs/import/reusable/ai-native-software-engineering-standard.md new file mode 100644 index 0000000..651efcf --- /dev/null +++ b/docs/import/reusable/ai-native-software-engineering-standard.md @@ -0,0 +1,350 @@ +# AI-Native Software Engineering Standard (AI-NSES) + +| 项 | 内容 | +| --- | --- | +| Version | 0.2 | +| Scope | 可复用软件工程标准 | +| Audience | 人类开发者、产品人员、架构师、AI Agent | + +## 1. Purpose + +AI-NSES 是一套面向 AI 协作开发的软件工程标准。 + +它不属于某一个具体项目,而是描述: + +- 项目应该如何组织。 +- 文档应该如何维护。 +- AI Agent 应该如何工作。 +- 开发流程应该如何闭环。 + +目标是让任何新的 AI Agent 在几分钟内理解项目,而不依赖历史聊天记录。 + +## 2. Design Philosophy + +### Documentation First + +复杂软件首先是知识,其次才是代码。代码是项目知识的一种实现形式。 + +### Context Driven + +Prompt 是临时的,Context 是长期资产。重要知识应该沉淀在仓库,而不是沉淀在聊天记录里。 + +### Living Documentation + +文档不是一次性产物。文档随着项目成长,开发结束时文档也应该同步结束。 + +### AI as Team Member + +AI 不是单纯的代码生成器,而是项目成员。它可以参与需求分析、产品设计、架构设计、开发、Review 和维护。 + +## 3. Core Principle + +Build a project that teaches AI. + +Do not teach AI every day. + +中文说明:让项目成为 AI 的长期记忆,而不是每次都重新解释项目。 + +## 4. Recommended Project Structure + +```text +project/ +├── AGENTS.md +├── CONTEXT.md +├── PROJECT_STATE.md +├── docs/ +│ ├── domain/ +│ ├── architecture/ +│ ├── workflows/ +│ ├── adr/ +│ ├── specs/ +│ └── guidelines/ +├── src/ +└── tests/ +``` + +中文说明:实际项目可以使用 `client/`、`server/`、`backend/`、`frontend/` 等目录,只要在 `CONTEXT.md` 或项目架构文档中说明映射关系即可。 + +## 5. Document Responsibilities + +### AGENTS.md + +读者:所有 AI Agent。 + +职责:规定 Agent 如何工作,包括默认工作流程、Debug 原则、Spec 原则、文档更新原则和 Review 原则。 + +更新频率:极低。 + +### CONTEXT.md + +读者:人类成员和 AI Agent。 + +职责:介绍项目整体背景,包括产品目标、技术栈、系统组成、当前模块和当前开发方向。 + +更新频率:低。 + +### PROJECT_STATE.md + +读者:人类成员和 AI Agent。 + +职责:记录当前项目状态,包括当前 Sprint、当前 Feature、当前 Priority、Known Issues 和 Next Steps。 + +更新频率:高。建议每完成一个 Feature 或 Checkpoint 更新一次。 + +### docs/domain/ + +读者:产品、业务、开发和 AI Agent。 + +职责:记录业务知识。建议一个业务对象一个文档,例如 `Guest.md`、`Hotel.md`、`Order.md`、`Reservation.md`、`Email.md`、`Task.md`、`PMS.md`。 + +内容包括定义、生命周期、业务规则和关系。禁止记录实现细节。 + +更新频率:低。 + +### docs/architecture/ + +读者:架构、开发和 AI Agent。 + +职责:记录系统设计,例如 `ARCHITECTURE.md`、`DATABASE.md`、`EVENTS.md`、`MODULES.md`。 + +重点描述模块边界、系统通信、数据库、事件和部署边界。 + +更新频率:低。 + +### docs/workflows/ + +读者:产品、业务、开发和 AI Agent。 + +职责:记录业务流程,例如 Email Processing、Reservation Sync、Task Generation、Webhook Processing。 + +重点描述输入、输出、状态流转、异常和补偿。 + +更新频率:中。 + +### docs/adr/ + +读者:架构、开发和 AI Agent。 + +职责:记录重要架构决策。每个重要设计决策一份文档。 + +建议格式:背景、为什么、备选方案、最终选择、影响。 + +规则:ADR 永远追加,不覆盖历史。 + +### docs/specs/ + +读者:产品、开发、测试和 AI Agent。 + +职责:记录具体功能规格。每一个 Feature 对应一个 Spec。 + +生命周期:Draft -> Approved -> Implemented -> Superseded / Archived。 + +规则:Spec 完成以后保留,不删除。 + +如果既有项目已经有 `docs/project/requirements/` 等历史目录,可以通过项目级 adoption 文档建立映射;但新增重要 Feature 仍应使用 Spec 模板的结构表达背景、目标、非目标、业务规则、接口或交互契约、验收标准和测试范围。 + +### docs/guidelines/ + +读者:开发、测试和 AI Agent。 + +职责:记录长期规范,例如 API、UI、CODING、TESTING、SECURITY。 + +更新频率:极低。 + +## 6. Knowledge Layers + +### Long-term Knowledge + +生命周期:整个项目。 + +包括 Vision、Domain、Architecture、Guidelines。 + +### Medium-term Knowledge + +生命周期:一个版本或一个阶段。 + +包括 Workflow、ADR、Spec。 + +### Short-term Knowledge + +生命周期:一个 Sprint 或一个开发周期。 + +包括 Current Sprint、Known Issues、Next Steps、Backlog。 + +## 7. Stable vs Dynamic Documents + +长期稳定: + +- `AGENTS.md` +- `CONTEXT.md` +- `docs/domain/` +- `docs/architecture/` +- `docs/guidelines/` + +中期演进: + +- `docs/workflows/` +- `docs/adr/` +- `docs/specs/` + +高频更新: + +- `PROJECT_STATE.md` + +中文说明:这样可以避免整个仓库每天发生无意义变化,也能让新 Agent 快速判断哪些文档代表长期事实,哪些文档代表当前状态。 + +## 8. Feature Lifecycle + +任何 Feature 默认按以下顺序推进: + +```text +Idea +-> Requirement +-> Discussion +-> Specification +-> Implementation +-> Verification +-> Documentation Update +-> Done +``` + +Documentation Update 属于 Definition of Done,不能跳过。 + +### 8.1 Definition of Ready + +进入实现前,复杂 Feature 或会影响跨模块协作的变更必须满足: + +- 需求来源明确,知道是谁提出、解决什么问题。 +- 目标和非目标明确,避免实现时扩大范围。 +- 受影响的用户、业务流程、接口、数据模型、权限、安全、审计或外部系统边界已经列出。 +- 前端、后端、测试或第三方系统的分工已经明确。 +- 验收标准已经可测试,至少有主要 Given / When / Then 场景。 +- 未确认问题已经列出;如果问题影响业务规则或安全边界,应先确认再实现。 + +简单 bugfix 或纯内部重构可以不用新建完整 Spec,但仍应在任务说明中写清目标、范围和验证方式。 + +### 8.2 Definition of Done + +Feature 完成前必须确认: + +- 实现满足本次 Spec、Change Request 或任务说明。 +- 相关自动化测试、类型检查、lint、构建或手工验证已经运行,或明确说明无法运行的原因。 +- 受影响的需求、Spec、Workflow、Domain、ADR、接口契约、安全边界和 Project State 已同步。 +- 对外或跨团队契约发生变化时,调用方文档和测试说明已同步。 +- 没有把 Secret、真实客户数据、构建产物、临时文件或无关本地变更纳入交付。 + +### 8.3 Change Request + +当一个已 Approved 或 Implemented 的 Spec 发生需求变更时,优先新增或更新 Change Request,而不是把讨论散落到聊天记录中。 + +Change Request 至少说明: + +- 变更背景。 +- 原规则。 +- 新规则。 +- 影响范围。 +- 迁移或兼容策略。 +- 前端、后端、测试和文档待办。 +- 验收标准。 + +小变更可以直接追加到原 Spec 的“变更记录”章节;跨前后端、权限、安全、数据模型或第三方契约的变更应单独成文。 + +### 8.4 Traceability Matrix + +复杂 Feature 应维护需求追踪表,用于连接需求、实现、测试和文档。 + +推荐字段: + +| 需求项 | 后端状态 | 前端状态 | 测试状态 | 文档位置 | 当前状态 | +| --- | --- | --- | --- | --- | --- | +| 示例需求 | Pending / Done / N/A | Pending / Done / N/A | Pending / Done / Blocked | Spec / Contract / README | Draft / Approved / Implemented | + +追踪表可以放在 Spec、Project State 或专项 checkpoint 文档中;关键是让新 Agent 能快速判断“文档写了但代码没做、代码做了但文档没写、后端做了但前端没做、前端需要但后端没做”。 + +### 8.5 Agent Handoff + +给 AI Agent 分派任务时,建议使用统一交接结构: + +- 背景:为什么做。 +- 目标:本次必须完成什么。 +- 边界:本次不做什么。 +- 必读文档:从入口文档到具体 Spec / Contract。 +- 影响范围:后端、前端、测试、第三方、数据、权限、安全、审计。 +- 验收标准:可执行或可观察的检查项。 +- 允许写操作:是否允许改代码、改文档、造数据、点确认按钮或清理测试数据。 +- 输出要求:代码位置、测试结果、文档更新清单、风险和下一步建议。 + +## 9. AI Working Principles + +- AI 应先理解,再开发。 +- 复杂需求默认先讨论,不立即编码。 +- 复杂功能默认先 Spec,不直接实现。 +- Bug 默认先定位,不猜测修复。 +- UI 默认保持一致性,不过度设计。 +- 代码修改前先确认目标、边界和验收标准。 +- 涉及接口、安全、权限、数据模型或外部系统时,先读相关契约文档。 +- 涉及跨 agent 协作时,先确认 Spec、Change Request 或 handoff 是否足够清楚。 + +## 10. Documentation Rules + +所有文档必须回答一个问题: + +未来的新 Agent 为什么需要阅读它? + +如果回答不了,就不要创建,也不要维护。 + +文档应该小、独立、易维护。 + +不要维护一个 8000 行的 `DOMAIN.md`。应该按对象拆分成 `Order.md`、`Guest.md`、`Hotel.md`、`Email.md`、`Task.md`。 + +通用标准只规定文档职责和流程。具体项目可以新增 Project Overlay,记录本项目专属的需求门禁、核心概念、交付格式和安全边界;Overlay 不应反向污染通用标准。 + +## 11. Documentation Update Rules + +完成 Feature 后必须检查: + +- Domain 是否需要更新。 +- Architecture 是否需要更新。 +- Workflow 是否需要更新。 +- ADR 是否需要新增。 +- Spec 是否需要改为 Implemented 或补充结果。 +- Project State 是否需要更新。 +- 安全、权限、接口契约是否需要同步。 +- Traceability Matrix 是否需要更新。 +- Change Request 是否需要关闭、合并到 Spec 或标记为 Superseded。 + +如果没有文档变化,应明确说明: + +```text +No documentation changes required. +``` + +不要为了修改而修改文档。 + +## 12. Success Criteria + +一个新的 AI Agent 进入项目后,阅读以下有限文档即可开始工作: + +```text +AGENTS.md +-> CONTEXT.md +-> PROJECT_STATE.md +-> 项目级 Overlay 或 Adoption 文档 +-> 相关 Domain +-> 相关 Workflow +-> 相关 Spec 或 ADR +``` + +无需阅读整个代码库。 + +无需依赖历史聊天记录。 + +## 13. Scope + +AI-NSES 不限制编程语言、框架、数据库或 AI 模型。 + +它适用于 Codex、Claude Code、Gemini CLI、Cursor,以及未来任何 AI Agent。 + +它描述的是软件工程,不是某个具体工具。 + +具体项目如果有更强的安全、合规、业务流程或多 agent 协作要求,应通过项目级 Overlay 补充,不直接把项目私有业务规则写入本通用标准。 diff --git a/docs/import/reusable/ai-native-templates/ADR.template.md b/docs/import/reusable/ai-native-templates/ADR.template.md new file mode 100644 index 0000000..3fa85ee --- /dev/null +++ b/docs/import/reusable/ai-native-templates/ADR.template.md @@ -0,0 +1,40 @@ +# ADR-编号 标题 + +| 项 | 内容 | +| --- | --- | +| 状态 | Proposed / Accepted / Superseded | +| 日期 | YYYY-MM-DD | +| 决策人 | 按实际填写 | + +## 1. 背景 + +说明为什么需要做这个决策。 + +## 2. 约束 + +- 约束 1: +- 约束 2: + +## 3. 备选方案 + +### 方案 A + +说明优点和缺点。 + +### 方案 B + +说明优点和缺点。 + +## 4. 最终选择 + +说明最终选择哪个方案。 + +## 5. 影响 + +- 正面影响: +- 负面影响: +- 后续动作: + +## 6. 历史说明 + +ADR 只追加,不覆盖历史。如果未来决策变化,新建 ADR 或标记被替代。 diff --git a/docs/import/reusable/ai-native-templates/AGENTS.template.md b/docs/import/reusable/ai-native-templates/AGENTS.template.md new file mode 100644 index 0000000..2018063 --- /dev/null +++ b/docs/import/reusable/ai-native-templates/AGENTS.template.md @@ -0,0 +1,39 @@ +# 项目协作与开发规范 + +## 1. 文档入口 + +- 项目背景:`CONTEXT.md` +- 当前状态:`PROJECT_STATE.md` +- 项目文档索引:`docs/README.md` 或项目自定义索引 +- 通用开发规范:按项目实际路径填写 + +## 2. 工作方式 + +- 先确认目标、边界和验收标准,再改代码或文档。 +- 每次只做一个明确 checkpoint。 +- 不修改与当前任务无关的用户变更。 +- 不回滚用户自己的改动,除非用户明确要求。 +- 遇到不确定的技术栈、接口契约、权限边界或数据模型,先确认再继续。 + +## 3. 文档更新原则 + +完成 Feature 后必须检查: + +- Domain 是否需要更新。 +- Architecture 是否需要更新。 +- Workflow 是否需要更新。 +- ADR 是否需要新增。 +- Spec 是否需要更新状态。 +- Project State 是否需要更新。 + +如果没有变化,明确说明: + +```text +No documentation changes required. +``` + +## 4. 测试与验证 + +- 修改后运行对应模块已有检查命令。 +- 如果检查命令尚未配置或因环境问题无法运行,必须明确说明原因。 +- 不假装测试通过。 diff --git a/docs/import/reusable/ai-native-templates/AGENT_HANDOFF.template.md b/docs/import/reusable/ai-native-templates/AGENT_HANDOFF.template.md new file mode 100644 index 0000000..d40c8bd --- /dev/null +++ b/docs/import/reusable/ai-native-templates/AGENT_HANDOFF.template.md @@ -0,0 +1,57 @@ +# Agent Handoff 标题 + +## 1. 背景 + +说明本次任务从哪里来,要解决什么问题。 + +## 2. 目标 + +- 必须完成事项 1: +- 必须完成事项 2: + +## 3. 边界 + +- 不做事项 1: +- 不做事项 2: + +## 4. 必读文档 + +1. `AGENTS.md` +2. `CONTEXT.md` +3. `PROJECT_STATE.md` +4. 项目文档索引 +5. 本次相关 Spec / Change Request / Contract + +## 5. 影响范围 + +| 范围 | 是否影响 | 说明 | +| --- | --- | --- | +| 后端 | 是 / 否 | | +| 前端 | 是 / 否 | | +| 数据库 | 是 / 否 | | +| 权限 / 安全 / 审计 | 是 / 否 | | +| 第三方系统 | 是 / 否 | | +| 测试数据 | 是 / 否 | | +| 文档 | 是 / 否 | | + +## 6. 验收标准 + +- Given / When / Then: +- Given / When / Then: + +## 7. 允许操作 + +- 是否允许改代码: +- 是否允许改文档: +- 是否允许造测试数据: +- 是否允许执行写操作: +- 是否允许清理数据: + +## 8. 输出要求 + +- 完成内容: +- 未完成内容: +- 代码或文档位置: +- 测试命令和结果: +- 风险: +- 下一步建议: diff --git a/docs/import/reusable/ai-native-templates/ARCHITECTURE.template.md b/docs/import/reusable/ai-native-templates/ARCHITECTURE.template.md new file mode 100644 index 0000000..b8d40f5 --- /dev/null +++ b/docs/import/reusable/ai-native-templates/ARCHITECTURE.template.md @@ -0,0 +1,31 @@ +# 架构说明 + +## 1. 系统目标 + +说明系统架构服务的产品目标和主要约束。 + +## 2. 模块边界 + +| 模块 | 职责 | 不负责 | +| --- | --- | --- | +| | | | + +## 3. 依赖方向 + +说明模块之间的依赖方向,避免双向依赖和跨层直连。 + +## 4. 数据边界 + +说明数据库、缓存、文件存储、消息队列和外部系统的数据边界。 + +## 5. 安全边界 + +说明鉴权、授权、租户隔离、审计和敏感数据处理方式。 + +## 6. 关键决策 + +列出相关 ADR 链接。 + +## 7. 演进计划 + +说明后续可能调整的方向和触发条件。 diff --git a/docs/import/reusable/ai-native-templates/CHANGE_REQUEST.template.md b/docs/import/reusable/ai-native-templates/CHANGE_REQUEST.template.md new file mode 100644 index 0000000..91f202a --- /dev/null +++ b/docs/import/reusable/ai-native-templates/CHANGE_REQUEST.template.md @@ -0,0 +1,61 @@ +# Change Request 标题 + +| 项 | 内容 | +| --- | --- | +| 状态 | Draft / Approved / Implemented / Superseded / Archived | +| 日期 | YYYY-MM-DD | +| 提出人 | 按实际填写 | +| 关联 Spec | 路径或编号 | +| 影响范围 | Backend / Frontend / Test / Docs / Security / Integration | + +## 1. 变更背景 + +说明为什么要改,当前问题是什么。 + +## 2. 原规则 + +- 原规则 1: +- 原规则 2: + +## 3. 新规则 + +- 新规则 1: +- 新规则 2: + +## 4. 影响范围 + +| 范围 | 是否影响 | 说明 | +| --- | --- | --- | +| 用户流程 | 是 / 否 | | +| 后端接口 | 是 / 否 | | +| 前端交互 | 是 / 否 | | +| 数据模型 | 是 / 否 | | +| 权限 / 安全 / 审计 | 是 / 否 | | +| 第三方契约 | 是 / 否 | | +| 测试数据 / 迁移 | 是 / 否 | | + +## 5. 兼容和迁移 + +- 是否兼容旧数据: +- 是否需要数据清理: +- 是否影响生产: +- 回滚或恢复方式: + +## 6. 验收标准 + +- Given / When / Then: +- Given / When / Then: + +## 7. Agent 待办 + +| Agent / 角色 | 待办 | 验证 | +| --- | --- | --- | +| 后端 | | | +| 前端 | | | +| 测试 | | | +| 文档 | | | + +## 8. 文档更新 + +- 需要更新的文档: +- 不需要更新的文档及原因: diff --git a/docs/import/reusable/ai-native-templates/CONTEXT.template.md b/docs/import/reusable/ai-native-templates/CONTEXT.template.md new file mode 100644 index 0000000..ea4fed6 --- /dev/null +++ b/docs/import/reusable/ai-native-templates/CONTEXT.template.md @@ -0,0 +1,50 @@ +# 项目上下文 + +## 1. 项目目标 + +说明项目要解决什么问题、主要服务谁、成功后用户会得到什么价值。 + +## 2. 当前系统组成 + +| 模块 | 中文说明 | +| --- | --- | +| `frontend/` 或 `client/` | 前端应用,负责展示、交互和调用本项目后端。 | +| `backend/` 或 `server/` | 后端服务,负责业务规则、数据、权限、安全和外部系统适配。 | +| `docs/` | 项目文档、规范、需求和架构说明。 | + +## 3. 技术栈 + +### 后端 + +- 语言: +- 框架: +- 数据库: +- 构建工具: +- 测试工具: + +### 前端 + +- 框架: +- 构建工具: +- UI 组件: +- 测试工具: + +## 4. 业务领域 + +列出核心业务对象,例如用户、订单、任务、酒店、邮件、支付、库存等。 + +## 5. 外部系统 + +列出外部系统、调用方向、鉴权方式和接口契约位置。 + +## 6. 当前开发方向 + +说明当前阶段最重要的开发目标和不做的事情。 + +## 7. 新 Agent 阅读顺序 + +1. `AGENTS.md` +2. `CONTEXT.md` +3. `PROJECT_STATE.md` +4. 项目文档索引 +5. 与当前任务相关的 Domain、Workflow、Spec、ADR diff --git a/docs/import/reusable/ai-native-templates/DOMAIN_OBJECT.template.md b/docs/import/reusable/ai-native-templates/DOMAIN_OBJECT.template.md new file mode 100644 index 0000000..616c1ea --- /dev/null +++ b/docs/import/reusable/ai-native-templates/DOMAIN_OBJECT.template.md @@ -0,0 +1,27 @@ +# 业务对象名称 + +## 1. 定义 + +说明这个业务对象是什么,不是什么。 + +## 2. 生命周期 + +说明对象从创建到结束的主要状态。 + +## 3. 业务规则 + +- 规则 1: +- 规则 2: +- 规则 3: + +## 4. 关系 + +说明它和其他业务对象的关系。 + +## 5. 禁止混淆 + +列出容易和它混淆的概念。 + +## 6. 非目标 + +说明本文不记录实现细节、表结构或接口字段。实现细节应放在 Architecture、Spec 或代码中。 diff --git a/docs/import/reusable/ai-native-templates/PROJECT_STATE.template.md b/docs/import/reusable/ai-native-templates/PROJECT_STATE.template.md new file mode 100644 index 0000000..689167b --- /dev/null +++ b/docs/import/reusable/ai-native-templates/PROJECT_STATE.template.md @@ -0,0 +1,39 @@ +# 项目当前状态 + +| 项 | 内容 | +| --- | --- | +| 最近更新 | YYYY-MM-DD | +| 当前分支 | 按实际填写 | +| 当前阶段 | 按实际填写 | +| 当前重点 | 按实际填写 | + +## 1. 当前 Feature 或 Checkpoint + +- 名称: +- 状态:Draft / In Progress / Blocked / Ready for Review / Done +- 目标: +- 验收标准: + +## 2. 当前优先级 + +1. 第一优先级: +2. 第二优先级: +3. 第三优先级: + +## 3. 已确认事实 + +- 记录对后续开发有影响的当前事实。 +- 只写仍然有效的事实,不写长篇历史。 + +## 4. Known Issues + +- 记录当前已知问题、风险和待验证点。 + +## 5. Next Steps + +- 下一步最小动作。 +- 下一个建议 checkpoint。 + +## 6. 文档同步提醒 + +完成 Feature 后检查 Domain、Architecture、Workflow、ADR、Spec、Project State 是否需要更新。 diff --git a/docs/import/reusable/ai-native-templates/README.md b/docs/import/reusable/ai-native-templates/README.md new file mode 100644 index 0000000..b399d16 --- /dev/null +++ b/docs/import/reusable/ai-native-templates/README.md @@ -0,0 +1,23 @@ +# AI-NSES 模板目录 + +本目录保存可复制到新项目的 AI-NSES 文档模板。 + +使用方式: + +1. 复制需要的模板到新项目对应位置。 +2. 删除模板中的示例说明。 +3. 补充项目真实信息。 +4. 在项目 `AGENTS.md` 和项目文档索引中加入入口。 + +模板清单: + +- `AGENTS.template.md`:AI Agent 协作入口模板。 +- `CONTEXT.template.md`:项目背景入口模板。 +- `PROJECT_STATE.template.md`:项目当前状态模板。 +- `DOMAIN_OBJECT.template.md`:业务对象文档模板。 +- `WORKFLOW.template.md`:业务流程文档模板。 +- `ADR.template.md`:架构决策记录模板。 +- `SPEC.template.md`:功能规格模板。 +- `CHANGE_REQUEST.template.md`:已确认需求的变更请求模板。 +- `AGENT_HANDOFF.template.md`:给后端、前端、测试或文档 agent 的任务交接模板。 +- `ARCHITECTURE.template.md`:架构说明模板。 diff --git a/docs/import/reusable/ai-native-templates/SPEC.template.md b/docs/import/reusable/ai-native-templates/SPEC.template.md new file mode 100644 index 0000000..68478f6 --- /dev/null +++ b/docs/import/reusable/ai-native-templates/SPEC.template.md @@ -0,0 +1,75 @@ +# Feature Spec 标题 + +| 项 | 内容 | +| --- | --- | +| 状态 | Draft / Approved / Implemented / Superseded / Archived | +| 日期 | YYYY-MM-DD | +| 负责人 | 按实际填写 | +| 需求来源 | 用户 / 客户 / 业务方 / 内部发现 | +| 关联 Change Request | 可为空 | + +## 1. 背景 + +说明为什么要做这个功能。 + +## 2. 目标 + +- 目标 1: +- 目标 2: + +## 3. 非目标 + +- 不做事项 1: +- 不做事项 2: + +## 4. 用户与场景 + +说明谁会使用这个能力,在哪些场景使用。 + +## 5. Definition of Ready + +- 需求来源已确认: +- 目标和非目标已确认: +- 影响范围已确认: +- 权限、安全、审计和数据边界已确认: +- 前后端 / 测试分工已确认: +- 未确认问题已列出: + +## 6. 业务规则 + +- 规则 1: +- 规则 2: + +## 7. 接口或交互契约 + +说明请求、响应、权限、安全、审计和兼容性要求。 + +## 8. 需求追踪表 + +| 需求项 | 后端状态 | 前端状态 | 测试状态 | 文档位置 | 当前状态 | +| --- | --- | --- | --- | --- | --- | +| 需求 1 | Pending / Done / N/A | Pending / Done / N/A | Pending / Done / Blocked | | Draft / Approved / Implemented | + +## 9. 验收标准 + +- Given / When / Then: +- Given / When / Then: + +## 10. 测试范围 + +- 单元测试: +- 集成测试: +- 手工验证: + +## 11. Definition of Done + +- 实现满足 Spec: +- 测试已运行或说明无法运行原因: +- 需求追踪表已更新: +- Project State 已更新: +- 接口、安全、权限、审计、集成契约已同步: +- 无 Secret、真实数据、构建产物或无关本地变更: + +## 12. 文档更新 + +完成后检查 Domain、Architecture、Workflow、ADR、Project State 是否需要更新。 diff --git a/docs/import/reusable/ai-native-templates/WORKFLOW.template.md b/docs/import/reusable/ai-native-templates/WORKFLOW.template.md new file mode 100644 index 0000000..99f26f3 --- /dev/null +++ b/docs/import/reusable/ai-native-templates/WORKFLOW.template.md @@ -0,0 +1,37 @@ +# 业务流程名称 + +## 1. 目标 + +说明这个流程要完成什么业务目标。 + +## 2. 触发条件 + +说明流程从哪里开始。 + +## 3. 输入 + +| 输入 | 中文说明 | 来源 | +| --- | --- | --- | +| | | | + +## 4. 输出 + +| 输出 | 中文说明 | 去向 | +| --- | --- | --- | +| | | | + +## 5. 正常流程 + +1. 步骤一。 +2. 步骤二。 +3. 步骤三。 + +## 6. 异常与补偿 + +- 异常: +- 补偿: +- 审计: + +## 7. 边界 + +说明哪些事情属于本流程,哪些事情不属于本流程。 diff --git a/docs/import/reusable/general-development-guidelines.md b/docs/import/reusable/general-development-guidelines.md new file mode 100644 index 0000000..0c48229 --- /dev/null +++ b/docs/import/reusable/general-development-guidelines.md @@ -0,0 +1,151 @@ +# 通用开发协作规范 + +## 1. 文档定位 + +本文记录可跨项目复用的开发习惯、协作方式、工程边界和 agent 工作规则。 + +新项目可以复制本文为根目录 `AGENTS.md`,再补充项目特有的技术栈、目录、启动命令和业务约束。 + +## 2. 工作方式 + +- 先确认目标、边界和验收标准,再写代码。 +- 大改动前先说明目标、范围和预计修改的文件。 +- 每次只做一个明确 checkpoint,不顺手扩展无关功能。 +- 不修改与当前任务无关的用户变更。 +- 不回滚用户自己的改动,除非用户明确要求。 +- 遇到不确定的技术栈、目录、接口契约或数据模型,先确认再继续。 +- 新项目初始化或技术栈升级前,必须检查前端、后端、构建工具、测试工具和运行时版本兼容性。 +- 输出分层结构、目录树、数据模型、字段映射或接口示例时,必须补充中文说明,不能只依赖英文命名表达业务含义。 + +## 3. 分支与提交 + +- 稳定分支保持可交付。 +- 集成分支用于合并多个功能。 +- `feature/*` 用于具体开发。 +- commit message 使用中文,清楚说明本次业务或技术变更。 +- 提交前检查工作区,避免误提交本地文件、Secret、构建产物或真实业务数据。 + +## 4. 项目结构 + +- 前端、后端、文档和接口契约应分目录管理。 +- 前端只负责展示、交互和调用本项目后端。 +- 后端负责数据库、Secret、外部系统适配、业务规则、审计和安全脱敏。 +- 文档目录应保存设计文档、导入规范、接口说明和重要决策记录。 +- 接口契约要有唯一来源,避免前端和后端各写一套互相漂移的定义。 +- 架构图、目录树、模块边界、表结构和接口字段说明应优先使用中文解释职责、依赖方向和业务含义。 + +## 5. 技术栈版本兼容性 + +新项目建立规范或引入依赖前,应先形成最小版本矩阵,确认前端、后端、运行时、构建工具、测试工具和接口契约工具之间没有明显冲突。 + +最低检查范围: + +- 后端语言版本、语法版本、框架版本、ORM / Mapper、数据库迁移工具、OpenAPI 工具和测试框架。 +- 前端 Node.js 版本、包管理器、框架、语言、构建工具、Lint、类型检查、测试框架和 UI 组件库。 +- 前后端接口契约工具,例如 OpenAPI 生成器、类型生成器、序列化格式和日期时间格式。 +- 本地开发、CI、部署环境的运行时版本是否一致。 + +落地要求: + +- 项目级文档必须写清楚当前采用的版本和兼容约束。 +- 版本结论应优先来自官方文档、项目依赖文件、锁文件或实际验证命令。 +- 发现版本冲突时,先调整版本方案,再生成代码或项目骨架。 +- 升级核心依赖前,应先检查 breaking changes、运行时要求和配套插件支持范围。 +- 无法确认兼容性时,应在项目文档中标记为待确认,不能假装已经验证。 + +## 6. 低耦合与可维护性 + +项目各功能模块必须保持低耦合、高内聚和高可维护性。新增功能时,应先明确模块边界、输入输出、依赖方向和验收标准,再开始实现。 + +落地要求: + +- 每个模块只负责一个清晰业务能力。 +- 依赖方向必须单向清晰,通用能力不能反向依赖具体业务。 +- 跨模块调用优先通过 Service、Port、API 契约或事件完成,不直接访问对方 Mapper、Entity 或内部实现。 +- 外部系统能力通过 Adapter 隔离,业务层只依赖稳定端口。 +- DTO、Entity、Domain、Response、外部系统 DTO 不混用。 +- 公共代码只有在出现真实重复和稳定语义后再抽取。 +- 禁止提前设计大而全的通用工具、通用服务或通用模型。 +- 单个文件、类、方法或组件过大时,应按业务职责拆分。 +- 修改一个功能时,不应迫使无关模块跟着改;频繁连锁修改说明边界需要重新设计。 +- 测试应覆盖模块公开行为和关键边界,避免只测试内部实现细节。 + +## 7. 前端规范 + +- 使用 TypeScript 时应开启严格类型检查。 +- 页面组件负责展示和交互,不直接承载复杂业务规则。 +- API 请求统一封装,不在组件里拼接后端 URL。 +- 服务端数据不要长期复制进客户端全局状态。 +- 前端不得保存或传递后端 Secret、Provider API Key、数据库凭证或客户渠道 Token。 +- 不使用中文或英文显示文案做业务判断。 + +## 8. 后端规范 + +- Controller 只处理 HTTP 契约、参数校验、权限入口和响应映射。 +- Service 编排业务流程、事务、领域对象、Repository 和外部端口。 +- 外部系统通过 Port / Adapter 隔离,业务层不直接依赖厂商 SDK 或外部 DTO。 +- 数据库变更必须可追踪;已发布 migration 不直接修改。 +- 新项目建表前必须明确数据库字符集和 collation 策略;MySQL 项目默认建议使用 `utf8mb4_bin`,避免外部 ID、Token、哈希、状态码或业务代码因大小写不敏感而误判。 +- 写接口要考虑幂等、并发版本、审计、失败恢复和脱敏。 + +Java 后端代码默认参考 Alibaba Java Coding Guidelines。可复用摘要见 `docs/import/reusable/alibaba-java-coding-guidelines-summary.md`。 + +后端基础结构、ID、审计字段、分页和 Mapper 规范可复用 `docs/import/reusable/backend-base-structure-pagination-guidelines.md`。 + +核心要求: + +- 命名清晰,不使用拼音、无意义缩写或随意缩写。 +- DTO、Entity、领域对象、VO / Response 不混用。 +- 不写魔法值,业务代码、状态码和固定字符串应抽为常量或枚举。 +- 集合、空值、字符串、时间、金额和 `BigDecimal` 使用安全写法。 +- 异常不吞掉,日志有上下文但不输出 Secret、Token 或个人敏感信息。 +- 分层结构、包结构、数据模型和字段映射必须配中文注释或中文说明,说明每层职责、依赖方向和关键字段含义。 +- 复杂流程、幂等、事务、外部字段映射和脱敏逻辑要写必要中文注释。 + +## 9. 接口契约 + +- API 返回稳定代码,不返回中文或英文文本作为前端业务判断依据。 +- 动态字段使用稳定 `fieldKey` 和可国际化 `labelKey`。 +- 错误响应不得回显 Secret、Token、原始消息正文、附件 URL 或个人敏感信息。 +- 接口变更前先确认领域模型、字段映射、前端影响和测试范围。 + +## 10. 安全规范 + +- Secret 只放环境变量、本地 `.env` 或部署平台 Secret。 +- 仓库只提交无真实值的 `.env.example`。 +- 禁止提交真实客户数据、Token、Cookie、API Key、数据库密码或支付信息。 +- 日志、错误响应和测试夹具不得暴露敏感信息。 +- `.DS_Store`、构建产物、依赖目录、IDE 临时文件和本地样本目录应忽略。 + +## 11. 测试与验证 + +- 修改后运行对应模块已有检查命令。 +- 如果检查命令尚未配置或因环境问题无法运行,必须明确说明原因。 +- 不假装测试通过。 +- 优先做小的可验证闭环,再逐步扩展业务能力。 +- 高风险改动需要补充更接近真实使用路径的测试。 + +## 12. Agent 协作规则 + +Coding agent 开始任务前应先读取: + +1. `AGENTS.md` +2. `README.md`,如果存在 +3. `docs/` 中与当前任务相关的文档 +4. 当前代码结构和最近 Git 状态 + +改文件前要说明计划;完成后要说明改了什么、如何验证、还有哪些风险或未完成项。 + +## 13. 新项目补充模板 + +复制到新项目后,建议补充: + +- 项目名称和业务定位。 +- 技术栈和版本。 +- 前后端与运行时版本兼容性检查结论。 +- 目录结构。 +- 分层结构、数据模型和字段映射的中文说明。 +- 本地启动、测试、构建命令。 +- 接口契约位置。 +- 不能随便改的目录或文件。 +- 当前第一个 checkpoint 和验收标准。 diff --git a/docs/import/字段组.xlsx b/docs/import/字段组.xlsx new file mode 100644 index 0000000000000000000000000000000000000000..565c70f2b7f28f1742b6b284b3f38bde9722b06f GIT binary patch literal 9179 zcmZ{K1zgkJ_df!HG>FnA6HuDbAl)@U7#$L$yCo%**aRgVDItuOMoL0rG}4Zc4k_v8 zH=gJB`aa*s|L?!Ow)@%L`@HWx=bn4dK08NC4Ga4L1_s7GjN|CXrl8iIWa2v*7z|h# z806?)V>wqBFNlkmx$a9hh^HyHpR-d^QTmsBC7bDyFUWXKc|@1^|$PdSl!6y+)(qo-1X4;&i$MzQK}I{N0*(%!I+X+AF!IgV@Q;Rcyq)-D%bz1Cw|;QtqtPPYOhSG&`;Ska-e3mMhxM4>)WQF62KAF%2O zLNlhZFk!uhu{N|M#kezXz*ObiSfbnEBD^S3^}t{#q-uC9)^-td=h zrSIEf+IdOen6vZ;rRm8~@cEOq6{CV)y#u9LrbWO<#_Q&LWVN-X!V4a$uU<&m=%w|E zAo%Z}@KGUR?mxhxOq2A&Y56tV&E_;wj76-sPjVb&nWThCyT@!_C@3D6yTZ&$_+?a) zrgWLjTg=e6nts#JkXvAu|L6Nc2Q~VbJITLZnOGAe&QtfKo>n*qM55L~0GyZi^KTFz zkKf+;A^dQi^0(Cc*ex=ap&8X9mxQM;;^EDS7Or?{JlTQJlk#N-No2n(@Oot3>eJ(Re+qdmt%YhKq*)C7xtVURP;nS#U$s!dv%xN=0k#xb(W6~AIthhJWaJ@ z&#_ED6{S(gVW-9PcVV$k$oR&IBIUhsB1Hs-l!)U^66cyz;>=z4r ziYrpyzt1ZiU|liCM4lE$-{a5G&|I`vNMOzN^9TCh2$NE)GIIj&$`SWjK6}(|u%tpp z!K>sn9q$wlb!Tm>NFF=-ti@l!orbU1fg8{1@aE$IiJ&#yA4eqD=qUo}qWy{|p7olg zgW&MPW~NE~HgjQS1j%flqi-QyGQJowbnad)0-uR>sH@}t{tjaU@9u8@`z6%gw7M*= zp|>MP^F+%ec3Jcj28xlvyY{=0pBgZlK4oBPjsQcCLPHDF>(6X<^OsW67Wf@mxuH}P zPt-2rET#<%e(J9_pBou$zN%xdgzQ-ZYxzT1Ic*4@o%x;VdmPc3VhiC&d;`GQ$iWh#|mt43^uG&|NouydqPtkR8fnP)#O6C`i+IMrMk zs-Nd>jS0xBU!N#fwn&0wHp6VmqAu(sow(c9U$Znfa&s{bcX9})VU02^y=s3%Yi`i@ z<`b2=amW>Hm3Vfxr19uILklinae2j}&0QJVj~VIpcSfHZFK|HVW?dD}qUzVB36XQ( z+ed=nWhGc4Hcq@xkaJA|!5HV=3xt$Eu2T}0aM{HBRcsOlNPgy7g(u<^(mc0)(oxFq zU)^5*rXK6pqaNDXXCB40dDv!CQa`)(G#`N^^^DUk$gpz5Ot-v?gFp2ywuy-bruTj$ zTetL>*-y1~sTt^2b|-!_S32_c$ZU$CrIw>V#y$#^HtHeW{S6TDpfFWIJDiIys zjL-mS>?6{|oUE&#YpSojA3tJSOj%BlK-hW>V5|3(dx>5}<9r5EnC_Qplu3Orb=eh~ zY&^qY{jNpwpogOe6#X9mIm(hUuF@^cO#)1%n0-onkg433Fl|;*x--64k~0BMSynNA(7m`oyrSc~L2^;CfzlU_ z1OgR)8t<<5n9pN|8+c=&{E9+uioISw!AlhTM&v16AHL>F`ZQSXLLBV6K*(xbC%w*4 z9?H=i!Yh#Kw0n>aEriLo&Ju)C>F8iMDpwfpeS97}*5>;VXY5y}s}qI)fE95@yZzd0 z%$Z+ean^wt#0y@91bMI7W<8DpyJ=_UB?s&$@SYN}8YsEqB zj?ra&rG6v1X!+gmUB+t)=jB~W&DwqiOE!K|)YJNuJ`>$%k{JBF6t(-NoH_RvHg4|x z`*dVf(0%5H4t3h-UkgP4dpdf(bb|bKEM_EVI7Xm@!i{QAGp+vm^O%OnyV;RD5aV@0 zhcW3E<;BY5Nngv7Vl5*?3bF})f)wf6@=L%pg?pP6)9w1=DXNb*EEM(qzA!%9i-=QL zy?iqD$+jPW(C6-MK3Rp6>Yu1TuN!^ybjrRaX-v*Rd=Unzc5i!pS0!Ot`tb@rC7j0k zX|r4nLt>;HC0^&=_g8)*FCMtF&3>kf+k2<~{l#-@Yt{#b#p4tYu-8w>O$#v~jJ*$& zI9J)&{UB%GqLOs(F>!yxRCRw4)33D)2eUl3GyD1$=ZCgYMHGDV8nMGkf3mUE%X?R( z?t5+MM!9C}^Y?fY{=5crN00oSm8TBP)C>gDeT=lp-p$8S?c(n_vVLO;Y>xBcFEDPT z@wu;9n#|%Q!~eV%7Q;&^CGWKp0r^6dueRW=bF~hOX81S9+JP8=D)-C z+jlaK9lV6EZ(>G*ylK{jWaAlxa9+3J|MKyr|HM~Q$|SxWd`*`!Oz-=B^k9^2!T(~Z zq@D>^t|j!S2qV-?$M&)?w7dh=!S94RR}GZBpP`_7q2-Px2eV><3}cBs)Dy_jbd{*k z;j0OZA3IT8;@CiD$5nF1C?V~@>S^sSx-f))gcw*@>zY!L=2i(O>U4EnRZ8^LW_H%K zIo4lzawNvChT(cUeSDRe5E4e+pu3RNfg@?#jPt(PH0Am(wmvJ~oHvA()5^%O*_JPb zj#F531xH@W@FKp6;rNF>g`GkQ%^^Z|J?H~{@Pjp~M zH+;7d{cQ#+epV=6Aw(hj#GcHz^-*!|hR%3N1->93-@zUHFHTQ?#M8z9$Q;TNyN?I- zWhM9fHRud1)R?$KaX%#` z8wZiX=2_RfYx8D-XG;f5Y341ahYKf^KB!8YO;$I`T+45k z7+24#3T}2%SBCUQ`zS+N{LdD*iw2Yiw6+U!Z3=xv9bz4Xl?HUSj}do^t_~0iMO;c; z+S@w_NYT}rHizy7{1VmnVNxF3b|?DYgf4TKFp2!kz{}r<-+!c4Z4XTEnkCvTNLESA z;leJXqkb;Gj(V{&kZ`(kQE+g$8AkA8tu#q8F(hW>_2V$wS!jwbM^G-wvZBIqzu=I` z0bxt?w0$PM0)eL6Bj#s^Z-4o%$`5^8UlbVnrdC8Q;fa-bwhDL64CShtd?p@45%MT1 zN{Xyy2~H_No(r9PM$S`ccGyqiN0Dlq{Sqg{!gCmnI(=;xbx(gTBGy*sP#wlCbD-Xo zhci;e`_s@VvvZaDa<57|kiS!5B6PX2VK&!dV{xbGt}tIB3t`x^(3(n#CV>~gN{i2UB6 z5HvFytt;klL9$1ELCGb>mWPi>sEksxA^Z<+Ry;=tC1QwEUlEWkMDFE#Nwzc7uq&qI z1Vt>e)I?<_-5Khy_H8U!rI`r{j(>oP+LUfaYz~0I8OB-or3oDb8aPtTC;JNCM7cqS zj6tjOcZ^$Z03Q9sH?X(bR{*#8VW}LO!y&Kh#K{87YwC?|uWS;0Js#w zo^OeHB!2{Ji0-XN9nB4(1;&++hu=wulz;o0D}qIL^$>n3$Y6Lx{tLX?BsOguDpsYm1ZL{8-a5r7D6WS z@qE7Fi#tOe`2nIPZ~GN&uwnfmMOL9v_o;S-le}(RvtAC)R&DAOzv*hOVEXVOB9#F= z+NplxQBOl%(tpnoPG73-kiS$Psl5<`X#2P!uh0T4v~5r~Ky~D7s0d0F=0o8-D{g|o zR5+Oszj!))M!h?fJa>^=8*338{GLrcZWso$bFBZkL6bMDopj{~g%`aS9e}+`omLr0 zT;(RNph_B#qM`C470hY;v{u&dXzV8&AZXnPV6W3>6$1b=s4c`A`7<(7n?-1>!1(|N zY1V9ZV04Dqu@QWumY4QJmOwys?^coym;;bSFC*Cq$a3z9k44H!^*<7ESSzM4*3@!Z0jEN zbOy}H<@yozbji7|F8RfERBy}InFCE$fN0Y@>lrC{P)Wa}Pejnxpk%GU&AT8c?JM!3 z=%5gJ$h`T+8Cz8xomq)dX*^Gn-WytROktNuT4BK(QkR*$Go`VJE6M_j5;`t}By6Jo z(u4&Y0)V(!jU~%`i5Ru{%)OR+dh?k#JUC=8eG~=X{UrGQ2>hNnBO7NGk0|lC0x=Kh z6J^0jpJS2bOCyj$M!%5l8|_`+TKBikglQ)b@b0*Ovd%BVvTh-eO&V8!zy5Srp%b8> z2*hpoC0oT?996G3EGv?p(JKVvvJ+#aLJ{cepwL(NkV_Qo4Lq+WJ9HO6 zuh$!CKTECB8ReWSdAPS*|?U2R;+1I?FRAQ zG_kXyHCnfGeF5$EIEfDyc?=YPIDv?<8qw{nu$$!26{Vi2<)tvsFhY1ZWUGR} zD|_Lz7NTFP(p@>bSsBj(DieX|merVx;jeN4Uf{Zki`x0+;k-o`jSp2O;DyMRH z=;$)xr6ZA505&WS^E?-OH$u<{=-YupzpKd}9^5KTwnc?}_7m zKM(J5RkwD-ncB*AV0whUoqqlh3YR>Vl!iwy?=99?e$SYEw{F$~IRajsk{$mRzc?M6 zfed&)pJidccoy9&IW*>}Z<8emTVDbn{$@Y#ffuEWbjN*+n;+@kUGHd_Eoo3w)rgWj zG?HhWmt*9~9ip&dERT^iQ(zPY9|6x#$qs*upP$kPO?4aV1|T_~+3?6P+PIPgHNhl} zKX`NhOBxnLRLK#nBobBdObIZdSC&ve*j1@s2gqT3bgJ zI=n+ZfxXk7LO!V$^>?c9(?UiSiS$LFUyl!0*YCn%wY)wWDndfZG-IOv9z`_PZuo$o z4}9G4OJh4HlZt=JkA?k|2V)lu7hBz>POs%XHc=6BDL3kl*BNK>1XC>(ef9*~d2`lc zOCU@X#==ab!OLN{Sa>lSOQ-xFtd%}hTiItT{adV&1Y$~Ito>Mli5*3it>29g^l`SJ z6JVLHZ+4(p1A)WUAr=w6?G{7Pdfk>BsBl$?JXeW5fZg$jzjujS<$*jYZ8;0Vu_~-V zNDb|fk&YtMRhHN*?Xe!&Np(~ECm(a@v_wTy!q8vGLAW~9B7$NG1g#ih z=FpW-qKNKI20VXg8Yf`{19FvF&vvc(pIwacj#SUA}*z@hA`7v=R-h%U2 z&~ussz$8pyKyEZ#*dRyjDX^&iNkNfzSEooA)V;fN%I*EQXif}N9&)gchZyE9U>w{y;atOXcvu=_S1b>@XoUUw zP+oD#zwehRe|Zb!RV15CMsVRZLN}X~M0Jw*CBBjI!Xmc`tqxKIU9?Y9dZXUR+swOe zO$SsKFCc5*IK1a6Qjk1L7R=lV5RrZ5&er-NwYkJut)bca4Zv|Gs+Xp2tV$oJRS-tM zh|oLBF*ukQSj{U^Z%pw1x+s#T1m_sL;)Rfx zZz$A*l!KdSAahVaQ;N@=8tpE%9tsrU{CV+={JmkCGLMb38t^yEO9$)cec>v1 zCI8YkkL;(K_}-^loY;qT2_siv$~G2nGLucNFg4&Fa)G6=DP>H-O~n+Ojjco_zfGqb zaUo93ikj)3Pq+3yK{AV$khXSC*|4H68X!z29q6kh+=Q<1^1|M&6nj2o^#?^ezU}G{ zCZ3EJ;mD00@(=PRdF-_arU7Ryt=ux8UGJ|&jUbuy?G1&F_)-*>4nm7Bg z-}AWZqH46~s^oVH$U-+hYX+m9FiwK|qu-Iuvu8)Zp&CLvRlX>q#*nmv8|(xD)`Q?)g>X?^w`;KtvOIz1DN)SOsG#hZKc*qao2?P4dP;Km{f-h+C=?|MZf(t1yof)kvFWT~^4_u|KE9QPFj z`=}C?c%ST2C7M=|+7?0zjFqh~3mad-G*o;f2q2{K=z&wpZwJl=G00>z<$_$;JQmdF zL<44JnU~Br?^#LnQ}Bp9}hZ!q-C2!reacHUc3aRz5!{h!tH9@hAQ{X2`9`B5}Bj5S70><+(G2X&$9UXY0OUXWL37V4LGKpe0bs`F5(R<+y@ zY@LaAHEpQA0*6^M`hIZZo#E#mPUWnQjIqu>;&4bBo~T9l3vTR~&&I!Xv~6W6f>^Dy z1*qKFe>lIl|J=j>2x>IL?%=LAsJo(00*QFg2Z~`43Yngy_YNuS50*gm2TQtjkZ(Cs z!XHlZQrJR!V-z6h-O2^zNG8+c6z51DBeZ#I{+%cvSSix8MbIUo`JwGLE^!V~BvC+B zp&t20)<&FwY6Bl5PLxv*UiXXDc<=DS{)Pcvo7vSFeuMl@LD=0NIpbe zf<#>OE)Z^fC+zxN|J*A)x_<&Q6fN+{b`NRseZE5k^%e0`Uw03EUVC`V98SF)?Zzr7 zu;CM@^iI`{lUZXM3L6~E%y8RPs(u+DcKAg?kcx1ql`ZJ_norrM2y7bnD8}sql zhx?R%gZdi6Y3m{lFJswbYAME4l%~()g`vi!%95<}Sp4i?~S$uPk$+J$AxHS_5YRyzM7*&tt z|6HJ$W;Yo)?p25MQoXhxm8(7@QQ2t-@}M}Tu?+|wP7Fhxznt=DXz2Wcl2fwOjO4E< zc#~P5eRV5RGoa7yCMvZR#BN%rIoM^kWXRvbJsTsH$LN&D1nzf6tD#lj&X06DC{NoI zp}5f4oeFBtaZGn6lGVZYDZYzL_#6LXlH$5N#JMr)el(z!sPk86-qYEeIlv zgKJThlWwvfS0khcNJCSxLQ@q#_AH^*Fvl91mQh6<89 z9cZd6x2?)juW6zNK@mzB!kwlLf^1_AUBBN`q(1i}+LQ#MoX;FZvJE|Dt8AidwnJ^4 z(?C$%Cba7Y5$W0<^k~~+?NQhzZ!cF#&<7AX(iMzy3b2F_>E%vm)u9c*B&9gJxFu+i zgVUYk&LfZ|`GrotlON007Vf3=_rdj9zzT zvj71AL23ZtH07%e)Ym5v;S=a!9_oic*_{jV_NvGB7!ee}I^pf>Twt`q0-Ff2c`L^A~@vh3w&WyOm z4*SgI^e>a^2WGQY$!ceCvPOE#saIH{6gaL{4gP55u;X!-@=9i~iKZ)^hw-|Hli=-RV~=^b&q*q8 zNWjh~6?~$Gs?5=NC9(5d)@k*(eqFEX%hLa02;`O*AabumH(T}`BM)p{msrF*@a|@s zZcagDRF-tl!=@ZhEbm^W+>Tn9lH+@q0w&&!2*hV&MpBBqj1tq$o%P*54(4dn5RssU zg~8PJDw^u&+~85nsg}y!CC3BWUxU%rs7val1R+WV0Ps?5h?o4|@t}PJkhgq&k;l8? zuh=Tm-z2svfa4s*d+*@jSGW{!v%PsTb|)gu2zLt?lQ%HbFx91PYVt9FBarsx@S7ML_EsQtELvsi%2jv zk*@Jp%t;1awnEnlELmBdQ<+fyGri37xs#|?E%!4P~Rkza=cljU?yEHbF1xxD$w zWug)NVT;zafJ+~I(c%ZnV|E2AS7YQtM0ptG3OfwpB0yoOC-ep;VR^mD`&`5L_yr?$ zYSQ-(T-kk?Mafdw6q>j=Jru37I04u3Sof-0TaCgPBSdl}lD5ffsZ+maY**cjLwNk| z;|%}0<%}$`q>VOvD}$K#(Xhy~83;{RT0Z0&c$m+X*sTPASKJMt0yCMp)?ke1EL?ZU zz@?q?-WAKG)zLo*FJH7J#|3lzc8UMvfL}G)1}Ek_s+3)Bi1Su9u<*HWfHXuodtC4F zO1dPEm*GV}D{Vq$t4yijF4&dNJZCHS?#|55x~Av)=`+?Zcl0wfw?1|{%bTIry$x74?hi&_=v{Sb@HO)Q2Xc%F z6?ikufpmMsNL@?V{mruNaS)_^!u$h?6*6$w{DF*3uLNn_B^b3u-Yr1RB1fhK2*LK* zsjqDK!_RAygZW-0IPD2Q&N^o|!1R8kT`C3-W1S+^%hiE2TDb~Rid@|IG_sa9lYi;B zRw$TE9}}Z*C|t4`Zs-(HvV!+JoO8#7T@DXaI-1*?TN%iJpBpsbt*ANh#NKz!9#V-J z-KaJ<7baFSh0sIktpN3mQndN#{%R0=b;kgdy93Wgc8K$yF?2PZq(}ObNi%!YbVjOT zB&C7xZ;k9(=AW95d!{rvw@ZJye_KoS*{v=BD;1krdtN6n&Z&quSTDO4AbcYWCex8A zRUD>8PxDeC3oJ=-!Ghsm*F6q5Nb}um z@)-57?>z*#Si%0mJAL_Pi%SYS@1H)Nu5!fpW+*sDM)5f5n=Ay#WL6@9m0?szpyVpp zW3{_WX7#5IR%rq)@zE@sNQOiI%>kKx%r1Fwy@4*LA>rFYnb~)DKVG%4bizKMp$DgA zEwffa8@1a_yQGMF2Qe`>I)pj$yNp109oQS)oFufHQebEFgA7*ck`=SRJB zUX6VE*OGsQRRx7o^xKY7_RRlN@`0gVh`(yRAj1$DLs7aTgRWLS&gM*=2|`pqbf?B< z;k@SuEe)}Ro>7OXglCo79*^uLHe4qqXL(4oc8r0}9RU?e{A(h~x#fZHEpXF2$nTp` zsiB?|QXd~X_X?Lt_*pONXUC z%mTA>p@>|nd>c;~A?M~F{CG9ar~Zx#EvsRH&5{_>rL#_^s8LUySY3>&c9_)bs$4Ml zk+|ml_xWGndwxql^TIm2|5FQ7Zd%C7eXftPUhDF5&5%>!vk!a;Ylpkjx^-*qA%ZU5 zHx%s(nN*dlAAf)RWs&&;>8APInn7IOq&$2j;LkSyErM8tzThi^L(s1|awSi;EKqWD zo02EizjO4@rj(L)1xZj~iakmkxD&*?a6ua@c!4(d4dZe!THvu_eT9(fD&~+sr(Xa) zNZKGtPT$_2ecC8Q2Yng2RSt->hq?VIi>xAy5tO~g$Od5=>;*ao`zHR>&}3(EL_hUq zPPj^;v9J0wm|}ex_QTG#S=q(zN2xV(iMrkH7B_HhjvWBVX!uPh7jruva8WHJ6V^dE zs8R2E{-C|#4YF|`C7XHga}-^W*Taa+jPR)6n$4%P2(%h*t+eT_b~%S28cPX=$smM; zoU^r6tD9mDzZ}GPp7yGy)jqa`^~gwOFI}=1)nm6R|0^kWTx-kJlucVn*|j`BNkO?I z5P>L4+oBl7<9?Nrv+Ogf1&(z9?S-#yO$q=_>A6)KBrNs;T)7d9JfDmPJJ;CR-P|r) zpMTYb*(UDhU$W=1Km4>tTV+A=StdvsICC+V6Kd?eS^J6>M*8+3e0KNZeS;(g zj%syV?R|5G>g->xJb|Jo@nyk?FZaD6W6pdj@lPt8OTN+>)(h%bF1BZ+V84HBUH@^X zw$rt`e0aX6a^LjT{LpqNsY%!3z=BjG+es+K@HJMNHGEm$oR&$U{n`vm8ByeusG4=E z`Qdo_kH=|t`^nG?TX_*)pOqOtjk__{#P+GUH1#pL#P-R#Ks0)UUV!@U#WYcarBSRx z*v~*M{ChMY{bDLWBJq{O#1c!nZZwmYnI@ zNDw*zV?-CSr6cs(^XT`4r8GH|4cyEMl}^>v=G>x-Rts0w|nr<`mYv+QoW>EUb;` znq=J*_KqXx=?$_jLZUNZ_ov~Ve0@N9@KAd_JmY>x`*1TWc)M$Z}$M; zePJkuvi`Sq=oY5bSkpxFw2px2$h$T6^GYD_#3gF5zM5q}E5i9KGg|*y3507MnvO|( z0xC`1016S~8Rwy~KTPmvWai?(9`w60uW0V|yJ*xM(O*0Rn4xzUYEqQrx;#Is=PFiS zBZT}y37k9YA`mb`1Gvjc2bQhWZx7P*C&~H{rxJ4CNInV*p^*CSanQL{e5Z z+?uFD=$wy|dysMb(|CR@wK<8+LB&VH5d55kSktO_!xtZ=A7{O9>33Hg(T8*a-$H$+ z)7<2SvuzNu`~M4#b+DM0rUM%{F}GTJ=>=Rj@X_x+y@u8wJlf?wqNoXw>fB!g-M=4u zDAhx`gaCDs$M#>2mNYp=$0VjbchwPs~Cd%0T|Bu`!{hS|APKgoKJ#J%GOVC3&VfG|3}13u7hZg QSJF}LHVVaCOvi8k0U#)GQ2+n{ literal 0 HcmV?d00001 diff --git a/docs/project/README.md b/docs/project/README.md new file mode 100644 index 0000000..91d55dd --- /dev/null +++ b/docs/project/README.md @@ -0,0 +1,50 @@ +# fire-safety-ymd 项目文档索引 + +本目录保存 fire-safety-ymd 专属知识。通用 AI-NSES 原文和模板保存在 `../import/reusable/`,不要在通用文件中写入本项目业务规则。 + +## 阅读入口 + +| 文档 | 用途 | 更新频率 | +| --- | --- | --- | +| [`../../AGENTS.md`](../../AGENTS.md) | Agent 工作方式、Go 规则和安全底线 | 低 | +| [`../../CONTEXT.md`](../../CONTEXT.md) | 项目目标、系统边界和技术栈 | 低 | +| [`../../PROJECT_STATE.md`](../../PROJECT_STATE.md) | 当前 checkpoint、风险和下一步 | 高 | +| [`ai-native-adoption.md`](ai-native-adoption.md) | AI-NSES 在本仓库的目录映射和采用状态 | 低 | +| [`ai-nses-project-overlay.md`](ai-nses-project-overlay.md) | fire-safety-ymd 对通用标准的项目覆盖规则 | 低 | +| [`backend-development-guidelines.md`](backend-development-guidelines.md) | Go 后端工程规范 | 低 | +| [`security-access-control-boundary.md`](security-access-control-boundary.md) | AI、MCP、数据库、身份和敏感数据边界 | 中 | +| [`integrations/superagent-openapi.md`](integrations/superagent-openapi.md) | SuperAgent Open API、Session、SSE 恢复与探针说明 | 中 | +| [`integrations/superagent-mcp-spatial.md`](integrations/superagent-mcp-spatial.md) | SuperAgent 空间只读 MCP、配置、工具和联调门禁 | 高 | +| [`operations/postgis-srid-4326.md`](operations/postgis-srid-4326.md) | 已确认 WGS84 数据的 SRID 元数据迁移、验证与权限边界 | 高 | +| [`operations/nginx-public-entry.md`](operations/nginx-public-entry.md) | `agent.nianxx.com` 到 Go 对话/MCP 的 HTTPS 反向代理示例 | 高 | +| [`operations/docker-test-deployment.md`](operations/docker-test-deployment.md) | `/home/firee-safety-ymd` 测试服务器的 Docker Compose、Nginx、验证与回滚手册 | 高 | +| [`../architecture/chat-api-v1.md`](../architecture/chat-api-v1.md) | 用户对话入口、内存 Session 映射、SSE 和三凭证边界 | 高 | +| [`../architecture/public-chat-entry-v1.md`](../architecture/public-chat-entry-v1.md) | DashScope 风格兼容入口、Nginx、会话映射与严格结果边界 | 高 | +| [`../architecture/spatial-mcp-v1.md`](../architecture/spatial-mcp-v1.md) | MCP、Service、Repository 与 8 张 PostGIS 表的架构 | 中 | +| [`../workflows/user-chat.md`](../workflows/user-chat.md) | 首轮/多轮对话、失败处理和 curl 联调流程 | 高 | +| [`../workflows/exercise-plan-evidence.md`](../workflows/exercise-plan-evidence.md) | 从地名候选或演练点坐标查询方案证据的流程与能力边界 | 中 | +| [`../specs/fire-safety-ymd-chat-api-v1.md`](../specs/fire-safety-ymd-chat-api-v1.md) | 默认关闭的用户对话 API v1 可验收契约 | 高 | +| [`../specs/fire-safety-ymd-dashscope-compatible-chat-v1.md`](../specs/fire-safety-ymd-dashscope-compatible-chat-v1.md) | 既有 `completion` 客户端的受限兼容契约 | 高 | +| [`../specs/fire-safety-ymd-superagent-openapi-connectivity.md`](../specs/fire-safety-ymd-superagent-openapi-connectivity.md) | SuperAgent Open API 连通性基线 Spec | 中 | +| [`../specs/fire-safety-ymd-superagent-mcp-spatial-readonly-v1.md`](../specs/fire-safety-ymd-superagent-mcp-spatial-readonly-v1.md) | 空间只读 MCP v1 的可验收契约 | 高 | +| [`../import/db-samples/README.md`](../import/db-samples/README.md) | 本地 SQL 样例的敏感数据与禁止执行规则 | 中 | + +## 通用 AI-NSES 资料 + +| 文档 | 用途 | +| --- | --- | +| [`../import/reusable/ai-native-software-engineering-standard.md`](../import/reusable/ai-native-software-engineering-standard.md) | 通用 AI-NSES 标准 | +| [`../import/reusable/general-development-guidelines.md`](../import/reusable/general-development-guidelines.md) | 跨项目开发协作规范 | +| [`../import/reusable/ai-native-templates/README.md`](../import/reusable/ai-native-templates/README.md) | ADR、Spec、Workflow、Handoff 等模板索引 | + +## 后续文档目录 + +以下目录已在出现真实内容时建立;仍不创建空占位文件: + +- `docs/domain/`:消防资源、火情位置、防火网格等稳定业务概念。 +- `docs/architecture/`:已有用户对话、DashScope 风格公网兼容入口和空间 MCP v1;后续保存系统、数据、模块和部署架构。 +- `docs/workflows/`:已有用户对话和演练方案证据查询流程;后续保存异常补偿流程。 +- `docs/adr/`:不可逆或跨模块的重要决策。 +- `docs/specs/`:可验收的功能规格;已有原生/兼容用户对话 API、SuperAgent Open API 和空间只读 MCP Spec。 + +新增文档必须能回答“未来的新 Agent 为什么需要阅读它”,并从本索引或上级文档建立入口。 diff --git a/docs/project/ai-native-adoption.md b/docs/project/ai-native-adoption.md new file mode 100644 index 0000000..312f6f9 --- /dev/null +++ b/docs/project/ai-native-adoption.md @@ -0,0 +1,85 @@ +# AI-NSES 采用说明 + +| 项 | 内容 | +| --- | --- | +| 项目 | `fire-safety-ymd` | +| 采用日期 | 2026-09-04 | +| 通用标准版本 | AI-NSES 0.2 | +| 当前采用阶段 | Active development | + +## 1. 目标 + +本项目采用 AI-NSES,将长期项目知识保存在仓库中,使人类开发者和新的 Agent 不依赖历史会话也能理解当前目标、代码边界、风险和下一步。 + +通用标准保持原样,项目差异放在根入口和 `docs/project/ai-nses-project-overlay.md`,避免 fire-safety-ymd 的规则反向污染可复用模板。 + +## 2. 本次接入的通用文件 + +模板来源快照:`/Users/andy/IdeaProjects/th-hotel-simple/docs/import/reusable`。 + +仅接入以下白名单内容: + +- `docs/import/reusable/ai-native-software-engineering-standard.md` +- `docs/import/reusable/general-development-guidelines.md` +- `docs/import/reusable/ai-native-templates/` 全目录 + +它们作为通用参考和新文档模板保存,不直接承担本项目当前状态说明。 + +## 3. 明确排除的来源内容 + +没有复制以下内容: + +- TH Hotel 的 `docs/project/`。 +- TH Hotel 的根 `AGENTS.md`、`CONTEXT.md`、`PROJECT_STATE.md`。 +- reusable 目录中的 Java、Spring、MyBatis、前端、分页基础结构和特定集成指南。 +- reusable 根 `README.md`,因为本 checkpoint 的复制白名单未包含它。 + +本仓库的项目级文档均按森林防火、Go、SuperAgent、MCP 和 PostGIS 的实际已知边界重新编写。 + +## 4. AI-NSES 目录映射 + +| AI-NSES 职责 | fire-safety-ymd 位置 | 采用状态 | +| --- | --- | --- | +| Agent 工作入口 | `AGENTS.md` | 已采用 | +| 长期项目背景 | `CONTEXT.md` | 已采用 | +| 动态项目状态 | `PROJECT_STATE.md` | 已采用 | +| 项目文档索引 | `docs/project/README.md` | 已采用 | +| 项目覆盖规则 | `docs/project/ai-nses-project-overlay.md` | 已采用 | +| 长期工程规范 | `docs/project/backend-development-guidelines.md` | 已采用 | +| 安全与权限边界 | `docs/project/security-access-control-boundary.md` | 已采用;Chat 静态联调门禁与 MCP 服务账号边界已实现,最终用户机制待设计 | +| Domain 文档 | `docs/domain/` | 按需创建 | +| Architecture 文档 | `docs/architecture/` | 已采用;已有原生/兼容用户对话入口和空间只读 MCP v1 | +| Workflow 文档 | `docs/workflows/` | 已采用;已有用户对话和演练方案证据查询流程 | +| ADR | `docs/adr/` | 首个架构决策时创建 | +| Spec | `docs/specs/` | 已采用;已有原生/兼容用户对话 API、SuperAgent Open API 与空间只读 MCP Spec | + +Go 源码使用 `cmd/`、`internal/` 和 `pkg/`,这是通用标准中 `src/` 的项目级映射。 + +## 5. 项目 Definition of Ready + +跨模块、数据、权限、安全或外部平台的 Feature 进入实现前必须明确: + +- 目标、非目标、调用方和可测试验收标准。 +- SuperAgent 或 MCP 的真实契约来源,不依据记忆猜测字段。 +- 身份、区域或租户范围如何进入可信服务端上下文。 +- PostgreSQL 表、字段、几何类型、SRID、数据时效和权限是否已经核验。 +- 错误、超时、重试、幂等、审计和敏感信息处理方式。 +- 所需测试层级和可用测试环境。 + +## 6. 项目 Definition of Done + +Feature 完成时至少满足: + +- 代码已 `gofmt`,相关测试和 `go test ./...` 已运行,或明确记录阻塞原因。 +- 接口只返回被授权、可追踪且不泄露敏感信息的数据。 +- 数据和空间计算结论有实际验证依据,没有把待确认项写成事实。 +- 相关 Spec、ADR、Workflow、安全文档和 `PROJECT_STATE.md` 已同步。 +- 没有引入 Secret、真实敏感样本、构建产物或无关用户改动。 +- 交付说明包含变更清单、验证结果、剩余风险和下一 checkpoint。 + +## 7. 维护方式 + +- 通用 AI-NSES 升级时先比较上游快照,再单独做文档升级 checkpoint。 +- 项目规则只修改项目入口或 Overlay,不直接编辑通用快照。 +- `PROJECT_STATE.md` 每个 checkpoint 更新;稳定背景变化时才更新 `CONTEXT.md`。 +- 采用状态或目录映射变化时更新本文。 diff --git a/docs/project/ai-nses-project-overlay.md b/docs/project/ai-nses-project-overlay.md new file mode 100644 index 0000000..d6667cf --- /dev/null +++ b/docs/project/ai-nses-project-overlay.md @@ -0,0 +1,74 @@ +# fire-safety-ymd AI-NSES 项目覆盖规则 + +## 1. 适用范围 + +本文补充通用 AI-NSES,只适用于 `fire-safety-ymd`。它记录森林防火、Go 后端、SuperAgent、MCP 和 PostGIS 带来的项目专属约束,不修改通用标准本身。 + +## 2. 规则优先级 + +从高到低: + +1. 用户在当前 checkpoint 中的明确要求。 +2. 根 `AGENTS.md`。 +3. 本项目 Overlay 与专项安全/接口契约。 +4. 通用 AI-NSES 和通用开发规范。 +5. 仅供参考的模板和历史讨论。 + +发现冲突时停止扩大实现,在交付中明确冲突和所采用的上级规则。 + +## 3. 核心系统边界 + +目标链路为: + +```text +用户侧应用 + -> fire-safety-ymd Go 服务 + -> SuperAgent + -> 本项目 MCP 工具 + -> PostgreSQL/PostGIS +``` + +- Go 服务拥有身份上下文、权限校验、外部调用适配、数据访问和审计责任。 +- SuperAgent 拥有自然语言理解、工具选择和答案组织责任,不拥有授权决策和数据库权限。 +- MCP 层提供有限的领域能力,不把历史表名、任意 SQL 或内部凭证暴露给模型。 +- PostgreSQL/PostGIS 提供空间事实,但记录存在不等于资源当前可用。 + +## 4. 森林防火事实规则 + +- 数据库事实、用户陈述、模型推断和建议必须可区分。 +- 查询结果应逐步具备来源、更新时间、状态、不确定性和是否需要现场确认等字段。 +- 未验证 PostGIS SRID 和几何类型前,不发布距离、范围包含或最近资源结论。 +- 防火通道只有几何和长度时,只能表达空间邻近,不能推断车辆可通行或计算可靠路线。 +- 水源或设施没有可用状态时,不能只因记录存在就宣称可用。 +- 联系人和电话属于受控数据,是否返回取决于服务端角色和区域权限。 + +## 5. AI 与应急安全规则 + +- 大模型不得虚构人员、设施、危险源、道路状态、距离或实时火情。 +- 高风险回答应明确数据时间和未知项,必要时提示现场确认。 +- AI 输出是辅助信息,不替代 119、人员撤离、消防控制室和现场指挥体系。 +- 工具失败、数据过期或权限不足时返回明确状态,不允许模型用猜测补齐。 +- 面向真实火情的分阶段方案、路径计算或资源调度属于后续高风险 Feature,必须先有单独 Spec、安全审查和测试方案。 + +## 6. Go 交付约束 + +- 保持 `cmd -> internal/app -> handler/service -> domain/repository` 的清晰依赖,不机械模仿 Java 分层。 +- 默认从 Go 标准库和小接口开始;新增框架、SDK 或通用抽象需要现实用例。 +- 所有外部系统通过适配边界接入,领域和服务层不得依赖厂商 DTO。 +- 每个 checkpoint 至少执行 `gofmt` 和 `go test ./...`;并发、数据库或网络代码按风险增加 race、集成和失败场景测试。 + +## 7. 文档门禁 + +以下变化必须同步文档: + +| 变化 | 必须更新 | +| --- | --- | +| SuperAgent 或 MCP 契约 | Spec、Architecture/Workflow、Project State | +| PostgreSQL schema、SRID 或查询语义 | 数据架构、相关 Domain、测试说明 | +| 身份、角色、区域/租户边界 | 安全文档、接口契约、测试场景 | +| 高风险应急建议逻辑 | Spec、Domain、Workflow、安全边界、ADR(如适用) | +| Go 核心依赖或 module path | 后端规范、ADR(如影响较大)、Project State | + +## 8. 当前例外与待确认 + +当前已实现无业务数据的 `/health`、默认关闭的原生用户对话 API、可选 DashScope 风格兼容入口、SuperAgent Open API 出站 Adapter/探针,以及默认关闭的空间只读 MCP/PostGIS 基线。对话入口只有独立静态联调凭证、精确 Origin、单进程内存会话和并发 Run 冲突,不是最终用户认证;兼容接口仍坚持严格成功后才返回正文。Docker/Compose 测试部署和 Nginx 公网配置目前均是待目标机验证的安全基线:应用只发布宿主机回环端口,Secret 在运行时注入且迁移凭证不得进入服务容器。MCP 已具备独立服务间 Bearer、显式服务端静态数据范围(数据库全范围或镇街白名单)、固定参数化查询和 readiness 门禁;4,048 条实库非空几何已通过受控事务补齐 SRID 4326 元数据,严格 readiness 和 7 个工具的本地真实查询均已通过,但真实消防 Profile、公网部署与 SuperAgent MCP 联调尚未完成。无效几何在不改写原始数据的前提下被查询排除并明确告警;动态授权、共享/持久会话、持久审计和生产安全控制仍未实现。 diff --git a/docs/project/backend-development-guidelines.md b/docs/project/backend-development-guidelines.md new file mode 100644 index 0000000..713c49f --- /dev/null +++ b/docs/project/backend-development-guidelines.md @@ -0,0 +1,141 @@ +# fire-safety-ymd Go 后端开发规范 + +## 1. 当前技术基线 + +| 项 | 当前选择 | 说明 | +| --- | --- | --- | +| Go module | `fire-safety-ymd` | 临时名称,正式仓库路径待确认 | +| Go 工具链 | 本地 `1.26.6`;Docker builder `1.26.8` | 本地测试已验证;镜像版本已显式固定,目标服务器构建待验证 | +| HTTP | 标准库 `net/http` | 当前无需 Web 框架 | +| 测试 | `testing`、`httptest` | 当前无第三方测试库 | +| 数据库 | PostgreSQL + PostGIS + `pgx/v5 v5.10.0` | 原生 pgxpool、固定只读 SQL;实库 SRID 严格 readiness 已通过 | +| AI 集成 | SuperAgent Open API Adapter + MCP `2025-06-18` | 协议代码已实现并默认关闭;真实环境待联调 | + +引入依赖前优先核对官方版本、Go 版本要求、维护状态、许可证和可替代性。核心依赖升级需要独立验证,不能只修改版本号。 + +## 2. 目录职责和依赖方向 + +```text +cmd/server 进程入口、信号处理、依赖装配 +internal/app 应用生命周期与顶层 wiring +internal/config 配置读取、默认值和校验 +internal/handler HTTP/MCP 入站契约和响应映射 +internal/service 用例编排、授权后的业务流程 +internal/domain 稳定领域概念、值对象和业务规则 +internal/repository PostgreSQL/PostGIS 等持久化适配 +internal/integration SuperAgent 等外部系统协议适配 +pkg 真正需要对外复用的稳定包 +``` + +依赖原则: + +- `cmd/server` 只装配,不承载业务逻辑。 +- Handler 不直接执行 SQL,不直接依赖 SuperAgent 厂商 DTO。 +- Service 通过小接口使用持久化或外部系统;接口优先定义在使用方附近。 +- Domain 不依赖 HTTP、数据库驱动、MCP SDK 或 SuperAgent SDK。 +- Repository 负责参数化查询、行映射和数据库错误转换,不组织自然语言回答。 +- 避免为了分层而创建一一映射的空结构;只有出现真实职责时才增加文件和抽象。 + +## 3. 包与 API 设计 + +- 包名短、小写、表达单一能力,不使用 `utils`、`common`、`base` 等无边界集合。 +- 导出标识符必须有稳定的外部使用场景;默认保持在 `internal` 或不导出。 +- 接口应小,由调用方定义;不要先为每个结构体创建接口。 +- 构造函数显式传入必需依赖,依赖缺失应尽早失败。 +- 避免包级可变状态、隐式 `init` 副作用和 service locator。 +- `context.Context` 作为有取消、超时或请求范围操作的第一个参数,不保存在长期结构体字段中。 +- 时间使用 `time.Time`/`time.Duration`;跨系统时间格式、时区和精度必须在契约中明确。 + +## 4. 错误处理 + +- 正常业务失败返回错误,不使用 `panic`。 +- 使用 `fmt.Errorf("动作和对象: %w", err)` 添加上下文并保留错误链。 +- 用 `errors.Is`/`errors.As` 判断语义,不比较错误文本。 +- Domain/Service 错误映射为稳定 API code;客户端不依赖中英文 message 做逻辑判断。 +- HTTP 响应和日志不回显 SQL、凭证、Token、完整个人信息或第三方原始敏感响应。 +- 不静默吞掉影响正确性的错误;确实只能忽略时要有注释或受控观测方式。 + +## 5. HTTP 规范 + +- 路由显式声明方法;不支持的方法返回 405。 +- JSON 响应设置 `Content-Type: application/json; charset=utf-8`。 +- Handler 只做协议解析、格式校验、可信身份上下文读取、调用 Service 和响应映射。 +- 请求体设置合理大小上限;服务端设置 Header、请求、空闲和优雅关闭超时。 +- 后续统一错误格式、请求 ID 和日志关联字段,并在 Spec/OpenAPI 中形成唯一契约来源。 +- `/health` 当前是 liveness,只说明进程可处理 HTTP;依赖就绪检查应使用独立 readiness 端点,不能改变现有语义而不更新契约。 + +## 6. 配置与 Secret + +- 环境变量名称使用项目明确前缀,例如 `FIRE_SAFETY_`。 +- 配置由 `internal/config` 一次读取、校验并作为不可变值注入。 +- Secret 不设置可误用的生产默认值,不写入仓库、命令参数、普通日志或错误响应。 +- 本地 `.env` 被忽略;需要示例时只提交无真实值且有注释的 `.env.example`。 +- 启动时验证必需配置,错误指出配置键但不输出其值。 + +## 7. PostgreSQL/PostGIS 规范 + +数据库接入遵守: + +- 当前使用 `pgx/v5` 原生 `pgxpool`;升级前核对 Go/PostgreSQL 支持范围、安全变更和回归测试。 +- 所有值使用参数绑定,表名和排序字段使用代码白名单;绝不拼接模型生成 SQL。 +- 使用最小权限数据库角色,MCP 第一阶段默认为只读。 +- 查询设置请求超时和数据库 statement timeout;分页、数量上限和空间半径必须有限制。 +- 事务边界由 Service 用例决定,Repository 不隐藏跨用例事务。 +- 已发布 migration 只追加,不原地篡改;导入原始 SQL 前单独审计 `DROP`、sequence、owner 和 extension 依赖。 +- `geometry` 的 SRID、实际类型和有效性先通过真实数据核验,再决定字段约束与查询。 +- 经纬度距离计算明确选择合适的 `geography`、投影转换或测地算法,并以米等业务单位返回。 +- 用 `ST_DWithin` 等可利用 GiST 索引的条件限制候选集,避免对整表逐行计算距离。 +- WKB、GeoJSON、经纬度和 MultiLineString/Polygon 映射必须有代表性测试。 + +## 8. SuperAgent 与 MCP 适配 + +- 外部 SDK 和 DTO 只存在于适配层;Service 使用项目自己的端口和领域对象。 +- 当前 SuperAgent Adapter 分离 `CreateSession` 与 `StreamMessage`;上层负责持久化本地会话映射,不能每轮隐式创建新 Session。 +- SSE 成功必须同时具备最终内容、成功 `run.completed` 和顶层 `end`;断流只恢复既有 Run,不重新 POST 消息。 +- DashScope 风格兼容 Handler 只做入站协议转换,与原生 `/api/chat` 共享 Chat Service;兼容 `session_id` 必须映射本地会话,不能直接暴露或接受 Provider Session。 +- 兼容 SSE v1 不透传 `message.delta`;只有严格成功后才以 `finish_reason=stop` 返回正文。若要真正逐字输出,必须先更新 Spec、安全失败语义和客户端验收条件。 +- 当前 MCP 只实现 `2025-06-18` 的 initialize、initialized notification、tools/list、tools/call 和同步 JSON 响应;扩展 SSE、session 或新协议版本前需更新 Spec。 +- 重试只用于可安全重试的操作,并设置次数、退避和总时限;避免重复创建会话或重复业务动作。 +- MCP 工具按领域能力命名,输入输出 schema 稳定、有限且有中文业务说明。 +- `user_id`、角色、租户/区域等授权上下文由服务端注入,不成为模型可自由选择的普通参数。 +- 工具结果保留结构化状态、来源和不确定性,不让 Agent 解析日志或数据库原始字段。 + +## 9. 测试与验证 + +基础命令: + +```bash +gofmt -w ./cmd ./internal +go test ./... +``` + +按风险增加: + +```bash +go vet ./... +go test -race ./... +``` + +测试要求: + +- 纯逻辑优先使用表驱动单元测试。 +- HTTP 使用 `httptest` 验证方法、状态码、Header、JSON 和错误边界。 +- Repository 集成测试使用隔离数据库和可重复迁移,不依赖个人长期数据库。 +- PostGIS 测试覆盖 SRID、点/线/面、边界点、空几何、无效几何、半径上限和排序稳定性。 +- 外部适配测试覆盖超时、取消、断流、限流、鉴权失败和不完整事件。 +- 测试夹具不得包含真实联系人、Token、生产坐标或客户数据。 + +## 10. 日志与可观测性 + +- 使用结构化日志字段表达 request、tool、duration、result 和 error category,避免拼接大段原始内容。 +- 不记录 Authorization、Cookie、数据库 DSN、完整用户问题中的敏感信息或完整 MCP 结果。 +- 后续为 HTTP、SuperAgent、MCP 和数据库调用建立统一 request/correlation ID。 +- 指标区分业务无结果、权限拒绝、上游失败、超时和内部错误,不能只看 HTTP 500 总数。 + +## 11. 提交前检查 + +- 仅包含本 checkpoint 文件,没有覆盖用户已有变更。 +- `gofmt` 和相关测试已运行。 +- 新依赖有理由且 `go.mod`/`go.sum` 一致。 +- 没有 Secret、真实敏感数据、构建产物或 IDE 文件。 +- Project State、接口、安全和相关设计文档已同步。 diff --git a/docs/project/integrations/superagent-mcp-client.example.json b/docs/project/integrations/superagent-mcp-client.example.json new file mode 100644 index 0000000..4451479 --- /dev/null +++ b/docs/project/integrations/superagent-mcp-client.example.json @@ -0,0 +1,33 @@ +{ + "format": "fire-safety-ymd-superagent-mcp-client/v1", + "mcpServers": { + "fire-safety-ymd-spatial-readonly": { + "transport": "http", + "protocolVersion": "2025-06-18", + "url": "https:///mcp", + "method": "POST", + "headers": { + "Authorization": "Bearer ${FIRE_SAFETY_MCP_AUTH_TOKEN}", + "Content-Type": "application/json", + "Accept": "application/json" + }, + "tools": { + "allow": [ + "fire_safety_search_place_candidates", + "fire_safety_resolve_incident_context", + "fire_safety_find_nearby_water_sources", + "fire_safety_find_command_post_candidates", + "fire_safety_list_nearby_access_lines", + "fire_safety_get_responsible_units", + "fire_safety_find_nearby_risk_areas" + ], + "writeTools": [] + }, + "runtimeNotes": { + "authTokenSource": "Use a separate per-environment high-entropy secret; never reuse the SuperAgent Open API Key.", + "scopeSource": "The all or town-allowlist data scope is configured on the fire-safety-ymd server and is never supplied by tool arguments.", + "safety": "Results are planning evidence only. Invalid source geometries are excluded, so results may be incomplete. Availability, passability, command-post suitability, live team positions and assembly sites require field confirmation." + } + } + } +} diff --git a/docs/project/integrations/superagent-mcp-spatial.md b/docs/project/integrations/superagent-mcp-spatial.md new file mode 100644 index 0000000..4097be0 --- /dev/null +++ b/docs/project/integrations/superagent-mcp-spatial.md @@ -0,0 +1,160 @@ +# SuperAgent 空间只读 MCP 接入指南 + +| 项 | 内容 | +| --- | --- | +| 状态 | 代码已实现;实库严格 readiness 与本地 7 工具冒烟已通过,公网/SuperAgent 联调待执行 | +| Endpoint | `POST /mcp` | +| MCP 版本 | `2025-06-18` | +| 传输 | 单 JSON 请求/响应;不提供服务端 SSE | +| 数据源 | PostgreSQL/PostGIS,固定只读查询 | + +## 调用边界 + +```mermaid +flowchart LR + SA["SuperAgent"] -->|"独立 MCP Bearer Token"| MCP["fire-safety-ymd /mcp"] + MCP -->|"schema 校验、超时、可信数据范围"| SVC["SpatialService"] + SVC -->|"固定参数化 SQL"| PG["PostgreSQL/PostGIS 只读角色"] + PG --> SVC --> MCP --> SA +``` + +SuperAgent Open API Key 用于本服务调用 SuperAgent;MCP Token 用于 SuperAgent 回调本服务。两者方向、权限和生命周期不同,必须使用不同 Secret。 + +## 首版工具 + +| 工具 | 数据表 | 能回答 | 不能回答 | +| --- | --- | --- | --- | +| `fire_safety_search_place_candidates` | 现有 8 张空间表 | 按名称、镇街或村庄搜索有界地点候选 | 任意地址地理编码、自动确定演练点 | +| `fire_safety_resolve_incident_context` | 防火网格 | 点位所在网格、镇街、区域标签 | 实时火情、负责人电话 | +| `fire_safety_find_nearby_water_sources` | 水源地、蓄水池 | 候选水源、距离、容量/源状态(有值时) | 当前可用、取水道路可达 | +| `fire_safety_find_command_post_candidates` | 检查站、瞭望哨 | 附近候选设施 | 自动确定指挥部、安全性/容量结论 | +| `fire_safety_list_nearby_access_lines` | 防火通道 | 附近已绘制通道、最近接入点 | 路线规划、车辆可通行、实时封路 | +| `fire_safety_get_responsible_units` | 防火网格 | 责任中队名称 | 队伍实时位置、战备状态、集结点 | +| `fire_safety_find_nearby_risk_areas` | 墓地坟区、林区工矿企业 | 周边风险区域与距离 | 实时危险程度 | + +工具不会返回负责人、书记、队长、值班人员、电话或图片字段。 + +## 地名输入的最小流程 + +`fire_safety_search_place_candidates` 接受 `place_name` 和可选 `limit`,名称长度为 2 至 100 字符,默认返回 10 条、最多 20 条。它对现有业务记录的名称、镇街和村庄字段做不区分大小写的包含匹配,不是完整地名库,也不调用外部地图服务。 + +返回值包含: + +- `matched_field`:命中了名称、镇街还是村庄。 +- `match_kind`:`exact` 或 `partial`。 +- `location_kind=recorded_point`:源记录本身是点。 +- `location_kind=representative_point`:源记录是线或面,只返回只读计算的代表点。 + +SuperAgent Profile 应遵循: + +1. 用户只给地名时,先调用该工具,不直接猜测坐标。 +2. 无结果时,请用户补充“镇街 + 村庄 + 具体地标”或在地图上选点。 +3. 多个候选时按名称、类型、镇街和村庄列出,让用户选择;不能静默使用第一条。 +4. 即使只有一个候选,也先回显给用户确认,再把确认后的 WGS84 坐标交给其他空间工具。 +5. `representative_point` 不能描述为地点中心、入口或真实演练点。 +6. 工具返回 `source_records_with_invalid_geometries_are_excluded` 时,必须说明结果可能不完整;无结果只能表示在有效记录中没有找到。 + +可加入 SuperAgent 系统提示词的最小片段: + +```text +当用户没有经纬度但提供了地名时,先调用 fire_safety_search_place_candidates。 +搜索结果只是地点候选:无结果时请用户补充地名或地图选点;多结果时列出候选并请用户选择;不得静默选择第一条。 +任何候选都要先回显名称、类型、镇街、村庄和坐标供用户确认。location_kind=representative_point 时必须说明它只是线面记录的代表点,不能当作真实演练点。 +只有用户确认坐标后,才调用网格、水源、指挥部候选、防火通道、责任中队和风险区域工具。 +所有空间工具都会排除无效几何。看到 source_records_with_invalid_geometries_are_excluded 时,要说明结果可能不完整;不得把无结果解释为原始数据库确认不存在。 +``` + +## 配置顺序 + +1. 为数据库创建只授予所需表 `SELECT` 的独立角色。 +2. 在本地未提交的 `.env` 中填写 `.env.example` 新增的 `FIRE_SAFETY_POSTGIS_*` 配置,先保持 `FIRE_SAFETY_MCP_ENABLED=false`。 +3. 加载环境并执行只读 readiness probe: + +```bash +set -a +source .env +set +a +go run ./cmd/postgis-probe +``` + +probe 只输出 PostGIS 版本、表行数、空/无效/越界几何计数、几何类型、SRID 和 GiST 索引状态,不输出业务记录、联系人、DSN 或 SQL。无效/空几何作为明确 warning 并从工具查询中排除,不阻塞 MCP;非 4326 SRID、异常类型或越界坐标仍阻塞启动。 + +4. 只有数据所有者确认所有几何确为 WGS84,且数据库通过单独、经审查的数据变更把非空几何元数据补充为 `SRID=4326` 后,才设置: + +```text +FIRE_SAFETY_POSTGIS_EXPECTED_SRID=4326 +``` + +如果 probe 显示 `SRID=0`,不要让应用在查询时静默 `ST_SetSRID`。本项目的数据提供方已确认源数据为 EPSG:4326 且无偏移;2026-09-05 已按 [`../operations/postgis-srid-4326.md`](../operations/postgis-srid-4326.md) 完成受控元数据迁移,迁移未改变坐标数值或丢弃 Z。重新导入原始 SQL 后仍需重新执行该迁移,不得假设导入文件自带 SRID。 + +5. 配置至少 32 个可打印 ASCII 字符的独立 MCP Token,并选择一种可信数据范围,最后启用 MCP。 + +若该 MCP 凭证获准读取当前数据库中 MCP 固定查询表的全部记录,使用显式全范围模式: + +```text +FIRE_SAFETY_MCP_AUTH_TOKEN=<独立高熵 Secret> +FIRE_SAFETY_MCP_SCOPE_MODE=all +FIRE_SAFETY_MCP_ALLOWED_TOWNS= +FIRE_SAFETY_MCP_ENABLED=true +``` + +若只获准读取部分镇街,保留默认白名单模式: + +```text +FIRE_SAFETY_MCP_AUTH_TOKEN=<独立高熵 Secret> +FIRE_SAFETY_MCP_SCOPE_MODE=town_allowlist +FIRE_SAFETY_MCP_ALLOWED_TOWNS=莒格庄镇,高陵镇 +FIRE_SAFETY_MCP_ENABLED=true +``` + +`all` 表示当前配置数据库中 MCP 固定查询表的所有记录,包括镇街字段为空的记录;它不表示最终用户已获得动态授权。`all` 模式不能同时保留镇街列表,模型和 MCP 工具参数均不能改变服务端选择的范围。 + +6. 启动服务。启用 MCP 时,服务会在监听 HTTP 前执行 readiness:SRID、类型或坐标范围不符合要求时拒绝启动;无效/空几何会输出 warning 并由查询排除。 + +## 协议探测 + +真实 Token 只放环境变量,不写命令历史或文档: + +```bash +curl -sS http://127.0.0.1:8080/mcp \ + -H "Authorization: Bearer ${FIRE_SAFETY_MCP_AUTH_TOKEN}" \ + -H 'Content-Type: application/json' \ + -H 'Accept: application/json' \ + -d '{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"manual-probe","version":"1"}}}' +``` + +随后发送 `notifications/initialized`、`tools/list` 和受控测试地名/坐标的 `tools/call`。联调不得使用真实火情或未授权精确坐标。 + +SuperAgent 侧配置模板见 [`superagent-mcp-client.example.json`](superagent-mcp-client.example.json)。 + +## 运行与安全说明 + +- `/mcp` 默认不注册;`GET /mcp` 返回 405,不提供 SSE stream。 +- 浏览器携带 `Origin` 的请求被拒绝,首版只支持服务到服务调用。 +- 数据范围由服务端配置注入;默认镇街白名单,数据库全范围必须显式选择,工具参数不能提供或扩大权限范围。 +- 数据库连接设置只读事务默认值和 statement/lock timeout;工具另有总超时、半径与数量限制。 +- 普通日志只记录 request ID、方法/工具名、耗时和结果类别,不记录坐标、参数或结果正文。 +- 所有工具结果包含 `source_records_with_invalid_geometries_are_excluded`;SuperAgent 必须把无结果表述为“有效记录中未找到”,不能据此断言原始数据库不存在相关记录。 +- HTTP 服务本身不终止 TLS;部署时必须通过受控网关或反向代理提供 HTTPS、网络白名单、限流和 Secret 轮换。 + +## 当前联调门禁 + +- 样例 SQL 不可执行,也不可提交;其中包含破坏性 DDL 和受限联系人数据。 +- 2026-09-05 已完成受控 SRID 元数据迁移和严格 audit:8 表共 4,055 条记录,4,048 条非空几何均为 SRID 4326,SRID/类型/范围硬门禁通过。7 条空几何、35 条无效面几何以及 7 张缺 GiST 索引表仍按预期告警;真实 `tools/call` 尚未执行,因此当前仍不能宣称 MCP 的业务结果已经联调验证。 +- 缺少适用于距离表达式的 GiST 索引时可做小数据开发联调,但生产前必须补齐并验证查询计划。 +- 地名包含匹配当前没有专用名称索引;真实数据量下先验证查询耗时,后续再决定标准地名表、别名词典或 `pg_trgm` 索引。 +- 用户身份、动态区域授权和持久审计仍是后续 checkpoint;当前一个 MCP Token 只对应一个静态数据库全范围或镇街白名单。 + +## 2026-09-05 本地实库冒烟结果 + +服务仅监听 `127.0.0.1:18080`,使用运行时只读 PostGIS 账号和显式 `all` 数据范围完成测试;未输出联系人、完整业务记录或测试坐标。 + +- `GET /health` 返回 200;无 Token 的 `POST /mcp` 返回 401。 +- `initialize` 协商 MCP `2025-06-18`;initialized notification 返回 202;`tools/list` 返回全部 7 个工具。 +- 使用“观水镇”得到 10 个地点候选,并选取一个精确匹配的 `recorded_point` 蓄水池记录,仅作为获授权测试坐标。 +- 网格上下文返回 1 条;水源、指挥部候选和通道各返回 10 条;责任中队返回 1 条;风险区域返回 9 条。 +- 7 个响应的 `structuredContent` 与文本 JSON 投影一致,空间参考均为 EPSG:4326,计数与数据数组一致,不包含已禁止的联系人类字段键。 +- 所有工具均保留 `source_records_with_invalid_geometries_are_excluded` 以及各自能力限制 warning。 +- 第二轮 7 个数据库工具调用分别约耗时 0.35 至 0.99 秒,均低于 5 秒工具超时;这只是当前小数据本地结果,不替代生产索引和并发验证。 +- 不存在的地名稳定返回 `no_results`;越界经度返回 `INVALID_ARGUMENT`;不可信 Origin 返回 403;`GET /mcp` 返回 405。 +- 运行日志只记录 request ID、工具名、结果类别和耗时,未记录坐标、请求参数或结果正文。 diff --git a/docs/project/integrations/superagent-openapi.md b/docs/project/integrations/superagent-openapi.md new file mode 100644 index 0000000..953b263 --- /dev/null +++ b/docs/project/integrations/superagent-openapi.md @@ -0,0 +1,125 @@ +# SuperAgent Open API 项目接入指南 + +| 项 | 内容 | +| --- | --- | +| 项目 | `fire-safety-ymd` | +| 协议基线 | TH Hotel 仓库保存的 2026-07-12 Open Agent API 资料 | +| 当前状态 | Go Adapter、默认关闭的原生用户对话 API 与可选兼容入口已完成模拟联调;上线前必须与当前 SuperAgent 环境重新联调 | + +## 1. 项目调用边界 + +```text +用户侧应用 + -> fire-safety-ymd POST /api/chat + 或 /api/v1/apps/{app_id}/completion + -> 内存会话与并发 Run 控制 + -> SuperAgent Client Port + -> SuperAgent Open API Adapter(本 checkpoint) + -> 已发布的消防 SuperAgent Profile +``` + +前端不得持有 Open API Key 或直接调用 SuperAgent。Adapter 只负责 Provider 协议,消防业务 Service 不直接解析 SSE 或依赖 Provider DTO。 + +## 2. Profile 与 Session + +Open API 请求不直接提交 `profile_id`。平台通过 `df_open_*` 外部应用 Key 对应的策略选择已发布 Profile,因此消防项目必须创建独立外部应用并绑定消防 Profile,不能复用 TH Hotel 的 Profile 或 Key。 + +`external_subject_id` 表示外部主体,不是 Profile ID。首版对话 API 使用服务端固定的非真实测试主体,并在单进程内存中保存: + +```text +本地用户 + 本地会话 + -> SuperAgent external_subject_id + -> SuperAgent session_id +``` + +同一多轮对话复用同一 Session;同一 Session 有 active Run 时返回 `409`,不能并发重发。该映射尚未持久化,重启和多实例切换后旧对话不可恢复;真实用户阶段必须改用经验证且假名化的主体,并补共享/持久会话存储。 + +## 3. Open API 流程 + +1. 创建 Session:`POST /api/open/agent-sessions`。 +2. 流式发送消息:`POST /api/open/agent-sessions/{session_id}/messages/stream?include_trace=true`。 +3. 初始流断开时查询 `GET {Content-Location}`。 +4. 使用 `GET {Content-Location}/events` 和 `Last-Event-ID` 恢复。 +5. 只有最终内容、成功 `run.completed` 和顶层 `end` 同时存在时返回成功。 + +如果未来需要主动取消,平台资料中的接口为: + +```text +POST /api/open/agent-sessions/{session_id}/runs/{run_id}/cancel +``` + +取消能力不属于本 checkpoint。 + +## 4. 外部应用要求 + +平台管理员需要为本项目确认: + +- 独立消防 Profile 已发布且 API exposure 已开启。 +- 外部应用 Key 已绑定目标 Profile。 +- 至少具备 `agent_sessions:create`、`agent_sessions:message` 和 `agent_sessions:read` scope。 +- 后续需要取消 Run 时增加 `agent_sessions:cancel`。 +- 需要公开 Trace 时,应用策略的 `trace_policy.enabled=true`。 +- 工具输入、输出和步骤信息默认只暴露 summary,不开放模型原始思考过程。 + +## 5. 本地配置 + +复制 `.env.example` 中的占位配置到本地 Secret 管理方式,不提交真实值。 + +主要变量: + +```text +FIRE_SAFETY_SUPERAGENT_ENABLED=true +FIRE_SAFETY_SUPERAGENT_BASE_URL=https:// +FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY= + +FIRE_SAFETY_CHAT_ENABLED=true +FIRE_SAFETY_CHAT_AUTH_TOKEN= +FIRE_SAFETY_CHAT_ALLOWED_ORIGINS=http://localhost:5173 +``` + +fire-safety-ymd 不读取通用 `DEERFLOW_OPEN_API_KEY`,避免开发机上其他项目的凭证被意外复用。 + +## 6. CLI 连通性探针 + +配置测试环境后执行: + +```bash +go run ./cmd/superagent-probe +``` + +探针只发送代码内固定的无敏感信息消息,输出最终回答和必要的安全元数据。它不启动业务聊天、不查询消防数据库、不写业务状态。 + +探针不会自动读取 TH Hotel 的 `DEERFLOW_*` 或其他项目变量。至少需要显式设置本项目的启用开关、Base URL 和 Open API Key;未配置时命令会在任何网络请求前退出。 + +禁止把真实火情、联系人、电话、精确受限位置、生产凭证或客户数据放进探针消息和 metadata。 + +## 7. 用户对话 API + +原生接口为 `POST /api/chat`,请求体只接受 `message` 和可选的本地 `conversation_id`,响应是 SSE。首次请求由服务端创建 SuperAgent Session 并返回随机对话 ID,后续轮次使用该 ID 复用上下文。客户端不能提交 Provider Session、主体、角色、区域或 metadata。 + +配置 `FIRE_SAFETY_CHAT_COMPAT_APP_ID` 后,同一 Chat Service 还提供 DashScope 风格受限兼容路径。它把 `input.prompt` / `input.session_id` 映射到上述本地会话语义,以 `event: result` 和最终 `finish_reason=stop` 返回严格完成的正文。该入口不是 DashScope 全量代理,也不会把 URL App ID 或客户端 session ID直接传给 Provider。 + +详细事件、失败语义和 curl 示例见: + +- [`../../architecture/chat-api-v1.md`](../../architecture/chat-api-v1.md) +- [`../../workflows/user-chat.md`](../../workflows/user-chat.md) +- [`../../specs/fire-safety-ymd-chat-api-v1.md`](../../specs/fire-safety-ymd-chat-api-v1.md) + +`FIRE_SAFETY_CHAT_AUTH_TOKEN` 是独立的首版联调凭证:原生接口将其作为 Bearer,兼容入口将其作为 `xtoken`。它不是最终用户登录;浏览器会暴露静态 Token,因此公网真实用户入口仍必须接入身份提供方、动态授权、限流和审计。 + +## 8. 与 MCP 的关系 + +Open API 与 MCP 是两条独立连接: + +- fire-safety-ymd -> SuperAgent:使用 Open API Key。 +- SuperAgent -> fire-safety-ymd `/mcp`:使用独立 MCP 凭证。 + +本文件记录第一条 Open API 接入。仓库现已另外实现默认关闭的只读空间 MCP 基线;两条系统连接仍使用不同凭证,不能把 Open API Key 当作 MCP Token。用户侧原生/兼容对话使用第三个独立联调凭证,三个值都不得复用。MCP 的部署与工具契约见 `superagent-mcp-spatial.md`;兼容入站契约见 [`../../specs/fire-safety-ymd-dashscope-compatible-chat-v1.md`](../../specs/fire-safety-ymd-dashscope-compatible-chat-v1.md)。 + +## 9. 当前限制 + +- 用户聊天 HTTP API 和单进程并发 Run 控制已实现,但默认关闭,只有静态测试 Bearer,没有最终用户认证或动态授权。 +- 本地会话未持久化;重启、多实例切换和上游失败后不能恢复旧 `conversation_id`,也尚无主动取消。 +- MCP/PostGIS 已完成本地实库冒烟,但尚未完成 SuperAgent 到公网 MCP 的联调。 +- 真实 Profile、Key、scope、Trace 策略和网络连通性必须在目标环境验证。 +- Provider 协议可能在 2026-07-12 资料后变化;出现差异时更新 Spec 和契约,不在 Adapter 中静默猜测。 diff --git a/docs/project/operations/docker-test-deployment.md b/docs/project/operations/docker-test-deployment.md new file mode 100644 index 0000000..aa47723 --- /dev/null +++ b/docs/project/operations/docker-test-deployment.md @@ -0,0 +1,379 @@ +# 测试服务器 Docker 部署手册 + +| 项 | 约定 | +| --- | --- | +| 目标服务器目录 | /home/firee-safety-ymd | +| 运行方式 | 宿主机 Nginx 终止 HTTPS,Docker Compose 运行单个 Go 服务 | +| 应用监听边界 | 只把容器端口发布到宿主机 127.0.0.1:8080 | +| 数据库 | 使用现有 PostgreSQL/PostGIS;本手册不创建数据库容器、不导入 SQL | +| 状态 | 测试环境部署基线;目标服务器、公网 DNS/TLS 和 SuperAgent 回调仍需现场验证 | + +本手册不包含真实 Token、API Key、数据库密码或旧 Nginx 配置内容。命令中的尖括号是服务器上需要替换的占位符;不要把 Secret 写进命令行、Nginx 文件、镜像构建参数或 Git。 + +## 1. 部署前提 + +本仓库在本手册编写时尚无 Git 提交。服务器不能直接 clone 当前工作区,必须由项目维护者先审阅变更、手动 commit 并 push 到远程仓库;本项目不会自动提交、推送或迁移服务器数据。以下命令以发布 `main` 分支为例。 + +在开发机提交前先检查暂存区;当前已有两份用户 Excel 处于暂存状态,`git commit` 会把它们一并提交,除非维护者明确取消暂存。不要把 `.env` 或被忽略的原始 SQL 强制加入 Git: + +~~~bash +cd /Users/andy/IdeaProjects/fire-safety-ymd +git status --short +git diff --cached --stat +# 审阅后通过 IDE 或逐项执行 git add,只暂存确定要发布的路径。 +git commit -m 'bootstrap fire-safety test deployment' +git push -u origin main +~~~ + +服务器需要具备: + +- Git、Docker Engine 和 Docker Compose v2(命令为 docker compose)。 +- 宿主机 Nginx、可由 Nginx 读取的 agent.nianxx.com 证书和私钥。 +- 到 PostgreSQL/PostGIS 的网络访问;数据库账号应为运行时只读账号。 +- 如由 SuperAgent 访问 MCP,防火墙、安全组和 TLS 必须允许按约定访问 443。 + +先确认工具版本,不要把配置文件内容打印到日志: + +~~~bash +git --version +docker --version +docker compose version +nginx -v +~~~ + +如果 Nginx 尚未安装,可按发行版选择一种方式;用户已安装时跳过: + +~~~bash +# Debian/Ubuntu +sudo apt-get update +sudo apt-get install -y nginx +~~~ + +~~~bash +# RHEL/CentOS/Fedora 系 +sudo dnf install -y nginx +~~~ + +安装后确认服务: + +~~~bash +sudo systemctl enable --now nginx +sudo systemctl status nginx --no-pager +~~~ + +## 2. 获取代码到 /home/firee-safety-ymd + +### 首次部署 + +维护者完成 commit/push 后,在服务器上执行。私有仓库使用服务器已有的 Git 凭证管理方式,不要把 Git 密码或 Token 写入命令: + +~~~bash +sudo mkdir -p /home/firee-safety-ymd +sudo chown "$(id -un):$(id -gn)" /home/firee-safety-ymd +git clone https://git.nianxx.cn/huangting/fire-safety-ymd.git /home/firee-safety-ymd +cd /home/firee-safety-ymd +git switch main +~~~ + +如果目录已经是本仓库的工作副本,先确认远程、分支和工作树,再只做快进更新: + +~~~bash +cd /home/firee-safety-ymd +git remote -v +git status --short +git fetch --prune origin +git switch main +git pull --ff-only origin main +~~~ + +若 git status --short 有服务器本地改动,先确认归属;不要用 git reset --hard 或覆盖本地文件排障。 + +## 3. 创建服务器 .env + +.env.example 只是字段说明,不能作为运行时 Secret 文件。首次部署时复制并收紧权限: + +~~~bash +cd /home/firee-safety-ymd +cp .env.example .env +chmod 600 .env +~~~ + +用服务器上的受保护编辑器填写 .env,至少确认以下值。三个凭证必须分别生成,不能相互复用: + +~~~text +FIRE_SAFETY_HTTP_ADDR=:8080 + +FIRE_SAFETY_SUPERAGENT_ENABLED=true +FIRE_SAFETY_SUPERAGENT_BASE_URL=https:// +FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY= + +FIRE_SAFETY_CHAT_ENABLED=true +FIRE_SAFETY_CHAT_AUTH_TOKEN=<独立的高熵xtoken> +FIRE_SAFETY_CHAT_SUBJECT_ID=fire-safety-ymd-chat-test-subject +FIRE_SAFETY_CHAT_COMPAT_APP_ID=<与Nginx路径完全一致的公开app-id> +# 仅浏览器实际 Origin;CLI/服务端调用可以留空 +FIRE_SAFETY_CHAT_ALLOWED_ORIGINS=https:// + +FIRE_SAFETY_MCP_ENABLED=true +FIRE_SAFETY_MCP_AUTH_TOKEN=<独立的高熵MCP-bearer> +FIRE_SAFETY_MCP_SCOPE_MODE=all +FIRE_SAFETY_MCP_ALLOWED_TOWNS= + +FIRE_SAFETY_POSTGIS_ENABLED=true +FIRE_SAFETY_POSTGIS_DSN=<只读PostgreSQL连接串> +FIRE_SAFETY_POSTGIS_MIGRATION_DSN= +FIRE_SAFETY_POSTGIS_EXPECTED_SRID=4326 +~~~ + +注意: + +- FIRE_SAFETY_CHAT_AUTH_TOKEN 只用于兼容对话入口的 xtoken;FIRE_SAFETY_MCP_AUTH_TOKEN 只用于 SuperAgent 调用 /mcp;FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY 只用于 Go 服务访问 SuperAgent。三者必须不同。 +- FIRE_SAFETY_CHAT_COMPAT_APP_ID 是公开路径标识,不是 Secret;它必须与 Nginx 的精确 location = /api/v1/apps//completion 完全一致。 +- 使用 FIRE_SAFETY_MCP_SCOPE_MODE=all 时,FIRE_SAFETY_MCP_ALLOWED_TOWNS 必须保持为空。all 是服务账号级的固定查询表范围,不是最终用户级授权。 +- FIRE_SAFETY_POSTGIS_MIGRATION_DSN 不应配置给运行服务;迁移凭证只供一次性迁移命令使用,完成后应移除。 +- FIRE_SAFETY_HTTP_ADDR 在容器内应为 :8080,安全边界由 Compose 的 127.0.0.1:8080:8080 和宿主机 Nginx 提供;不要在 Compose 场景设为容器内的 127.0.0.1:8080。 +- 不要用没有 --quiet 的 docker compose config 或 docker inspect 把完整环境渲染到终端、CI 日志或工单;检查 .env 权限仍为 600。 + +## 4. 容器访问现有 PostgreSQL/PostGIS + +本项目不在 Compose 中创建或初始化数据库。数据库继续由现有 PostgreSQL/PostGIS 运维,Go 服务只使用已完成 SRID 确认的只读数据源。 + +如果 PostgreSQL 在另一台机器: + +- DSN 使用数据库私网 DNS 或私网地址,不使用容器内的 127.0.0.1。 +- 数据库防火墙只允许测试服务器访问 5432;如启用 TLS,按数据库证书要求配置 DSN 的 sslmode,不要为了省事关闭证书校验。 + +如果 PostgreSQL 与 Docker 在同一台宿主机: + +- 容器内的 127.0.0.1 指向容器自身;DSN 使用 Compose 声明的 host.docker.internal(或运维明确提供的 Docker 网关地址)。 +- Linux Compose 需要有 host.docker.internal:host-gateway 映射;本仓库的 Compose 配置包含该映射时才能使用这个名称。 +- PostgreSQL 的 listen_addresses、pg_hba.conf 和主机防火墙只允许 Docker 网段中的只读账号访问,不要发布 5432 到公网。 +- 不要把迁移账号放入 .env;数据库配置变更先由管理员验证最小权限。 + +服务启动时进行 MCP PostGIS readiness 校验。SRID、几何类型或 WGS84 范围硬门禁失败会阻止服务装配;空/无效几何和缺 GiST 索引是当前已知 warning。GET /health 只是进程存活检查,不代表数据库、SuperAgent 或 MCP 已可用,必须同时检查 Compose 日志和实际接口。 + +## 5. 构建、启动和本机检查 + +以下命令在服务器项目目录执行。Compose 构建上下文不把 .env 作为 Dockerfile 构建参数: + +~~~bash +cd /home/firee-safety-ymd +docker compose config --quiet +docker compose build --pull +docker compose up -d +docker compose ps +docker compose logs --tail=100 api +curl --fail http://127.0.0.1:8080/health +~~~ + +预期结果:容器为 Up(健康状态由 Compose 显示),health 返回 HTTP 200;有限日志不应出现 DSN、密码、Token 或 API Key。docker compose config --quiet 只验证配置,不输出展开后的 Secret。若 Compose 不支持 --quiet,升级 Compose 或使用不会回显结果的等价校验方式。 + +失败时只读取有限日志: + +~~~bash +docker compose ps +docker compose logs --tail=200 api +~~~ + +不要把完整日志粘贴到公共工单;先删去 URL 凭证、连接串、Header、联系人和精确位置。readiness 报错时核对数据库地址、只读权限和 FIRE_SAFETY_POSTGIS_EXPECTED_SRID=4326,不要关闭门禁或修改原始几何。 + +## 6. 安装/替换宿主机 Nginx 配置 + +仓库模板为 deploy/nginx/fire-safety-ymd.conf.example,不含真实 Secret。它把精确 Chat 兼容路径、/mcp 和可选 /health 反代到 127.0.0.1:8080,其他路径返回 404;详细 SSE 边界见 nginx-public-entry.md。 + +先只查找拥有 agent.nianxx.com 的配置文件名,避免使用 nginx -T 把旧文件中的 DashScope Key 打到终端: + +~~~bash +sudo grep -RIl \ + 'server_name[[:space:]]\+agent\.nianxx\.com' \ + /etc/nginx/conf.d /etc/nginx/sites-enabled 2>/dev/null +~~~ + +旧配置迁移前: + +1. 记录旧配置文件路径和归属;不要把旧配置内容复制到仓库或新模板。 +2. 如需回滚,将旧文件以 root 权限备份到不被 Nginx include 的目录并限制为 600;若含真实 DashScope Key,应由凭证所有者安排轮换。 +3. 把旧文件移出 *.conf include(例如改为 .disabled),避免同一域名出现两个冲突的 server 块。移动前确认目标是上一步查到的精确路径,不要批量操作整个 /etc/nginx。 + +部署模板的示例(目录因发行版而异;使用 sites-enabled 时按实际 include 位置调整): + +~~~bash +sudo install -o root -g root -m 0644 \ + deploy/nginx/fire-safety-ymd.conf.example \ + /etc/nginx/conf.d/fire-safety-ymd.conf +sudo vi /etc/nginx/conf.d/fire-safety-ymd.conf +~~~ + +编辑时完成以下非 Secret 替换: + +- 将 replace-with-fire-safety-app-id 替换为 .env 中的 FIRE_SAFETY_CHAT_COMPAT_APP_ID。 +- 确认证书路径、域名和上游 127.0.0.1:8080 符合测试服务器。 +- 保留精确 location、SSE 的 proxy_buffering off/超时设置和其他路径的 404;不要恢复旧配置的 location /、DashScope proxy_pass、Authorization Key 或 Nginx if Token 比较。 +- limit_req_zone 必须位于 Nginx http 上下文;模板假定由 conf.d/*.conf 在 http {} 中 include。若 include 层级不同,把指令移到合法的 http 上下文。 + +确认旧配置已禁用且新文件不含 Secret 后,再验证并 reload: + +~~~bash +sudo nginx -t +sudo systemctl reload nginx +~~~ + +如果 Nginx 不由 systemd 管理,使用发行版的 reload 方式;reload 前必须先通过 nginx -t。 + +## 7. 公网 Chat SSE 冒烟 + +先验证宿主机本地 /health 和 HTTPS /health,再用不含敏感信息的问题验证 SSE。将 Token 安全注入当前 shell(例如受保护的 Secret 管理器或交互式读取),不要把字面量写入 shell 历史: + +~~~bash +read -rsp 'Chat xtoken: ' CHAT_XTOKEN +printf '\n' +~~~ + +首轮请求: + +~~~bash +curl -N --fail \ + -H "xtoken: $CHAT_XTOKEN" \ + -H 'Content-Type: application/json' \ + -H 'Accept: text/event-stream' \ + --data '{"input":{"prompt":"观水镇附近有哪些水源候选?"},"parameters":{}}' \ + 'https://agent.nianxx.com/api/v1/apps//completion' +~~~ + +确认能读取 event: result,最后一个成功事件的 output.finish_reason 为 stop。只有 stop 事件的 output.text 可当作完成;中途断流、HTTP 错误或只有 finish_reason=null 不得当作可信答案。output.session_id 是本项目本地会话 ID,不是 SuperAgent 内部 Session ID。 + +后续轮次使用上一轮成功返回的本地会话 ID: + +~~~bash +curl -N --fail \ + -H "xtoken: $CHAT_XTOKEN" \ + -H 'Content-Type: application/json' \ + -H 'Accept: text/event-stream' \ + --data '{"input":{"prompt":"再说明候选的限制","session_id":""},"parameters":{}}' \ + 'https://agent.nianxx.com/api/v1/apps//completion' +~~~ + +至少验证:错误 xtoken 为 401;不存在的 app ID 和未列出的路径为 404。不要使用真实联系人、电话或精确敏感位置作为冒烟问题。 + +完成 Chat 测试后清理当前 shell 中的临时变量: + +~~~bash +unset CHAT_XTOKEN +~~~ + +## 8. SuperAgent MCP 回调配置与冒烟 + +在 SuperAgent 的工具/MCP 配置中新增远程 MCP 服务: + +~~~text +URL: https://agent.nianxx.com/mcp +Method: POST +Authorization: Bearer +~~~ + +该 Bearer 必须是 .env 中独立的 FIRE_SAFETY_MCP_AUTH_TOKEN,不要填 Chat xtoken 或 SuperAgent Open API Key。SuperAgent 还可能需要公网来源 IP 白名单、TLS CA 或自定义 Header,按实际平台和网络策略确认。 + +在服务器或受控终端用独立变量测试初始化: + +~~~bash +read -rsp 'MCP bearer: ' MCP_BEARER +printf '\n' +curl --fail \ + -H "Authorization: Bearer $MCP_BEARER" \ + -H 'Content-Type: application/json' \ + --data '{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"manual-probe","version":"1"}}}' \ + 'https://agent.nianxx.com/mcp' + +curl --fail --output /dev/null --write-out 'initialized HTTP %{http_code}\n' \ + -H "Authorization: Bearer $MCP_BEARER" \ + -H 'Content-Type: application/json' \ + --data '{"jsonrpc":"2.0","method":"notifications/initialized"}' \ + 'https://agent.nianxx.com/mcp' + +curl --fail \ + -H "Authorization: Bearer $MCP_BEARER" \ + -H 'Content-Type: application/json' \ + --data '{"jsonrpc":"2.0","id":2,"method":"tools/list"}' \ + 'https://agent.nianxx.com/mcp' +~~~ + +确认第一步返回 `protocolVersion: 2025-06-18`,第二步返回 HTTP 202,第三步列出 7 个固定工具;再用无敏感信息的已知地名/点位做少量只读调用,核对正确 Bearer 成功、错误/缺失 Bearer 为 401、warnings 保留、非法参数不会变成任意 SQL。查询到资源记录不代表实时可用、路线已规划或正式集结点。 + +完成 MCP 测试后清理当前 shell 中的临时变量: + +~~~bash +unset MCP_BEARER +~~~ + +## 9. 更新、重启与回滚 + +### 发布新版本 + +维护者先将目标版本 commit/push 后,服务器快进更新并重建: + +~~~bash +cd /home/firee-safety-ymd +git fetch --prune origin +git switch main +git pull --ff-only origin main +docker compose config --quiet +docker compose build --pull +docker compose up -d --remove-orphans +docker compose ps +docker compose logs --tail=100 api +~~~ + +更新会重启单个 Go 进程,现有内存会话丢失;客户端按会话失效处理并重新开始对话。当前部署不支持多实例共享会话,也没有持久化历史或审计。 + +### 回滚 + +回滚到上一个已验证的 Git commit 或镜像版本。先保留现状信息并确认目标版本;不要用破坏性 reset 覆盖未知的服务器本地改动: + +~~~bash +cd /home/firee-safety-ymd +git status --short +git log --oneline -5 +~~~ + +由维护者确认目标 commit(例如 ``)后,先以 detached HEAD 检出该已验证版本,再重新构建启动;这是显式回滚,不会自动覆盖服务器本地改动: + +~~~bash +git fetch --prune origin +git switch --detach +docker compose build --pull +docker compose up -d --remove-orphans +docker compose ps +~~~ + +回滚结束并准备恢复正常发布轨道时,再切回维护者指定的分支: + +~~~bash +git switch main +git pull --ff-only origin main +~~~ + +如果是 Nginx 配置导致故障,先恢复 root-only 备份,确认旧域名配置没有重复 include,再执行: + +~~~bash +sudo nginx -t +sudo systemctl reload nginx +~~~ + +回滚后重新检查本机和公网 /health、Chat 401/404 边界以及 MCP 初始化。docker compose down 只停止容器,不删除代码或数据库;明确需要停服时再执行: + +~~~bash +cd /home/firee-safety-ymd +docker compose down +~~~ + +## 10. 完成判定与当前限制 + +本手册完成不代表公网部署已完成。现场交付至少应记录: + +- 目标机 Compose 配置校验、build、up、健康检查和无 Secret 的有限日志结果。 +- 目标机 nginx -t 和 reload 成功;443 证书、DNS、安全组及旧 location / 已确认不再生效。 +- Chat 首轮/多轮 SSE、错误凭证、未知 app/path 的实际 HTTPS 响应。 +- SuperAgent -> /mcp 的独立 Bearer、TLS/网络白名单、工具发现和至少一个受控只读工具调用。 +- 数据库未公开 5432,运行账号保持只读,PostGIS readiness warning 和 35 条无效面几何缺口已记录。 + +若目标服务器无法使用 host.docker.internal、Compose 不支持 --quiet、Nginx include 目录不同、证书路径不同或 SuperAgent 对 MCP 的认证格式不同,应先记录实际环境并调整部署 Spec;不要通过放开端口、写入 Token、关闭 readiness 或恢复全路径反代规避问题。 diff --git a/docs/project/operations/nginx-public-entry.md b/docs/project/operations/nginx-public-entry.md new file mode 100644 index 0000000..90dd3c2 --- /dev/null +++ b/docs/project/operations/nginx-public-entry.md @@ -0,0 +1,175 @@ +# Nginx 公网入口(受控联调示例) + +| 项 | 内容 | +| --- | --- | +| 目标 | 让 `agent.nianxx.com` 通过 HTTPS 访问本项目的兼容对话接口和 MCP | +| 状态 | 示例配置;公网机器、证书、网络白名单和真实鉴权仍待联调 | +| 上游 | 仅反代本机 `127.0.0.1:8080` | + +## 1. 与旧配置的变化 + +旧配置把整个 `/` 转发到 DashScope,并在 Nginx 中比较 `xtoken`、注入 DashScope API Key。新边界是: + +```text +公网 HTTPS + -> Nginx(TLS、精确路径、限流、SSE 传输设置) + -> 127.0.0.1:8080(fire-safety-ymd) + -> Go 侧校验 xtoken / MCP Bearer +``` + +Nginx 不再调用 DashScope,也不保存或注入 SuperAgent Open API Key、Chat `xtoken`、MCP Bearer 或数据库凭证。Go 服务负责把用户兼容请求转发到 SuperAgent,并负责 `/mcp` 的独立 Bearer 校验。这样可避免把 Secret 写入 Nginx 配置,也避免使用 Nginx `if` 比较 Secret。 + +## 2. 可复制的配置示例 + +配置文件位于 [`../../../deploy/nginx/fire-safety-ymd.conf.example`](../../../deploy/nginx/fire-safety-ymd.conf.example)。它只公开以下路径: + +| 路径 | 用途 | 鉴权与限制 | +| --- | --- | --- | +| `/api/v1/apps//completion` | 截图所示的 DashScope-compatible 用户对话 SSE | Go 校验 `xtoken`;请求体 128 KiB;示例限流 5 req/s、burst 20 | +| `/mcp` | SuperAgent 调用本项目的 MCP | Go 校验独立 `Authorization: Bearer`;请求体 256 KiB;示例限流 20 req/s、burst 40 | +| `/health` | 可选进程存活检查 | 不访问数据库;如不希望公开可删除该 location | +| 其他路径 | 不对外提供 | Nginx 固定返回 404 | + +将示例放入 Nginx 的 `http` 配置范围(例如 `/etc/nginx/conf.d/fire-safety-ymd.conf`)前,必须完成以下替换: + +1. 把 `location = /api/v1/apps/replace-with-fire-safety-app-id/completion` 中的占位符替换为安全的 app ID。 +2. 该 app ID 必须与 Go 服务的 `FIRE_SAFETY_CHAT_COMPAT_APP_ID` 完全一致;两处不一致会导致请求被拒绝。 +3. 确认证书路径 `/cert/agent.nianxx.com.pem` 和 `/cert/agent.nianxx.com.key` 在公网机器上存在,并由 Nginx 进程可读。 +4. 直接在宿主机运行 Go 时,确认服务只监听 `127.0.0.1:8080`;使用本仓库 Compose 时,容器内监听 `:8080`,但端口映射必须保持为 `127.0.0.1:8080:8080`。两种方式都不要把 8080 或 PostgreSQL 5432 暴露到公网。 +5. 按环境配置 DNS、云防火墙/安全组和 SuperAgent 对 `/mcp` 的来源 IP/TLS 要求;这些部署事实尚未由本项目验证。 + +`limit_req_zone` 必须位于 Nginx `http` context,不能放进 `server` 或 `location`。示例文件假定它被 `conf.d/*.conf` 从 `http {}` 中 include;如果部署系统不是这样 include,应把两条 `limit_req_zone` 指令单独移到 `http {}`,并保留 `server`/`upstream` 在合法上下文。 + +## 3. Go 服务配置 + +在服务器的未提交 Secret 管理方式(例如权限收紧的 `.env`、systemd EnvironmentFile 或容器 Secret)中配置,不要写入 Nginx 文件或 Git: + +```text +# 宿主机直接运行时使用 127.0.0.1:8080;Compose 会覆盖为容器内的 :8080, +# 并只把端口发布到宿主机回环地址。 +FIRE_SAFETY_HTTP_ADDR=127.0.0.1:8080 + +FIRE_SAFETY_SUPERAGENT_ENABLED=true +FIRE_SAFETY_SUPERAGENT_BASE_URL=https:// +FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY= + +FIRE_SAFETY_CHAT_ENABLED=true +FIRE_SAFETY_CHAT_AUTH_TOKEN= +FIRE_SAFETY_CHAT_COMPAT_APP_ID= + +FIRE_SAFETY_MCP_ENABLED=true +FIRE_SAFETY_MCP_AUTH_TOKEN= +FIRE_SAFETY_POSTGIS_ENABLED=true +FIRE_SAFETY_POSTGIS_DSN= +``` + +当前仍是受控联调静态门禁: + +- `FIRE_SAFETY_CHAT_AUTH_TOKEN` 由用户请求的 `xtoken` Header 携带,必须与 SuperAgent Open API Key、MCP Token 使用不同值。 +- `FIRE_SAFETY_MCP_AUTH_TOKEN` 只用于 SuperAgent -> Go `/mcp`,不应复用 Chat Token。 +- Chat、SuperAgent、MCP 和 PostGIS 的完整配置校验以 `.env.example` 和对应项目文档为准。 +- 浏览器会看到 `xtoken`;它不能代表最终用户身份、角色、租户或数据授权。公网真实用户入口仍需身份提供方、动态授权、限流、Secret 轮换和持久审计。 + +测试服务器使用 Docker Compose 的完整目录、启动、更新和回滚步骤见 [`docker-test-deployment.md`](docker-test-deployment.md)。 + +## 4. SSE 代理边界 + +兼容对话响应是长连接 SSE。示例对该精确路径设置: + +- `proxy_http_version 1.1`、清空 `Connection`,保持上游流式连接。 +- `proxy_buffering off`、`proxy_cache off`、`proxy_request_buffering off` 和 `gzip off`,并返回 `X-Accel-Buffering: no`。 +- `proxy_read_timeout`/`proxy_send_timeout` 为 660 秒,覆盖 Go 默认 10 分钟总运行时限并留出少量余量;如提高 Go 的运行时限,必须同步提高 Nginx 超时。 +- `proxy_next_upstream off`,避免 POST 流中断后 Nginx 自动重试而产生重复 SuperAgent Run。 +- 限流拒绝使用 HTTP 429 和稳定 JSON;客户端仍需把网关层非 SSE 响应视为失败,不能等待 `finish_reason=stop`。 +- 透传 `Host`、`X-Forwarded-For`、`X-Forwarded-Proto`、`X-Forwarded-Host`、`X-Forwarded-Port`、`X-Forwarded-Server` 和 `X-Request-ID`。 +- 不设置 `proxy_set_header xtoken ...` 或 `proxy_set_header Authorization ...`;客户端送来的 `xtoken` 原样到 Go,由 Go 做常量时间比较,MCP Bearer 同理。 + +MCP 是普通 JSON 请求,不需要 SSE 的关闭响应缓冲设置;示例使用 30 秒读写超时,仍关闭上游自动重试。 + +## 5. 启用与检查 + +启动 Go 服务前,在服务器的进程环境中加载 Secret。不要把 Secret 直接写进命令行或 shell 历史;可以在受保护的环境文件中加载后,再通过 Header 环境变量展开: + +直接在宿主机运行时可使用 `go run ./cmd/server`;测试服务器的默认方式是由 Compose 启动镜像,不在宿主机安装或运行 Go 工具链。 + +在 reload 前先检查 Nginx 配置;示例机器没有安装 Nginx 时,本地无法代替目标服务器完成这一检查: + +```bash +sudo nginx -t +sudo nginx -s reload +``` + +### 首轮兼容对话 SSE + +以下命令不会把 Token 字面量写入命令历史;`CHAT_XTOKEN` 应由受保护的环境文件或 Secret 管理器注入当前 shell: + +```bash +curl -N --fail \ + -H "xtoken: ${CHAT_XTOKEN}" \ + -H 'Content-Type: application/json' \ + -H 'Accept: text/event-stream' \ + --data '{"input":{"prompt":"观水镇附近有哪些水源候选?"},"parameters":{}}' \ + 'https://agent.nianxx.com/api/v1/apps//completion' +``` + +客户端应读取 `event: result` 的 SSE 数据,等待 `output.finish_reason` 为 `stop` 后再把最终 `output.text` 视为完成;中途断流不能当作可信答案。首轮返回的 `output.session_id` 是本项目的本地会话 ID,后续是否能继续多轮取决于兼容接口的契约,不能自行把它当作 SuperAgent Provider Session。 + +### 多轮兼容对话 SSE + +使用上一轮服务返回并确认成功的会话 ID;不要把 SuperAgent 的内部 Session ID 放入客户端请求: + +```bash +curl -N --fail \ + -H "xtoken: ${CHAT_XTOKEN}" \ + -H 'Content-Type: application/json' \ + -H 'Accept: text/event-stream' \ + --data '{"input":{"prompt":"再说明这些候选的限制","session_id":""},"parameters":{}}' \ + 'https://agent.nianxx.com/api/v1/apps//completion' +``` + +如果目标客户端不接受 `text/event-stream` 或只需要一次性 JSON,应先以接口 Spec 为准;不要让 Nginx 擅自将 SSE 缓冲成普通 JSON。 + +### MCP 探活/初始化 + +MCP 请求需要独立 Token,具体 JSON-RPC body 以 MCP Spec 为准: + +```bash +curl --fail \ + -H "Authorization: Bearer ${MCP_BEARER}" \ + -H 'Content-Type: application/json' \ + --data '{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"manual-probe","version":"1"}}}' \ + 'https://agent.nianxx.com/mcp' + +curl --fail --output /dev/null --write-out 'initialized HTTP %{http_code}\n' \ + -H "Authorization: Bearer ${MCP_BEARER}" \ + -H 'Content-Type: application/json' \ + --data '{"jsonrpc":"2.0","method":"notifications/initialized"}' \ + 'https://agent.nianxx.com/mcp' + +curl --fail \ + -H "Authorization: Bearer ${MCP_BEARER}" \ + -H 'Content-Type: application/json' \ + --data '{"jsonrpc":"2.0","id":2,"method":"tools/list"}' \ + 'https://agent.nianxx.com/mcp' +``` + +预期 `initialize` 返回协议版本 `2025-06-18`,initialized notification 返回 HTTP 202,`tools/list` 返回 7 个固定工具。只收到 HTTP 200 但 JSON-RPC 中含 `error` 也属于失败,不能把它当作 MCP 已就绪。 + +## 6. 验收边界与回滚 + +完成 `nginx -t` 和 reload 后,应至少验证: + +1. `/health` 返回 200(如果保留可选 location)。 +2. 兼容对话首轮在正确 `xtoken` 下返回 SSE,末尾出现 `finish_reason: stop`;错误 Token 返回 401,且 Go 不会创建 SuperAgent Run。 +3. 兼容对话多轮使用返回的本地会话 ID;服务重启或会话失效后按契约返回会话不存在,不把 ID 当成 Provider Session。 +4. `/mcp` 使用独立 Bearer,错误或缺失 Bearer 返回 401;Nginx 没有注入任何 MCP Token。 +5. 未列出的路径(例如 `/api/chat`、`/api/v1/apps/other/completion` 和 `/anything`)返回 404。 + +如果 reload 后发现兼容路由、证书或 SSE 行为异常,先恢复上一个已验证的 Nginx 配置并保留 `nginx -t` 输出;不要通过放开 `location /`、关闭 Go 鉴权或把 Secret 写入 Nginx 来排障。 + +## 7. 未确认事项 + +- 目标公网机器的 Nginx 版本、include 层级、TLS 终止位置和证书续期方式;示例同时监听 80 做 HTTPS 跳转,安全组需按实际策略决定是否允许 80。 +- `agent.nianxx.com` 的 DNS、安全组、反向代理来源 IP 和 SuperAgent 对 MCP 回调的网络策略。 +- 兼容接口在目标 SuperAgent/Profile 中对 `input.session_id`、`parameters` 和 SSE `result` 事件的最终协议;实现以本项目 Spec 和实际联调为准。 +- 生产最终用户认证、动态授权、会话共享/持久化、主动取消、配额、审计和 Secret 轮换。 diff --git a/docs/project/operations/postgis-srid-4326.md b/docs/project/operations/postgis-srid-4326.md new file mode 100644 index 0000000..f6d0de3 --- /dev/null +++ b/docs/project/operations/postgis-srid-4326.md @@ -0,0 +1,81 @@ +# PostGIS SRID 4326 元数据迁移 + +| 项 | 内容 | +| --- | --- | +| 目标 | 将 8 张已确认 WGS84、当前 SRID 0 的源表几何标记为 EPSG:4326 | +| 状态 | 2026-09-05 已执行并通过迁移后严格 readiness | +| 命令 | `go run ./cmd/postgis-srid-migrate` | +| 数据决定 | 不修复 35 条无效几何,不改变 7 条空几何,不丢弃 Z | + +## 已确认前提 + +- 数据提供方确认 8 张表的源坐标系均为 EPSG:4326,且无坐标偏移。 +- 2026-09-05 只读预检发现 4,048 条非空几何全部为 SRID 0,坐标数值未越过 WGS84 范围。 +- `st_2_xianyoufanghuotongdao.geom` 是不限定 typmod 的 `geometry`,包含 1,026 条二维和 2 条三维非空几何。 +- 其余 7 张表是二维 `geometry(Geometry)` typmod,typmod SRID 为 0;必须由表所有者或等效迁移角色改为 `geometry(Geometry,4326)`。 +- 原始导入 SQL 已由用户另行保存。执行前仍建议对当前数据库创建快照,因为原始导入文件不包含导入后的潜在变化。 + +## 本次执行结果 + +- 使用与运行时只读账号分离的表所有者 `admin` 执行。 +- 4,048 条非空几何全部从 SRID 0 补齐为 SRID 4326;迁移后 SRID 0 计数为 0。 +- 7 张二维表改为 `geometry(Geometry,4326)`;防火通道保持裸 `geometry`。 +- 防火通道的 1,026 条二维和 2 条三维非空几何分布未改变。 +- 无 SRID WKB 指纹、总行数、空值、几何类型、有效性、范围和 Z/M 计数迁移前后一致,工具报告 `geometry_payload_unchanged=true`。 +- 只读严格 probe 报告 `MCP spatial validation passed`;35 条无效面几何、7 条空几何和 7 张缺 GiST 索引表继续作为预期 warning。 + +## 安全行为 + +迁移工具固定允许 8 张 `public` 表,不接受任意表名或目标 SRID。默认模式只读;写入必须同时提供: + +```text +--apply --confirm-source-crs=EPSG:4326 +``` + +写入在一个 serializable 事务中完成,并使用 advisory lock、5 秒锁等待上限和 120 秒语句上限。事务内会: + +1. 再次确认只出现 SRID 0/4326、类型符合表契约且坐标未越界。 +2. 将 7 张旧二维 typmod 表原位改为 `geometry(Geometry,4326)`,使用 `ST_SetSRID` 补充元数据。 +3. 保持防火通道列为裸 `geometry`,只对 SRID 0 的记录执行 `ST_SetSRID(geom,4326)`,保留混合 2D/Z。 +4. 比较迁移前后的无 SRID WKB 指纹、行数、空值、有效性、类型、坐标维度以及 Z/M 计数。 +5. 只有所有检查通过才提交;任一检查失败会回滚整个事务。 + +工具不会调用 `ST_Transform`、`ST_MakeValid` 或 `ST_Force2D`。 + +## 凭证边界 + +运行时 `FIRE_SAFETY_POSTGIS_DSN` 应继续使用只读角色。迁移时另行在未提交的 `.env` 中配置: + +```text +FIRE_SAFETY_POSTGIS_MIGRATION_DSN=postgresql://<临时迁移账号>:@:/ +``` + +该账号必须能更新全部 8 张表,并能修改 7 张旧 typmod 表的 `geom` 列;通常需要使用表所有者或数据库管理员临时执行。迁移完成后应从 `.env` 和部署环境移除该凭证。 + +## 执行步骤 + +```bash +set -a +source .env +set +a + +# 只读预检 +go run ./cmd/postgis-srid-migrate + +# 显式写入 +go run ./cmd/postgis-srid-migrate \ + --apply \ + --confirm-source-crs=EPSG:4326 + +# 使用运行时只读账号执行严格验证 +FIRE_SAFETY_POSTGIS_EXPECTED_SRID=4326 \ + go run ./cmd/postgis-probe +``` + +预期迁移后 8 张表的 `srids` 均为 `[4326]`。35 条无效几何、7 条空几何及缺少 GiST 索引仍应作为 warning 出现。 + +## 回滚边界 + +- 提交前的任何失败自动整笔回滚。 +- 提交后的恢复优先使用执行前数据库快照或用户保存并核验过的导入源。 +- 不提供面向未来数据的盲目 `4326 -> 0` 回滚命令,避免把迁移后新写入且本来正确的数据错误标记为未知 SRID。 diff --git a/docs/project/security-access-control-boundary.md b/docs/project/security-access-control-boundary.md new file mode 100644 index 0000000..d173b19 --- /dev/null +++ b/docs/project/security-access-control-boundary.md @@ -0,0 +1,133 @@ +# 安全与访问控制边界 + +| 项 | 内容 | +| --- | --- | +| 项目 | `fire-safety-ymd` | +| 状态 | SuperAgent 出站、用户对话 API v1 与空间只读 MCP 代码控制已实现;真实环境、最终用户授权和持久审计待完成 | +| 最近更新 | 2026-09-05 | + +## 1. 目的 + +本文定义用户侧应用、Go 服务、SuperAgent、MCP 工具和 PostgreSQL/PostGIS 之间的信任边界。当前已实现服务间 MCP 基线控制,但不代表最终用户鉴权、实库 readiness、生产网络或持久审计已经完成。 + +## 2. 当前实现状态 + +| 能力 | 状态 | 当前含义 | +| --- | --- | --- | +| `GET /health` | 已实现 | 公开的进程存活响应,不访问业务数据 | +| 用户认证 | 仅首版联调门禁 | `/api/chat` 使用独立静态 Bearer;可选兼容路径用同一信任方向的 `xtoken`;它们不是最终用户身份,身份提供方和正式 Token 格式待确认 | +| 角色/区域/租户授权 | 部分实现 | MCP 使用显式服务端静态范围:默认镇街白名单,或经授权的数据库全范围;最终用户动态身份/角色/租户仍未实现 | +| SuperAgent Open API Adapter | 已实现、模拟测试通过 | 默认关闭;Secret 仅由项目环境变量注入;真实环境未联调 | +| 用户聊天与会话映射 | 已实现、默认关闭 | 原生 `POST /api/chat` 和可选 DashScope 风格 `completion` 路径、严格 SSE 最终回答、单进程内存映射和同会话并发冲突;无持久化、跨实例恢复或最终用户授权 | +| MCP 工具 | 已实现、默认关闭 | 7 个固定只读工具(含地名候选搜索);独立 Bearer、body/Origin/schema/半径/数量/超时控制;实库严格 readiness 和本地真实工具查询已通过,公网/SuperAgent 待联调 | +| PostgreSQL/PostGIS | SRID 严格 readiness 已通过 | pgxpool、连接级默认只读、固定参数化 SQL;4,048 条非空几何已补齐 SRID 4326,35 条无效几何继续排除并告警 | +| 测试环境容器/公网入口 | 配置基线已建立、目标机待验证 | 单实例非 root 容器,只发布宿主机回环端口;Nginx 终止 TLS 并只公开精确路径;运行时不接收迁移凭证 | +| 审计存储 | 未实现 | 事件、字段和保留周期待确认 | + +## 3. 信任区域 + +| 区域 | 信任判断 | 主要控制 | +| --- | --- | --- | +| 用户输入 | 不可信 | 身份认证、大小限制、格式校验、提示注入隔离 | +| Chat 静态 Bearer | 仅受控联调可信 | 独立 Secret、精确 CORS Origin、默认关闭;不能充当用户身份或字段级授权 | +| Go 服务端身份上下文 | 授权事实的唯一入口 | Token 验证、角色与区域解析、不可被模型覆盖 | +| SuperAgent 输入与输出 | 外部、非权威 | 最少数据、超时、输出约束、工具侧重新授权 | +| MCP 参数 | 不可信,即使由 Agent 生成 | schema 校验、范围上限、服务端注入授权上下文 | +| Repository 查询 | 受控代码 | 参数化 SQL、白名单、只读角色、超时和行数上限 | +| PostgreSQL/PostGIS | 业务事实来源但可能陈旧或不完整 | 权限、数据质量、来源与更新时间标记 | +| 日志与审计 | 受限数据面 | 脱敏、访问控制、保留和完整性策略 | + +## 4. 身份与授权原则 + +- 客户端提交的 `user_id`、角色、区域或租户声明不能直接作为授权依据。 +- 首版 `/api/chat` DTO 根本不接受 `user_id`、`external_subject_id`、角色、区域、Provider Session 或 metadata;服务端固定测试主体不能用于用户级授权。 +- Go 服务从经过验证的凭证中建立可信身份上下文,并将授权范围与请求生命周期绑定。 +- SuperAgent 不做最终授权;每次 MCP 调用都在服务端重新校验工具权限和数据范围。 +- 模型可提供地点、半径、资源类型等业务查询条件,但不能自由指定要冒充的用户、租户或权限。 +- 默认拒绝;新增工具、字段或数据表必须进入显式权限矩阵。 +- 联系人、电话和精确敏感位置按字段级权限控制,不因同一记录的普通字段可见而自动可见。 +- `FIRE_SAFETY_MCP_SCOPE_MODE=all` 是 MCP 固定查询表的数据库全范围授权,只能在该 MCP 凭证确实获准读取这些表全部记录时显式使用;它不会从业务数据自动推导授权。 +- `all` 与镇街白名单不能同时配置,避免操作者误以为白名单仍限制结果;两种模式都不能由模型或工具参数修改。 +- 地名搜索只读取已允许的名称、镇街、村庄和几何字段,仍应用相同服务端数据范围;不读取联系人、电话或任意表。 +- 地名匹配输出属于候选。线和面的代表点只用于帮助用户识别记录,未经用户确认不得作为演练点继续查询或形成距离结论。 + +## 5. MCP 工具边界 + +MCP 第一阶段已按以下只读边界实现: + +- 工具对应稳定领域能力,不提供 `execute_sql`、任意表查询、文件读取或通用 HTTP 代理。 +- JSON schema 限制字符串长度、枚举、坐标范围、半径、分页和最大结果数。 +- Repository 只访问允许的 schema、表和字段,动态标识符来自代码白名单。 +- 输出只包含回答所需字段,并表达数据来源、更新时间、状态和不确定性。 +- 权限不足与“没有数据”使用不同稳定状态,避免 Agent 猜测。 +- 调用记录至少关联 request ID、可信主体、工具名、授权范围、耗时、结果类别和数据源版本;审计日志不保存无必要的完整敏感正文。 + +当前日志已记录 request ID、操作、耗时和结果类别;可信主体、数据源版本和持久审计尚未实现,因此生产前仍需补齐。 + +任何写工具、资源调度或状态变更都需要新的 Spec、幂等设计、人工确认边界、审计和安全 Review,不属于默认扩展。 + +## 6. PostgreSQL/PostGIS 边界 + +- 应用数据库角色使用最小权限;MCP 查询角色默认只授予必要表或视图的 `SELECT`。 +- 数据迁移凭证与运行时只读凭证分离;2026-09-05 的一次性 SRID 迁移使用表所有者执行,完成后不应保留在应用部署环境。 +- Compose 从未提交的 `.env` 注入运行时配置,但会把 `FIRE_SAFETY_POSTGIS_MIGRATION_DSN` 强制覆盖为空;镜像构建不得复制 `.env`、证书、SQL 导出或 Excel 数据。 +- 应用容器只发布宿主机 `127.0.0.1:8080`,公网只经 Nginx 的 HTTPS 精确路径进入;8080 和数据库端口不得加入公网安全组。 +- 使用参数化查询,不把用户或模型内容拼入 SQL、表名、排序或空间表达式。 +- 设置连接、查询和 statement timeout;限制半径、结果数和空间复杂度,防止昂贵查询。 +- 优先通过项目领域视图或明确映射屏蔽历史物理表名和不一致字段。 +- 验证 extension、SRID、几何类型、有效性、索引和单位后才开放空间结论。当前无效几何不参与结论:原始记录保留,readiness 告警,固定查询使用 `ST_IsValid` 排除,工具响应声明结果可能不完整。 +- 原始 SQL 中的 `DROP TABLE`、owner、sequence 和 extension 操作必须在隔离环境审计,不对现有数据库直接执行。 +- 数据库错误对外映射为稳定错误类别,不返回 DSN、SQL 文本、schema 细节或堆栈。 + +## 7. SuperAgent 数据边界 + +- 只发送完成当前请求所需的最少上下文,不默认发送联系人、电话、完整数据库记录或长期会话历史。 +- SuperAgent 凭证仅保存在服务端 Secret,按环境隔离并支持轮换。 +- Open API Key 不得复用 TH Hotel 或其他项目凭证;本项目只读取 `FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY`。 +- Provider 的 `Content-Location` 必须与配置 Base URL 同源;客户端禁止跟随重定向,避免跨域转发 Authorization。 +- SSE 只有在最终内容、成功 `run.completed` 和顶层 `end` 同时存在时成功;断流恢复不得重发原消息。 +- 公开 Trace 投影不包含原始工具输入、输出、Header、Cookie 或 Provider 原始 payload。 +- 明确平台的数据留存、训练使用、区域、子处理方和删除能力后,才能发送受限数据。 +- 流式事件和工具调用都按不可信外部数据解析,设置消息大小、事件类型和状态机约束。 +- 上游答案返回用户前应保留安全提示、数据时间和工具错误状态,不能把模型文本升级为权威事实。 +- `/api/chat` 不转发消息 delta;只有 Adapter 的严格成功条件全部满足后才发送最终 `message` 和 `done`。公开进度只保留安全的 `run.*` / `tool.*` 事件、工具名和状态。 +- 兼容 `completion` 路径同样不转发消息 delta;它先返回无正文的 `finish_reason="null"`,严格成功后才发送 `finish_reason="stop"` 和最终正文。URL App ID 只是公开路由标识,不是授权依据。 +- 兼容路径的 `xtoken` 复用 Chat 入站凭证而不是 SuperAgent Key。Nginx 不保存、比较或注入任何应用 Secret,只把 Header 交给 Go 校验。 +- 同一内存会话只允许一个活动 Run。流失败或结果不确定时删除本地映射,避免继续复用可能仍在运行的 Provider Session。 +- 当前会话只在单进程内存中保存且有数量/TTL 上限;不保存消息历史。重启或多实例切换会导致旧对话不可用,不得向用户承诺持久会话。 + +## 8. 敏感数据与日志 + +至少按以下类别处理: + +| 数据 | 基线分类 | 日志策略 | +| --- | --- | --- | +| API Key、Token、Cookie、数据库密码 | Secret | 禁止记录 | +| 联系人、电话、用户标识 | 个人/受限 | 默认脱敏,不记录完整值 | +| 精确设施或风险区域坐标 | 业务敏感 | 按权限最少披露,不记录完整结果集 | +| 用户问题和对话 | 可能含敏感信息 | 默认不记录原文,使用摘要或分类字段 | +| 工具名、耗时、结果类别 | 运行元数据 | 可记录,不附敏感 payload | + +生产前需确定数据分类负责人、日志访问角色、保留周期、删除流程和安全事件响应方式。 + +当前 Chat 日志只记录 request ID、结果类别、是否复用和耗时,不记录消息、回答、对话 ID、Provider Session 或 Trace payload。`docs/import/db-samples/*.sql` 含真实联系人、电话和精确坐标,仅作为本地只读输入并由 Git 忽略;不得执行或进入版本历史。长期测试数据必须另做脱敏 fixture。 + +## 9. 应急场景安全 + +- 系统只提供辅助信息,不替代报警、撤离和现场指挥。 +- 数据缺失、陈旧、无状态或工具失败时必须显式说明,不能生成看似确定的资源可用结论。 +- 涉及路线、危险源或处置禁忌的功能需要专门规则、权威来源、版本和测试,不仅依赖模型提示词。 +- 任何可能延误报警或撤离的交互流程都必须在实现前进行安全审查。 + +## 10. 待确认决策 + +- 身份提供方、Token 验证方式和服务间认证。 +- `/api/chat` 与兼容 `completion` 路径从静态联调凭证迁移到真实用户身份的方案,以及共享会话、主动取消、限流和滥用防护。 +- 是否存在多租户;若不存在,区域/组织数据范围如何表达。 +- 角色与字段级权限矩阵,尤其是联系人和精确位置。 +- SuperAgent 的部署方、数据处理条款、鉴权、MCP 回调认证与网络边界。 +- 数据库网络拓扑、只读账号、视图策略和审计能力。 +- 对话、工具调用和安全审计的存储位置与保留周期。 +- 面向真实火情功能的责任主体、免责声明和人工确认机制。 + +这些问题影响实现方向,应在对应 Feature 的 Definition of Ready 阶段确认,不在代码中自行假设。 diff --git a/docs/specs/fire-safety-ymd-chat-api-v1.md b/docs/specs/fire-safety-ymd-chat-api-v1.md new file mode 100644 index 0000000..4f9042a --- /dev/null +++ b/docs/specs/fire-safety-ymd-chat-api-v1.md @@ -0,0 +1,164 @@ +# 用户对话 API v1 Spec + +| 项 | 内容 | +| --- | --- | +| 状态 | Implemented | +| 日期 | 2026-09-05 | +| 负责人 | fire-safety-ymd 后端 | +| 需求来源 | 用户要求“先做对话 API” | +| 关联 Change Request | 无 | + +## 1. 背景 + +项目已经具备 SuperAgent Open API 出站 Adapter 和空间只读 MCP,但用户侧应用还没有安全的服务端对话入口。前端不能直接持有 SuperAgent Open API Key,也不能自行指定 SuperAgent Session、用户身份或授权范围。 + +本 checkpoint 提供一个最小可联调的对话 API。它由 Go 服务创建并复用 SuperAgent Session,通过 SSE 返回安全进度和最终回答。真实用户认证、动态区域授权、数据库会话持久化和生产审计后续单独设计。 + +## 2. 目标 + +- 新增默认关闭的 `POST /api/chat`。 +- 使用独立静态 Bearer Token 保护首版测试入口,不复用 SuperAgent Key 或 MCP Token。 +- 首轮创建随机本地 `conversation_id` 和 SuperAgent Session;后续轮次复用映射。 +- 同一 `conversation_id` 同时只允许一个活动 Run,冲突时返回 `409`。 +- 通过 SSE 返回对话 ID、安全进度、严格完成后的最终回答和完成事件。 +- 不把消息正文、回答正文、Secret、原始工具参数或输出写入日志。 +- 对输入大小、总运行时间、内存会话数量、空闲过期时间和浏览器 Origin 设上限。 + +## 3. 非目标 + +- 不实现注册、登录、JWT、SSO、角色、租户或最终用户级数据授权。 +- 不接受客户端提交的 `user_id`、`external_subject_id`、SuperAgent `session_id`、角色或区域范围。 +- 不持久化对话;进程重启或多实例切换后,旧 `conversation_id` 不可继续使用。 +- 不返回逐字 token、模型原始思考、原始 Trace、工具输入或工具输出。 +- 不实现历史消息查询、会话列表、主动取消、重试队列、限流或 WebSocket。 +- 不把 Agent 回答写入消防业务事实。 + +## 4. HTTP 契约 + +### 4.1 请求 + +```http +POST /api/chat +Authorization: Bearer +Content-Type: application/json +Accept: text/event-stream + +{ + "message": "观水镇附近有哪些可用水源?", + "conversation_id": "conv_..." +} +``` + +- `message` 必填,去除首尾空白后不能为空,并服从 SuperAgent 单条消息字节上限。 +- `conversation_id` 首轮省略;后续使用服务返回的随机值。 +- 未知字段、多个 JSON 值、非法媒体类型和超大请求体均拒绝。 +- 客户端不得提交身份、权限、Provider Session 或任意 metadata。 + +### 4.2 成功响应 + +响应媒体类型为 `text/event-stream`。事件顺序如下: + +```text +event: conversation +data: {"conversation_id":"conv_...","reused":false} + +event: progress +data: {"event":"tool.started","tool_name":"fire_safety_search_place_candidates","status":"running"} + +event: message +data: {"conversation_id":"conv_...","answer":"..."} + +event: done +data: {"conversation_id":"conv_...","run_id":"...","usage":{"input":0,"output":0,"total":0}} +``` + +- `conversation` 在 Provider Session 准备完成后首先发送。 +- `progress` 只包含经过现有 Adapter 清洗的事件名、工具名和状态;不包含文本、ID、参数或工具结果。 +- 没有业务事件时,服务每 15 秒发送一个 SSE comment heartbeat,保持长连接且不携带业务数据。 +- `message` 只在 Adapter 同时确认最终内容、成功 `run.completed` 和顶层 `end` 后发送。 +- `done` 是成功终止事件。 + +### 4.3 错误响应 + +在 SSE 开始前,错误使用 HTTP 状态码和稳定 JSON 错误,例如: + +```json +{ + "error": { + "code": "CHAT_AUTH_INVALID", + "message": "Chat authentication failed." + }, + "request_id": "..." +} +``` + +SSE 开始后的 Provider 失败使用终止 `error` 事件,不返回部分回答。主要错误类别: + +- `CHAT_REQUEST_INVALID`:输入或 JSON 不合法,HTTP 400。 +- `CHAT_AUTH_INVALID`:Bearer 无效,HTTP 401。 +- `CHAT_ORIGIN_FORBIDDEN`:浏览器 Origin 未获准,HTTP 403。 +- `CHAT_CONVERSATION_NOT_FOUND`:会话不存在、已过期或服务已重启,HTTP 404。 +- `CHAT_CONVERSATION_BUSY`:同一会话已有活动 Run,HTTP 409。 +- `CHAT_CAPACITY_REACHED`:内存会话达到上限,HTTP 503。 +- `CHAT_UPSTREAM_TIMEOUT`、`CHAT_UPSTREAM_UNAVAILABLE`、`CHAT_UPSTREAM_PROTOCOL_ERROR`、`CHAT_RUN_FAILED`:上游运行失败;SSE 尚未开始时使用 502/504,否则发送 `error` 事件。 + +## 5. 会话与并发规则 + +- 映射仅保存在当前 Go 进程内:`conversation_id -> SuperAgent session_id`。 +- `conversation_id` 由加密安全随机数生成,不包含用户、镇街或业务语义。 +- 新会话使用服务端固定测试主体 `FIRE_SAFETY_CHAT_SUBJECT_ID` 创建;客户端不能覆盖。 +- 每次发送消息使用新的幂等键和请求关联 ID。 +- 同一会话的 Run 使用互斥占用;不同会话可并发。 +- 空闲会话超过 TTL 后惰性清理;活动 Run 不清理。 +- 达到最大会话数后先清理过期会话,仍满则拒绝创建。 + +固定测试主体只是首版联调边界,不代表真实用户认证,也不能用于用户级授权或审计。 + +## 6. 配置契约 + +| 环境变量 | 默认值 | 说明 | +| --- | --- | --- | +| `FIRE_SAFETY_CHAT_ENABLED` | `false` | 对话 API 总开关;启用时要求 SuperAgent 同时启用 | +| `FIRE_SAFETY_CHAT_AUTH_TOKEN` | 空 | 独立高熵静态 Bearer,至少 32 个可打印 ASCII 字符 | +| `FIRE_SAFETY_CHAT_SUBJECT_ID` | `fire-safety-ymd-chat-test-subject` | 服务端固定测试主体,不使用真实用户标识 | +| `FIRE_SAFETY_CHAT_ALLOWED_ORIGINS` | 空 | 逗号分隔的精确 HTTP(S) Origin;空值只允许无 Origin 的服务端/CLI 调用 | +| `FIRE_SAFETY_CHAT_MAX_BODY_BYTES` | `131072` | HTTP JSON 请求体上限,最大 1 MiB | +| `FIRE_SAFETY_CHAT_RUN_TIMEOUT` | `10m` | 创建 Session 加单轮 Run 的总时限,最大 30 分钟 | +| `FIRE_SAFETY_CHAT_SESSION_TTL` | `30m` | 空闲内存会话保留时间,最大 24 小时 | +| `FIRE_SAFETY_CHAT_MAX_SESSIONS` | `1000` | 单进程最大内存会话数,最大 10000 | + +`FIRE_SAFETY_CHAT_AUTH_TOKEN` 必须分别不同于 `FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY` 和 `FIRE_SAFETY_MCP_AUTH_TOKEN`。 + +## 7. CORS 与鉴权边界 + +- 无 `Origin` 的 CLI/服务端请求可以进入 Bearer 校验。 +- 带 `Origin` 的浏览器请求必须精确匹配配置列表;不支持 `*`。 +- 预检只允许 `POST` 以及 `Authorization`、`Content-Type` Header。 +- 不启用 Cookie 身份或 CORS credentials。 +- 静态 Bearer 只适合首版受控联调;公网最终用户入口必须接入真实身份认证、速率限制和动态授权。 + +## 8. 日志与数据边界 + +- 只记录 request ID、结果类别、是否复用会话和耗时。 +- 不记录 Bearer、Open API Key、消息、回答、Provider payload、坐标、工具参数或工具输出。 +- 对外错误不包含上游响应正文、URL、Session ID、堆栈或 Secret。 +- SuperAgent 回答属于辅助内容,不自动升级为权威事实,也不替代报警、撤离和现场指挥。 + +## 9. 验收标准 + +- Given 对话 API 未启用,When 请求 `/api/chat`,Then 返回 404 且不创建 SuperAgent Client 调用。 +- Given 未授权或 Origin 不在白名单,When 请求 API,Then 在读取和转发消息前拒绝。 +- Given 首轮合法消息,When Provider 严格成功,Then 返回新 `conversation_id`、最终回答和 `done`。 +- Given 后续合法消息,When 传入同一 `conversation_id`,Then 复用原 SuperAgent Session。 +- Given 同一会话已有活动 Run,When 再次发送,Then 返回 `409` 且不发第二个上游请求。 +- Given Provider 流不完整或 Run 失败,When 处理结束,Then 不发送 `message` 或 `done`,只返回安全错误。 +- Given 服务重启、会话过期或随机 ID 不存在,When 继续对话,Then 返回 `404` 而不把客户端值当作 Provider Session。 +- Given 无真实网络和 Secret,When 执行自动化测试,Then 使用本地 fake/模拟 Provider 并全部通过。 + +## 10. Definition of Done + +- 配置、Service、SuperAgent 适配、HTTP Handler 和应用装配满足本文契约。 +- 覆盖鉴权、CORS、输入、SSE、复用、并发、过期、容量、上游错误和敏感信息边界测试。 +- `gofmt`、`go test ./...`、`go test -race ./...` 和 `go vet ./...` 通过。 +- `.env.example`、集成指南、安全边界、项目索引和 `PROJECT_STATE.md` 同步。 +- 不提交真实 Secret,不自动创建 Git 提交。 diff --git a/docs/specs/fire-safety-ymd-dashscope-compatible-chat-v1.md b/docs/specs/fire-safety-ymd-dashscope-compatible-chat-v1.md new file mode 100644 index 0000000..2dffc2d --- /dev/null +++ b/docs/specs/fire-safety-ymd-dashscope-compatible-chat-v1.md @@ -0,0 +1,143 @@ +# DashScope 风格兼容对话 API v1 Spec + +| 项 | 内容 | +| --- | --- | +| 状态 | Implemented | +| 日期 | 2026-09-05 | +| 负责人 | fire-safety-ymd 后端 | +| 需求来源 | 复用既有客户端的 `/api/v1/apps/{app_id}/completion`、`xtoken` 与 `event: result` 契约 | + +## 1. 背景 + +既有用户端按照 DashScope Application HTTP 形式调用固定应用路径,并以 SSE `result` 事件读取 `output.session_id`、`finish_reason` 和 `text`。本项目原生 `/api/chat` 的请求体与事件名称不同。为避免客户端一次性重写,本 checkpoint 在同一 Chat Service 之上增加一个受限兼容适配层,并保留原生接口。 + +本契约参考阿里云官方 Application API 的路径、`input.prompt`、`input.session_id` 和 SSE 结果外形,但不是 DashScope 全量代理或完整复刻。 + +## 2. 目标 + +- 新增可选的 `POST /api/v1/apps/{app_id}/completion`。 +- 接收既有客户端的 `xtoken`、`input.prompt`、可选 `input.session_id` 和 `parameters` 空对象。 +- 返回 `event: result` SSE;首个事件提供会话 ID,严格成功事件以 `finish_reason: "stop"` 终止。 +- 复用 `/api/chat` 的会话容量、TTL、同会话并发排斥、SuperAgent 严格完成判定和错误分类。 +- Nginx 只终止 TLS、限制路径、限流并反代本 Go 服务,不保存或注入任何应用 Secret。 + +## 3. 非目标与兼容边界 + +- 不移除或改变原生 `POST /api/chat`。 +- 不实现 DashScope 的 `messages`、`biz_params`、文件、图像、知识库、思考过程或任意 `parameters`。 +- 不代理到 DashScope,也不把兼容 `app_id` 当成真正的 SuperAgent Profile ID。 +- 不向客户端暴露 Provider Session、Run、Trace、工具参数/结果或半截回答。 +- 不提供非流式 JSON 模式;即使请求未带 `X-DashScope-SSE: enable`,本兼容路径也始终返回 SSE,以匹配现有客户端。 +- `parameters.incremental_output` 仅为输入兼容而接受。v1 只有严格完成后的一个正文结果,因此 `true` 与 `false` 不改变正文分片方式。 +- 静态 `xtoken` 只用于受控联调,不是最终用户身份或动态授权。 + +## 4. 配置与路由 + +| 环境变量 | 默认值 | 说明 | +| --- | --- | --- | +| `FIRE_SAFETY_CHAT_COMPAT_APP_ID` | 空 | 非空时注册兼容路径;1 至 128 个 ASCII 字母、数字、下划线或连字符 | +| `FIRE_SAFETY_CHAT_AUTH_TOKEN` | 空 | 兼容路径期望的 `xtoken`,仍必须与 SuperAgent Key、MCP Token 不同 | + +兼容路径只有在 `FIRE_SAFETY_CHAT_ENABLED=true` 且 App ID 非空时注册。App ID 是用于路径匹配的公开标识,不是 Secret;不匹配的路径返回 404。 + +## 5. 请求契约 + +```http +POST /api/v1/apps/fire-safety-public-app/completion +xtoken: +Content-Type: application/json +Accept: text/event-stream + +{ + "input": { + "prompt": "杨家盘瞭望哨 3 公里内的水源?" + }, + "parameters": {} +} +``` + +后续轮次把前一轮成功响应的 `output.session_id` 放回 `input`: + +```json +{ + "input": { + "prompt": "再说明这些候选的限制", + "session_id": "conv_..." + }, + "parameters": { + "incremental_output": true + }, + "debug": {} +} +``` + +约束: + +- `input.prompt` 必填,非空,并服从现有消息大小上限。 +- `input.session_id` 只能是本服务此前返回的本地会话 ID;服务端仍负责映射 Provider Session。 +- `parameters` 可省略、为空对象,或只包含布尔型 `incremental_output`。 +- `debug` 可省略或只能是空对象。 +- 未知字段、`null`、错误类型、尾随 JSON 和超大请求体均拒绝。 +- 兼容 DTO 不接收用户身份、角色、区域、租户、Provider Session 或 metadata。 + +## 6. 成功 SSE + +Prepare 成功后先发送会话事件。`finish_reason` 的初始值按官方示例使用字符串 `"null"`,不是 JSON `null`: + +```text +id: 1 +event: result +:HTTP_STATUS/200 +data: {"output":{"session_id":"conv_...","finish_reason":"null"},"usage":{},"request_id":"..."} +``` + +只有 SuperAgent Adapter 同时验证最终内容、成功 `run.completed` 和顶层 `end` 后,才发送终止事件: + +```text +id: 2 +event: result +:HTTP_STATUS/200 +data: {"output":{"session_id":"conv_...","finish_reason":"stop","text":"..."},"usage":{"models":[{"input_tokens":12,"output_tokens":34,"model_id":"qwen-plus-latest"}]},"request_id":"..."} +``` + +- 两个事件的 `session_id` 和 `request_id` 必须相同。 +- `session_id` 是本地随机会话 ID,不是 Provider Session。 +- `model_id` 在上游未提供有效模型名时省略;token 数量是当前 Provider 返回的单轮聚合值。 +- 等待期间每 15 秒可发送不含业务数据的 SSE comment heartbeat。 +- 客户端只有看到 `event: result` 且 `output.finish_reason == "stop"` 才能把本轮视为成功。 + +## 7. 错误 + +SSE 开始前使用 HTTP 状态和 DashScope 风格安全 JSON: + +```json +{"code":"CHAT_AUTH_INVALID","message":"Chat authentication failed.","request_id":"..."} +``` + +SSE 开始后的失败发送 `event: error`,只包含稳定 code、通用 message、request ID 和本地 session ID;不发送 `finish_reason: "stop"`,也不返回任何已接收的部分回答。 + +沿用原生 Chat 的主要状态:400 输入错误、401 token 错误、403 Origin 错误、404 App/会话不存在、409 会话忙、503 容量不足,以及 502/504 上游失败或超时。 + +## 8. CORS、Nginx 与 Secret + +- 浏览器 Origin 必须精确出现在 `FIRE_SAFETY_CHAT_ALLOWED_ORIGINS`;不允许 `*` 或 credentials。 +- 预检只允许 `POST` 以及 `Content-Type`、`xtoken`、`X-DashScope-SSE`、`X-Request-ID`。 +- Nginx 示例只公开精确兼容路径、`/mcp` 和可选 `/health`,其余路径返回 404。 +- Nginx 不比较或注入 `xtoken`、SuperAgent Open API Key、MCP Bearer 或数据库凭证;Header 原样交给 Go 验证。 +- 对话 SSE 必须关闭代理缓冲、缓存、gzip 和上游自动重试,并让代理超时覆盖 Chat 总运行时限。 + +## 9. 验收标准 + +- Given App ID 未配置,When 请求兼容路径,Then 路由返回 404,原生 `/api/chat` 行为不变。 +- Given App ID 或 `xtoken` 错误,When 请求,Then 在读取问题和调用 SuperAgent 前拒绝。 +- Given 首轮请求成功,When 读取 SSE,Then 先收到字符串 `"null"` 会话事件,再收到同会话 `"stop"` 最终正文。 +- Given 后续请求携带成功返回的 `session_id`,When 调用,Then 复用同一服务端会话映射。 +- Given Provider 流失败,When SSE 已开始,Then 收到安全 `error`,不收到部分正文或 `stop`。 +- Given Nginx 配置生效,When 请求未列出的路径,Then 不会转发到 Go 或外部 DashScope。 + +## 10. 协议来源 + +- [Alibaba Cloud Model Studio Application API reference](https://www.alibabacloud.com/help/en/model-studio/application-api-reference) +- [Call a Model Studio application by using an API](https://www.alibabacloud.com/help/en/model-studio/application-calling-guide) + +外部文档会演进;本项目以本 Spec 和自动化测试定义的受限兼容子集为准。 diff --git a/docs/specs/fire-safety-ymd-superagent-mcp-spatial-readonly-v1.md b/docs/specs/fire-safety-ymd-superagent-mcp-spatial-readonly-v1.md new file mode 100644 index 0000000..f8c933b --- /dev/null +++ b/docs/specs/fire-safety-ymd-superagent-mcp-spatial-readonly-v1.md @@ -0,0 +1,210 @@ +# 森林防火空间只读 MCP v1 Spec + +| 项 | 内容 | +| --- | --- | +| 状态 | Implemented;live readiness 与本地 7 tools/call 已通过,SuperAgent 回调待联调 | +| 日期 | 2026-09-05 | +| Checkpoint | `fire-safety-ymd-superagent-mcp-spatial-readonly-v1` | +| 需求来源 | 用户提供 8 张 PostgreSQL/PostGIS 表结构与每表 2 条样例,要求先基于现有数据建设 MCP | + +## 1. 背景与已核对输入 + +`docs/import/db-samples/` 中包含以下表的结构和样例: + +- `st_2_mpslfh_t_slfh_syd`:水源地,点。 +- `st_2_xianyouxushuichiguan`:蓄水池,点。 +- `st_2_xianyoufanghuotongdao`:防火通道,线。 +- `st_2_fanghuojianchazhan`:防火检查站,点。 +- `st_2_fanghuoliaowangshao`:防火瞭望哨,点。 +- `st_2_fanghuowangge`:防火网格,面。 +- `st_2_linqugongkuangqiye`:林区工矿企业,面。 +- `st_2_mudifenqu_mian`:墓地坟区,面。 + +SQL 文件仅作为结构和数据契约参考。文件包含 `DROP TABLE` 等导出语句,本项目不会执行这些文件。 + +当前样例足以确定首版字段映射,但不能证明全库的数据完整性、坐标系、几何类型一致性、索引、时效或资源当前可用性。尤其是: + +- 所有 `geom` 列均声明为宽泛的 `geometry(GEOMETRY)`,没有 typmod SRID。 +- 现场全量导入进一步确认:`st_2_xianyoufanghuotongdao.geom` 的源记录混合二维与 Z 维度,原 `geometry(GEOMETRY)` 二维 typmod 会拒绝 Z;用户改用裸 `geometry` 保留原始维度后报告重导成功。该导入反馈仍需只读 probe 核验最终数量和分布。 +- 样例几何十六进制是未携带 SRID 的 WKB;坐标数值及水源表独立经纬度字段看起来符合 WGS84,但这只能作为待验证线索。 +- 样例 DDL 只有防火网格明确包含 GiST 几何索引。 +- 多张表含姓名、电话等受限字段;首版 MCP 不返回这些字段。 + +## 2. 目标 + +- 在现有 Go HTTP 服务内提供默认关闭的 `POST /mcp`。 +- 与已经验证的 SuperAgent 配置保持兼容,采用 MCP `2025-06-18`、JSON-RPC 2.0 和单请求 JSON 响应。 +- 通过独立 Bearer Token 鉴权;不得复用 SuperAgent Open API Key。 +- 只提供固定、参数化、只读、按服务端可信数据范围执行的空间查询工具;默认使用镇街白名单,数据库全范围必须显式启用。 +- 支持在现有业务记录中按地名搜索有界候选,让没有坐标的用户先确认一个候选,再进入距离和包含关系查询。 +- 使用 PostgreSQL/PostGIS 作为事实来源,并在启用 MCP 前校验 PostGIS、表、SRID、几何类型和有效性。 +- 工具结果明确表达数据来源、生成时间、限制与需要现场确认的事项。 +- 提供一个只输出 schema/几何统计、不输出业务记录或联系人数据的 PostGIS readiness probe。 + +## 3. 非目标 + +- 不开放任意 SQL、任意表查询、文件读取或通用 HTTP 代理。 +- 不写数据库,不调度人员或资源,不生成/下发指挥命令。 +- 不声称防火通道是可通行路线,不实现路网拓扑、最短路、坡度或实时封路计算。 +- 不声称网格队伍的实时位置或集结点;现有表只提供责任队伍和区域信息。 +- 不依据资源记录存在就宣称水源、检查站或瞭望哨当前可用。 +- 不返回负责人、队长、值班人员、书记、联系电话或图片地址。 +- 不提供任意地址、山名或道路的外部地理编码,不实现完整地名库、拼音/别名纠错,也不自动选定地名候选。 +- 不在本 checkpoint 内建立最终用户身份、动态角色或多租户模型;首版使用每环境 MCP 凭证和服务端静态数据库全范围/镇街白名单。 + +## 4. MCP 传输与安全契约 + +### 4.1 支持的方法 + +- `initialize` +- `notifications/initialized` +- `tools/list` +- `tools/call` + +服务端声明协议版本 `2025-06-18` 和 `tools` capability。首版是无会话、非流式的 Streamable HTTP 子集:`POST` 返回 `application/json`;`GET /mcp` 返回 `405 Method Not Allowed`。 + +### 4.2 入站控制 + +- MCP 默认关闭,关闭时不注册 `/mcp`。 +- 只接受 `Content-Type: application/json`。 +- 使用 `Authorization: Bearer `,常量时间比较。 +- Token 至少 32 个可打印 ASCII 字符,且不能等于 SuperAgent Open API Key。 +- 请求体默认最大 256 KiB,最大可配置 1 MiB。 +- 首版只接受服务到服务请求;带非空 `Origin` 的请求拒绝,避免浏览器和 DNS rebinding 风险。 +- 不接受模型提供的用户、角色、租户或授权范围。范围模式和可选镇街列表只来自服务端配置。 +- 每次工具调用有硬超时;日志只记录 request ID、方法/工具名、耗时和结果类别,不记录参数、坐标、Token、SQL 或结果正文。 + +### 4.3 服务端数据范围 + +- `FIRE_SAFETY_MCP_SCOPE_MODE=town_allowlist` 是默认值;启用 MCP 时要求 `FIRE_SAFETY_MCP_ALLOWED_TOWNS` 非空,并按数据库镇街字段精确过滤。 +- `FIRE_SAFETY_MCP_SCOPE_MODE=all` 必须显式配置;它授权查询当前配置数据库中 MCP 固定查询表的全部镇街以及镇街字段为空的记录,此时 `FIRE_SAFETY_MCP_ALLOWED_TOWNS` 必须为空。 +- `all` 不会放宽只读、固定 SQL、空间半径、结果数量、字段脱敏、SRID 或 readiness 限制。 +- 业务表中的镇街值不能自动决定授权范围;范围只能由可信服务端配置建立。 + +## 5. MCP 工具 + +除地名候选工具外,所有工具的坐标参数使用 WGS84:`longitude` 范围 `[-180, 180]`,`latitude` 范围 `[-90, 90]`。需要距离的工具接收 `radius_meters` 和 `limit`,同时有工具级默认值和硬上限。 + +### 5.1 `fire_safety_search_place_candidates` + +输入 `place_name`(去除首尾空白后 2 至 100 字符)和可选 `limit`(默认 10、最大 20),在 8 张现有业务表的名称、镇街和村庄字段中做不区分大小写的文字包含匹配。百分号和下划线按普通文字处理,不作为 SQL 通配符。 + +返回稳定资源类型、记录 ID、名称、镇街、村庄、命中字段、`exact | partial`、WGS84 坐标和 `location_kind`: + +- 点记录返回 `recorded_point`。 +- 线记录返回首个组成线的起点,面记录返回 `ST_PointOnSurface` 代表点,两者标记为 `representative_point`。 + +候选不等于用户位置。无结果时 Agent 必须请用户补充名称或地图选点;多结果时必须请用户选择;即使只有一个候选,也要先回显确认,不能直接调用后续距离工具。该工具不返回联系人字段,不调用外部地图服务。 + +### 5.2 `fire_safety_resolve_incident_context` + +输入演练点坐标,查询覆盖该点的防火网格,返回网格 ID、镇街、区域标签。用于建立后续查询上下文,不返回个人信息。 + +### 5.3 `fire_safety_find_nearby_water_sources` + +合并水源地和蓄水池,按球面距离返回候选水源。返回名称/标签、类型、镇街、村、坐标、距离、容量(数据存在时)、源表报告状态和未经解释的源时间原值(数据存在时)。结果固定标记“当前可用性未验证”;时间原值的单位和时区待数据所有者确认。 + +### 5.4 `fire_safety_find_command_post_candidates` + +合并防火检查站和防火瞭望哨,返回附近候选设施、坐标、距离及源表报告状态。结果只代表空间候选,必须由现场核验安全、通信、容量、可达性和火势上风向等条件。 + +### 5.5 `fire_safety_list_nearby_access_lines` + +返回附近防火通道名称、按 WGS84 geometry 计算的米制长度、最近接入点和未经解释的源更新时间原值。该工具不称为 route planner;它不返回“推荐路线”或“可通行”结论。 + +### 5.6 `fire_safety_get_responsible_units` + +根据坐标查询覆盖网格及其防火中队名称。明确返回 `live_location_available=false` 和 `assembly_site_available=false`,因为当前表没有队伍实时位置、战备状态或正式集结点字段。 + +### 5.7 `fire_safety_find_nearby_risk_areas` + +合并墓地坟区和林区工矿企业,返回点位周边风险区域类别、名称/标签、镇街、村、方位、是否覆盖演练点和距离。联系人字段不返回。 + +## 6. 统一结果语义 + +工具成功结果使用: + +```json +{ + "status": "ok | no_results", + "data": {}, + "metadata": { + "generated_at": "RFC3339 UTC", + "data_sources": [], + "spatial_reference": "EPSG:4326", + "result_count": 0 + }, + "warnings": [] +} +``` + +`structuredContent` 保存上述对象,同时在 `content[0].text` 返回同一对象的 JSON 文本,兼容只读取文本内容的 MCP 客户端。 + +`data_sources` 只使用稳定领域名称(如 `water_source`、`fire_grid`),不向模型暴露历史物理表名。 + +`warnings` 使用 `results_limited_to_server_authorized_towns` 表示镇街白名单模式,使用 `results_include_all_towns_in_configured_database` 表示显式数据库全范围模式,避免 Agent 混淆结果范围。所有工具同时返回 `source_records_with_invalid_geometries_are_excluded`,防止 Agent 把排除后的无结果解释成数据库确认不存在相关记录。 + +工具级失败仍使用 JSON-RPC 成功响应中的 `isError=true`,结构化错误只公开稳定错误码: + +- `INVALID_ARGUMENT` +- `TOOL_NOT_FOUND` +- `DATA_SOURCE_UNAVAILABLE` +- `QUERY_TIMEOUT` +- `INTERNAL_ERROR` + +数据库 DSN、SQL、表结构细节和原始驱动错误不得进入 MCP 响应。 + +## 7. PostgreSQL/PostGIS 契约 + +- 使用 `pgx/v5` 原生连接池;当前仅面向 PostgreSQL,不引入 ORM。 +- 连接设置 `default_transaction_read_only=on`、`statement_timeout` 和应用名。 +- SQL 固定在 Repository,所有值使用 `$n` 参数;只有代码内表白名单可参与标识符拼接。 +- 查询始终应用服务端范围、半径和返回数量上限;`town_allowlist` 使用参数化镇街数组,`all` 使用服务端布尔参数显式跳过镇街条件,模型不能提供该参数。 +- 地名查询使用参数化的文字包含匹配,先按精确名称、精确村庄、精确镇街,再按前缀和普通包含排序;不拼接用户输入。线面候选的二维代表点只在只读查询中计算,不修改原始几何。 +- MCP v1 只支持已实库确认的 EPSG:4326。启用 MCP 时 `FIRE_SAFETY_POSTGIS_EXPECTED_SRID` 必须显式配置为 `4326`。 +- readiness 校验要求每张已用表存在、PostGIS 可用、非空几何的 SRID 均为 4326、类型符合点/线/面预期且坐标不越过 WGS84 范围;表/SRID/类型/越界不符合时服务拒绝启用 MCP。 +- 无效几何、空几何、空表或缺少空间索引作为 readiness warning。固定查询使用 `ST_IsValid` 排除无效几何,不自动调用 `ST_MakeValid`,不修改原始事实;生产前应评估数据缺口并补齐查询表达式适用的 GiST 索引。 +- 原始防火通道列保持裸 `geometry` 以容纳混合二维/Z。不得用 `ST_Force2D` 或重写 WKB 修改原始数据;确需二维投影时仅在只读查询/派生层显式执行并测试其结果语义。 +- 不自动执行样例 SQL、SRID 修复或索引 DDL。 + +## 8. 配置 + +| 环境变量 | 默认值 | 启用时要求 | +| --- | --- | --- | +| `FIRE_SAFETY_MCP_ENABLED` | `false` | 显式为 `true` | +| `FIRE_SAFETY_MCP_AUTH_TOKEN` | 空 | 必填,独立高熵 Token | +| `FIRE_SAFETY_MCP_SCOPE_MODE` | `town_allowlist` | `town_allowlist` 或 `all` | +| `FIRE_SAFETY_MCP_ALLOWED_TOWNS` | 空 | `town_allowlist` 时必填;`all` 时必须为空 | +| `FIRE_SAFETY_MCP_MAX_BODY_BYTES` | `262144` | `1..1048576` | +| `FIRE_SAFETY_MCP_TOOL_TIMEOUT` | `5s` | `>0` 且不超过 `30s` | +| `FIRE_SAFETY_POSTGIS_ENABLED` | `false` | MCP 启用时必须为 `true` | +| `FIRE_SAFETY_POSTGIS_DSN` | 空 | PostGIS 启用时必填;Secret | +| `FIRE_SAFETY_POSTGIS_EXPECTED_SRID` | 空/`0` | MCP 启用时必须显式为 `4326` | +| `FIRE_SAFETY_POSTGIS_CONNECT_TIMEOUT` | `5s` | `>0` 且不超过 `30s` | +| `FIRE_SAFETY_POSTGIS_QUERY_TIMEOUT` | `3s` | `>0` 且不超过 `30s` | +| `FIRE_SAFETY_POSTGIS_MAX_CONNS` | `4` | `1..20` | + +## 9. 验收标准 + +- MCP 关闭时 `/mcp` 不暴露,健康检查保持可用。 +- MCP 启用但 Token、合法范围、PostGIS 或显式 SRID 缺失时配置加载失败;`all` 与非空镇街列表同时出现时失败,错误不含 Secret。 +- 无 Token、错误 Token、错误 Content-Type、非空 Origin、超大 body 和非法 JSON 均被稳定拒绝。 +- `initialize`、initialized notification、`tools/list` 和 7 个 `tools/call` 契约通过本地测试。 +- 工具 schema 限制地名长度、经纬度、半径、数量和未知字段;工具不接受授权身份参数。 +- Service 测试证明可信数据库全范围/镇街白名单来自构造时配置,并覆盖互斥校验、无结果、超时和仓储失败。 +- Repository 只使用参数化值,MCP 结果不包含联系人字段。 +- readiness 测试证明无效几何不阻塞启动、会产生排除 warning,且全部固定查询显式包含 `ST_IsValid`;SRID、类型和越界仍为硬失败。 +- 地名 Service/Repository 测试覆盖首尾空白、长度、控制字符、结果上限、服务端范围、8 张固定表和候选确认 warning。 +- readiness probe 不输出业务行、电话、负责人、DSN 或 SQL。 +- `gofmt`、`go test -count=1 ./...`、`go test -race -count=1 ./...` 和 `go vet ./...` 通过。 + +## 10. 上线前未确认项 + +- 实库源 CRS 已由数据提供方确认为 EPSG:4326 且无偏移;2026-09-05 已通过单独审核、显式确认和单事务迁移为 4,048 条非空几何补齐 SRID 4326,严格 readiness 已通过。重新导入无 SRID 的原始 SQL 时仍必须重新经过该流程,应用查询不得静默赋值。 +- 实库 35 条无效面几何按首版决策排除;需评估由此造成的网格、责任单位和风险区域覆盖缺口是否满足生产要求。 +- 各资源状态字段的枚举、更新时间含义和数据刷新责任人。 +- MCP 回调网络地址、TLS、SuperAgent 实际 Header 行为及 Token 轮换方式。 +- 最终用户身份、区域权限、精确位置权限和审计保留策略。 +- 真实数据量下地名重复率、字段质量和无索引包含搜索性能;后续是否引入标准地名表、别名词典、`pg_trgm` 或经审批的外部地理编码服务。 +- 路线规划需要的路网拓扑、路面/宽度/坡度/车辆限制、实时封路、火场和天气数据。 +- 队伍集结需要的正式集结点、实时位置、战备状态、装备和容量数据。 diff --git a/docs/specs/fire-safety-ymd-superagent-openapi-connectivity.md b/docs/specs/fire-safety-ymd-superagent-openapi-connectivity.md new file mode 100644 index 0000000..da029c7 --- /dev/null +++ b/docs/specs/fire-safety-ymd-superagent-openapi-connectivity.md @@ -0,0 +1,159 @@ +# SuperAgent Open API 连通性基线 Spec + +| 项 | 内容 | +| --- | --- | +| 状态 | Implemented | +| 日期 | 2026-09-04 | +| 负责人 | fire-safety-ymd 后端 | +| 需求来源 | 用户 checkpoint `fire-safety-ymd-superagent-openapi-connectivity` | +| 关联 Change Request | 无 | + +## 1. 背景 + +fire-safety-ymd 需要由 Go 后端调用既有 SuperAgent 平台,并在后续让 SuperAgent 通过本项目 MCP 工具查询 PostgreSQL/PostGIS 消防数据。 + +TH Hotel 项目已经验证了“创建 Open Agent Session、流式发送消息、解析公开 Trace、严格判断完成状态和断流恢复”的调用形态。本 checkpoint 只迁移协议经验和安全边界,使用 Go 独立实现,不复制 Java、酒店业务、AgentBus、邮件、OSS 或任务结果写入逻辑。 + +## 2. 目标 + +- 建立独立、可测试的 Go SuperAgent Open API Adapter。 +- 分离 `CreateSession` 与 `StreamMessage`,为后续一个本地对话复用一个 SuperAgent Session 做准备。 +- 使用严格 SSE 成功条件,拒绝部分回答和不完整协议结果。 +- 初始 SSE 断流后通过既有 Run 恢复,不重新发送原始消息。 +- 提供默认关闭的 CLI 连通性探针,只发送固定无敏感信息消息。 +- 所有自动化测试使用本地模拟 Provider,不需要真实 Secret 或网络。 + +## 3. 非目标 + +- 不实现浏览器或业务聊天 API。 +- 不持久化本地会话与 SuperAgent Session 的映射。 +- 不实现 `/mcp` endpoint 或消防工具。 +- 不连接 PostgreSQL/PostGIS。 +- 不创建或配置 SuperAgent Profile、外部应用、API Key 或 MCP Server。 +- 不把 SuperAgent 输出写入消防业务事实。 + +## 4. 用户与场景 + +当前用户是开发和联调人员:在安全配置测试环境后,通过 CLI 探针验证 Go 服务能创建 Session、发送无敏感信息消息并取得严格完成的最终回答。 + +后续的用户对话 API v1 已在独立 Spec 中实现对该 Adapter 的复用、单进程会话映射、并发冲突和客户端 SSE;真实身份、持久化与主动取消仍需另行设计。 + +## 5. Definition of Ready + +- 目标与非目标已确认:是。 +- 协议基线已确认:以 TH Hotel 仓库保存的 2026-07-12 SuperAgent Open API 文档和当前实现为输入,真实环境上线前重新验证。 +- 权限和安全边界已确认:Secret 仅由本项目环境变量注入;没有用户业务数据进入探针。 +- 外部依赖已确认:实现只使用 Go 标准库。 +- 未确认问题已列出:Profile、外部应用、scope、真实 Base URL 和 API Key 均由平台管理员后续提供。 + +## 6. 协议与客户端契约 + +### 6.1 创建 Session + +```text +POST /api/open/agent-sessions +``` + +请求包含: + +- `external_subject_id`:调用方提供的稳定、非敏感主体标识。 +- `idempotency_key`:创建 Session 的稳定幂等键。 +- `metadata`:不包含 Secret 的关联元数据。 + +客户端接受响应中的 `session_id`,并兼容 `id` 字段。 + +### 6.2 流式发送消息 + +```text +POST /api/open/agent-sessions/{session_id}/messages/stream?include_trace=true +``` + +请求包含 `message`、稳定 `idempotency_key` 和安全 `metadata`。Session ID 由上层明确传入,Adapter 不在每条消息前隐式创建新 Session。 + +### 6.3 鉴权与关联 Header + +- `Authorization: Bearer ` +- `X-Request-ID: <稳定请求关联 ID>` +- `X-CSRF-Token` 与 `Cookie: csrf_token=...` 使用同一个每请求随机值。 +- SSE 使用 `Accept: text/event-stream` 和 `Cache-Control: no-cache`。 + +## 7. SSE 成功与恢复规则 + +一次调用只有同时满足以下条件才成功: + +1. 收到最终内容;优先使用 `message.final`,兼容累计 `message.delta` 和历史 `messages`/`values` AI 消息。 +2. 收到 `run.completed` 且 `status=success`。 +3. 收到顶层 `event: end`。 +4. 没有收到顶层 `error` 或 `run.failed`。 + +解析器必须支持 `event:`、多行 `data:`、`id:`、心跳注释和空行分帧,并按 SSE event ID 去重。 + +初始流提前结束时: + +- 不重新 POST 消息。 +- 从 `Content-Location` 获取同源 Run URL;必要时使用已解析的 Run ID 构造 URL。 +- 查询 Run 状态,再订阅 `{run_url}/events`。 +- 恢复请求携带 `Last-Event-ID`。 +- 使用有上限的指数退避并服从调用方 context 取消。 +- 恢复耗尽或进入失败终态时返回明确错误,不返回部分答案。 + +## 8. 配置契约 + +| 环境变量 | 默认值 | 说明 | +| --- | --- | --- | +| `FIRE_SAFETY_SUPERAGENT_ENABLED` | `false` | 总开关,默认不调用真实 Provider | +| `FIRE_SAFETY_SUPERAGENT_BASE_URL` | 空 | SuperAgent Open API Base URL | +| `FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY` | 空 | Secret,启用时必填 | +| `FIRE_SAFETY_SUPERAGENT_CONNECT_TIMEOUT` | `15s` | 建连超时 | +| `FIRE_SAFETY_SUPERAGENT_RECOVERY_MAX_ATTEMPTS` | `5` | SSE 恢复最大次数;最大可配置为 20 | +| `FIRE_SAFETY_SUPERAGENT_RECOVERY_INITIAL_BACKOFF` | `250ms` | 首次恢复退避 | +| `FIRE_SAFETY_SUPERAGENT_MAX_MESSAGE_BYTES` | `65536` | 单条消息 UTF-8 字节上限;最大可配置为 16 MiB | +| `FIRE_SAFETY_SUPERAGENT_PROBE_SUBJECT_ID` | 固定探针主体 | 仅 CLI 探针使用 | +| `FIRE_SAFETY_SUPERAGENT_PROBE_TIMEOUT` | `10m` | 探针总超时 | + +启用时配置缺失或不合法必须启动失败;错误不得包含 API Key。 + +## 9. 安全规则 + +- 不读取 TH Hotel 的环境变量或 Secret,不共享 API Key。 +- Base URL 必须是绝对 HTTP/HTTPS URL,且不得带 userinfo、query 或 fragment。 +- `Content-Location` 和恢复 URL 必须与配置 Base URL 同源,防止向其他主机转发 Authorization。 +- 错误只暴露安全状态和经过限制的 Provider error code,不返回响应正文。 +- 客户端不记录消息、Header、Cookie、API Key 或 SSE 原始数据。 +- Probe 使用固定安全消息;不得用它发送真实火情、联系人或生产数据。 + +## 10. 需求追踪表 + +| 需求项 | 后端状态 | 测试状态 | 文档位置 | 当前状态 | +| --- | --- | --- | --- | --- | +| 环境配置与安全默认值 | Done | Passed | 本文第 8 节 | Implemented | +| CreateSession | Done | Passed | 本文第 6.1 节 | Implemented | +| StreamMessage 与严格成功条件 | Done | Passed | 本文第 6.2、7 节 | Implemented | +| SSE 断流恢复 | Done | Passed | 本文第 7 节 | Implemented | +| CLI 探针 | Done | Compile passed;live test pending | 项目集成指南 | Implemented | + +## 11. 验收标准 + +- Given 未启用 SuperAgent,When 执行真实调用,Then 返回受控禁用错误且不发起网络请求。 +- Given 配置完整,When 创建 Session,Then请求具备鉴权、幂等、关联和 CSRF Header,且能解析 Session ID。 +- Given完整 SSE,When 同时收到最终内容、成功完成和 `end`,Then 返回最终回答与安全元数据。 +- Given SSE 缺少任一成功条件,When 无法恢复,Then 返回协议错误且不返回部分回答。 +- Given 初始 SSE 提前结束且存在 Run URL,When 恢复,Then只 GET Run/events、携带 `Last-Event-ID`,消息 POST 次数仍为 1。 +- Given未设置真实环境变量,When 运行全部测试,Then 不访问外网且测试通过。 + +## 12. 测试范围 + +- 配置默认值、合法值和错误值。 +- Session 请求路径、Header、CSRF、请求体和响应解析。 +- 当前 Trace SSE、历史 messages/values 兼容、事件 ID 去重和多行 data。 +- 缺少 final/completed/end、顶层 error、run.failed 和非法 JSON。 +- 提前 EOF 恢复、Last-Event-ID、同源 URL 和不重复 POST。 +- HTTP 非 2xx、超大控制响应和 context 取消。 + +## 13. Definition of Done + +- 实现满足本文契约。 +- `gofmt`、`go test ./...` 和 `go vet ./...` 通过。 +- 没有真实 Secret、用户数据、构建产物或无关用户变更。 +- 项目索引、集成指南、安全边界和 `PROJECT_STATE.md` 已同步。 +- 真实环境未配置时明确说明未做 live connectivity test。 diff --git a/docs/workflows/exercise-plan-evidence.md b/docs/workflows/exercise-plan-evidence.md new file mode 100644 index 0000000..5df2c82 --- /dev/null +++ b/docs/workflows/exercise-plan-evidence.md @@ -0,0 +1,59 @@ +# 演练方案证据查询流程 + +## 目标 + +本流程说明 SuperAgent 如何从用户确认的地名候选或演练点坐标取得结构化事实,再组织“候选方案”。MCP 返回证据和限制,不自动选择地点或作出现场指挥决定。 + +```mermaid +sequenceDiagram + participant U as 用户 + participant SA as SuperAgent + participant MCP as fire-safety-ymd MCP + participant PG as PostGIS + + U->>SA: 提供地名或演练点坐标与分析目标 + alt 只提供地名 + SA->>MCP: search_place_candidates + MCP->>PG: 在获授权业务记录中做有界名称匹配 + PG-->>MCP: 记录点/代表点候选 + MCP-->>SA: 候选 + 匹配类型 + location_kind + warning + SA-->>U: 展示候选并请求确认 + U->>SA: 确认一个候选坐标或改为地图选点 + else 已提供明确坐标 + SA->>SA: 校验 WGS84 坐标格式 + end + SA->>MCP: resolve_incident_context + MCP->>PG: 查询覆盖网格(可信数据库全范围或镇街白名单内) + PG-->>MCP: 网格事实 + MCP-->>SA: 网格 + 数据限制 + par 候选资源 + SA->>MCP: find_nearby_water_sources + SA->>MCP: find_command_post_candidates + SA->>MCP: list_nearby_access_lines + and 责任与风险 + SA->>MCP: get_responsible_units + SA->>MCP: find_nearby_risk_areas + end + MCP-->>SA: 结构化候选、来源、距离、warning + SA-->>U: 区分数据库事实、未知项和需现场确认的候选方案 +``` + +## 对原始四类问题的支持度 + +| 问题 | 当前输出 | 结论边界 | +| --- | --- | --- | +| 指挥部设置到哪个 | 检查站/瞭望哨空间候选 | 不能自动定点;需核验安全、上风向、通信、容量和可达性 | +| 水源地在哪里 | 水源地/蓄水池坐标、距离、容量和源状态(有值时) | 不能宣称当前有水、可取水或道路可达 | +| 上山路线怎么规划 | 附近防火通道及最近接入点 | 当前不支持路线;缺少拓扑、坡度、宽度、路面、车辆限制、实时封路、火势和天气 | +| 各队伍在哪里集结 | 覆盖网格记录的责任中队名称 | 当前不支持实时位置或集结点;缺少正式集结点、队伍位置、战备和装备数据 | + +## Agent 回答要求 + +- 明确标注“数据库事实”“基于距离的候选”“当前数据不支持”和“必须现场确认”。 +- 只有地名时先调用地名候选工具;零结果时请用户补充名称或地图选点,多结果时列出候选,禁止静默选择第一条。 +- 所有空间工具都会排除无效几何;收到 `source_records_with_invalid_geometries_are_excluded` 时,零结果只能表述为“有效记录中未找到”,不能断言数据库不存在相关资源或责任区域。 +- `location_kind=representative_point` 只帮助识别线面记录,不得未经确认直接用于后续距离分析;记录点同样应回显名称和位置供用户确认。 +- 无结果时报告 `no_results`,不得自行补造附近资源。 +- 工具失败、超时或权限范围外时,不使用模型常识替代数据库事实。 +- 不把联系人、电话或内部物理表名转述给普通用户。 +- 涉及真实火情时,AI 结果只作辅助,不替代报警、撤离和现场指挥。 diff --git a/docs/workflows/user-chat.md b/docs/workflows/user-chat.md new file mode 100644 index 0000000..eb64ee5 --- /dev/null +++ b/docs/workflows/user-chat.md @@ -0,0 +1,83 @@ +# 用户对话 API 工作流 + +## 1. 首轮对话 + +客户端向 `POST /api/chat` 发送 `message`,不传 `conversation_id`。服务端创建 SuperAgent Session,随后通过 SSE 返回: + +1. `conversation`:包含新生成的 `conversation_id` 和 `reused=false`。 +2. 零个或多个 `progress`:只用于展示运行/工具进度。 +3. `message`:严格完成后的最终回答。 +4. `done`:包含同一 `conversation_id`、安全 Run ID 和 token usage。 + +客户端只有在收到 `done` 后才把本轮标记为成功,并保存 `conversation_id`。 + +## 2. 后续对话 + +后续请求同时发送 `message` 和前一轮保存的 `conversation_id`。成功流中的 `conversation` 事件返回 `reused=true`,说明复用了原 SuperAgent Session 上下文。 + +同一对话在收到 `done` 或终止 `error` 前不得再次提交。若服务端返回 `CHAT_CONVERSATION_BUSY`,客户端应保留当前流并稍后重试,不能自动改用同一个问题创建多轮并发 Run。 + +## 3. 失败处理 + +- HTTP JSON 错误:SSE 尚未开始;按 HTTP 状态和 `error.code` 处理。 +- SSE `error`:流已经开始,但本轮没有可信最终回答;不得把此前进度当作答案。 +- `CHAT_CONVERSATION_NOT_FOUND`:会话已过期、服务重启或请求落到其他实例;提示用户上下文已失效,并在用户确认后省略 `conversation_id` 创建新对话。 +- 上游超时、协议错误、Run 失败或客户端中途断开:当前实现会使会话映射失效,以免继续复用可能仍有活动 Run 的 Provider Session。 +- 网络断开且没有看到 `done`:结果未知;首版不自动重放原消息,避免重复 Run。 + +## 4. curl 联调 + +先把 `.env` 显式加载到当前 shell,再启动服务;Go 程序不会自动读取 `.env`: + +```bash +set -a +source .env +set +a +go run ./cmd/server +``` + +另开终端发起首轮请求: + +```bash +curl -N \ + -H "Authorization: Bearer ${FIRE_SAFETY_CHAT_AUTH_TOKEN}" \ + -H "Content-Type: application/json" \ + -H "Accept: text/event-stream" \ + --data '{"message":"观水镇附近有哪些水源候选?"}' \ + http://127.0.0.1:8080/api/chat +``` + +从 `conversation` 或 `done` 事件复制 ID 后测试下一轮: + +```bash +curl -N \ + -H "Authorization: Bearer ${FIRE_SAFETY_CHAT_AUTH_TOKEN}" \ + -H "Content-Type: application/json" \ + -H "Accept: text/event-stream" \ + --data '{"message":"再说明这些候选的限制","conversation_id":"conv_替换为上一轮返回值"}' \ + http://127.0.0.1:8080/api/chat +``` + +浏览器前端还必须把其精确 Origin 加入 `FIRE_SAFETY_CHAT_ALLOWED_ORIGINS`,例如 `http://localhost:5173`。静态 Chat Bearer 会被浏览器用户看到,因此只适用于受控联调,不能直接作为公网最终用户鉴权。 + +## 5. 既有 DashScope 风格客户端 + +设置 `FIRE_SAFETY_CHAT_COMPAT_APP_ID` 后,同一个 Chat Service 还会注册: + +```text +POST /api/v1/apps/{FIRE_SAFETY_CHAT_COMPAT_APP_ID}/completion +``` + +首轮请求使用 `xtoken: ${FIRE_SAFETY_CHAT_AUTH_TOKEN}`: + +```bash +curl -N \ + -H "xtoken: ${FIRE_SAFETY_CHAT_AUTH_TOKEN}" \ + -H "Content-Type: application/json" \ + --data '{"input":{"prompt":"杨家盘瞭望哨 3 公里内的水源?"},"parameters":{}}' \ + http://127.0.0.1:8080/api/v1/apps/fire-safety-public-app/completion +``` + +兼容流先返回 `finish_reason: "null"` 和本地 `session_id`,严格成功后返回 `finish_reason: "stop"` 与最终 `text`。后续轮次把该 `session_id` 放入 `input.session_id`。客户端不得使用 URL 中的 App ID、静态 token 或 session ID推断身份与权限,也不得把未出现 `stop` 的断流结果当作成功答案。 + +公网 Nginx 示例和完整 curl 见 [`../project/operations/nginx-public-entry.md`](../project/operations/nginx-public-entry.md)。兼容字段的唯一项目契约见 [`../specs/fire-safety-ymd-dashscope-compatible-chat-v1.md`](../specs/fire-safety-ymd-dashscope-compatible-chat-v1.md)。 diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..5b5dcbc --- /dev/null +++ b/go.mod @@ -0,0 +1,13 @@ +module fire-safety-ymd + +go 1.26.6 + +require github.com/jackc/pgx/v5 v5.10.0 + +require ( + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect + golang.org/x/sync v0.17.0 // indirect + golang.org/x/text v0.29.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..c0e505b --- /dev/null +++ b/go.sum @@ -0,0 +1,26 @@ +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/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0= +github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= +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/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug= +golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= +golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk= +golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/app/app.go b/internal/app/app.go new file mode 100644 index 0000000..3ba0b84 --- /dev/null +++ b/internal/app/app.go @@ -0,0 +1,184 @@ +// Package app assembles the application and owns process-level lifecycles. +package app + +import ( + "context" + "errors" + "fmt" + "log" + "net/http" + "time" + + "fire-safety-ymd/internal/config" + "fire-safety-ymd/internal/domain" + "fire-safety-ymd/internal/handler" + "fire-safety-ymd/internal/integration/superagent" + "fire-safety-ymd/internal/repository" + "fire-safety-ymd/internal/service" +) + +const ( + readHeaderTimeout = 5 * time.Second + readTimeout = 15 * time.Second + writeTimeout = 35 * time.Second + idleTimeout = 60 * time.Second + shutdownTimeout = 10 * time.Second + readinessTimeout = 30 * time.Second +) + +// Application owns the HTTP server and coordinates graceful shutdown. +type Application struct { + server *http.Server + close func() +} + +// New wires the application's current dependencies. +func New(ctx context.Context, cfg config.Config) (*Application, error) { + routerOptions := handler.RouterOptions{} + closeDependencies := func() {} + + if cfg.Chat.Enabled { + if !cfg.SuperAgent.Enabled { + return nil, fmt.Errorf("initialize chat: SuperAgent must be enabled") + } + superAgentClient, err := superagent.NewHTTPClient(superagent.Config{ + Enabled: cfg.SuperAgent.Enabled, + BaseURL: cfg.SuperAgent.BaseURL, + APIKey: cfg.SuperAgent.OpenAPIKey, + ConnectTimeout: cfg.SuperAgent.ConnectTimeout, + RecoveryMaxAttempts: cfg.SuperAgent.RecoveryMaxAttempts, + RecoveryInitialBackoff: cfg.SuperAgent.RecoveryInitialBackoff, + MaxMessageBytes: cfg.SuperAgent.MaxMessageBytes, + }) + if err != nil { + return nil, fmt.Errorf("initialize SuperAgent client: %w", err) + } + chatAgent, err := superagent.NewChatAgentAdapter(superAgentClient) + if err != nil { + return nil, fmt.Errorf("initialize SuperAgent chat adapter: %w", err) + } + chatService, err := service.NewChatService(chatAgent, service.ChatOptions{ + ExternalSubjectID: cfg.Chat.SubjectID, + MaxMessageBytes: cfg.SuperAgent.MaxMessageBytes, + SessionTTL: cfg.Chat.SessionTTL, + MaxSessions: cfg.Chat.MaxSessions, + }) + if err != nil { + return nil, fmt.Errorf("initialize chat service: %w", err) + } + chatHandler, err := handler.NewChatHandler(chatService, handler.ChatOptions{ + AuthToken: cfg.Chat.AuthToken, + AllowedOrigins: cfg.Chat.AllowedOrigins, + MaxBodyBytes: cfg.Chat.MaxBodyBytes, + RunTimeout: cfg.Chat.RunTimeout, + Logger: log.Default(), + }) + if err != nil { + return nil, fmt.Errorf("initialize chat handler: %w", err) + } + routerOptions.Chat = chatHandler + if cfg.Chat.CompatAppID != "" { + dashScopeChatHandler, err := handler.NewDashScopeChatHandler(chatService, handler.DashScopeChatOptions{ + AppID: cfg.Chat.CompatAppID, + AuthToken: cfg.Chat.AuthToken, + AllowedOrigins: cfg.Chat.AllowedOrigins, + MaxBodyBytes: cfg.Chat.MaxBodyBytes, + RunTimeout: cfg.Chat.RunTimeout, + Logger: log.Default(), + }) + if err != nil { + return nil, fmt.Errorf("initialize DashScope-compatible chat handler: %w", err) + } + routerOptions.DashScopeChat = dashScopeChatHandler + } + } + + if cfg.MCP.Enabled { + log.Printf("MCP data scope mode=%s", cfg.MCP.ScopeMode) + postGIS, err := repository.OpenPostGIS(ctx, repository.PostGISOptions{ + DSN: cfg.PostGIS.DSN, + MaxConns: cfg.PostGIS.MaxConns, + ConnectTimeout: cfg.PostGIS.ConnectTimeout, + QueryTimeout: cfg.PostGIS.QueryTimeout, + }) + if err != nil { + return nil, fmt.Errorf("initialize PostGIS: %w", err) + } + closeDependencies = postGIS.Close + + readinessCtx, cancel := context.WithTimeout(ctx, readinessTimeout) + report, err := postGIS.ValidateForMCP(readinessCtx, cfg.PostGIS.ExpectedSRID) + cancel() + if err != nil { + postGIS.Close() + return nil, fmt.Errorf("validate PostGIS for MCP: %w", err) + } + for _, warning := range report.Warnings { + log.Printf("PostGIS readiness warning=%s", warning) + } + + spatialService, err := service.NewSpatialService(postGIS, domain.SpatialScope{ + AllTowns: cfg.MCP.ScopeMode == config.MCPScopeModeAll, + AllowedTowns: cfg.MCP.AllowedTowns, + }, cfg.PostGIS.QueryTimeout) + if err != nil { + postGIS.Close() + return nil, fmt.Errorf("initialize spatial service: %w", err) + } + mcpHandler, err := handler.NewMCPHandler(spatialService, handler.MCPOptions{ + AuthToken: cfg.MCP.AuthToken, + MaxBodyBytes: cfg.MCP.MaxBodyBytes, + ToolTimeout: cfg.MCP.ToolTimeout, + Logger: log.Default(), + }) + if err != nil { + postGIS.Close() + return nil, fmt.Errorf("initialize MCP handler: %w", err) + } + routerOptions.MCP = mcpHandler + } + + return &Application{ + server: &http.Server{ + Addr: cfg.HTTPAddress, + Handler: handler.NewRouter(routerOptions), + ReadHeaderTimeout: readHeaderTimeout, + ReadTimeout: readTimeout, + WriteTimeout: writeTimeout, + IdleTimeout: idleTimeout, + }, + close: closeDependencies, + }, nil +} + +// Run starts the HTTP server and shuts it down when ctx is canceled. +func (a *Application) Run(ctx context.Context) error { + defer a.close() + + serverErr := make(chan error, 1) + go func() { + serverErr <- a.server.ListenAndServe() + }() + + select { + case err := <-serverErr: + return normalizeServerError(err) + case <-ctx.Done(): + shutdownCtx, cancel := context.WithTimeout(context.Background(), shutdownTimeout) + defer cancel() + + if err := a.server.Shutdown(shutdownCtx); err != nil { + return fmt.Errorf("shut down HTTP server: %w", err) + } + + return normalizeServerError(<-serverErr) + } +} + +func normalizeServerError(err error) error { + if err == nil || errors.Is(err, http.ErrServerClosed) { + return nil + } + + return fmt.Errorf("serve HTTP: %w", err) +} diff --git a/internal/app/app_test.go b/internal/app/app_test.go new file mode 100644 index 0000000..82f8cc4 --- /dev/null +++ b/internal/app/app_test.go @@ -0,0 +1,116 @@ +package app + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "fire-safety-ymd/internal/config" +) + +func TestNewKeepsExternalIntegrationsDisabledByDefault(t *testing.T) { + application, err := New(context.Background(), config.Config{HTTPAddress: ":0"}) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + healthRequest := httptest.NewRequest(http.MethodGet, "/health", nil) + healthResponse := httptest.NewRecorder() + application.server.Handler.ServeHTTP(healthResponse, healthRequest) + if healthResponse.Code != http.StatusOK { + t.Fatalf("health status = %d, want %d", healthResponse.Code, http.StatusOK) + } + + mcpRequest := httptest.NewRequest(http.MethodPost, "/mcp", nil) + mcpResponse := httptest.NewRecorder() + application.server.Handler.ServeHTTP(mcpResponse, mcpRequest) + if mcpResponse.Code != http.StatusNotFound { + t.Fatalf("MCP status = %d, want %d while disabled", mcpResponse.Code, http.StatusNotFound) + } + chatRequest := httptest.NewRequest(http.MethodPost, "/api/chat", nil) + chatResponse := httptest.NewRecorder() + application.server.Handler.ServeHTTP(chatResponse, chatRequest) + if chatResponse.Code != http.StatusNotFound { + t.Fatalf("chat status = %d, want %d while disabled", chatResponse.Code, http.StatusNotFound) + } + if application.server.ReadHeaderTimeout <= 0 || application.server.ReadTimeout <= 0 || application.server.WriteTimeout <= 0 || application.server.IdleTimeout <= 0 { + t.Fatalf("HTTP timeouts are incomplete: %#v", application.server) + } +} + +func TestNewWiresChatAPIToSuperAgent(t *testing.T) { + provider := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "Bearer provider-open-api-key" { + t.Errorf("provider Authorization = %q", r.Header.Get("Authorization")) + } + switch r.URL.Path { + case "/api/open/agent-sessions": + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprint(w, `{"session_id":"provider-session-1"}`) + case "/api/open/agent-sessions/provider-session-1/messages/stream": + if r.URL.Query().Get("include_trace") != "true" { + t.Errorf("include_trace = %q", r.URL.Query().Get("include_trace")) + } + w.Header().Set("Content-Type", "text/event-stream") + _, _ = fmt.Fprint(w, "event: trace\ndata: {\"event\":\"message.final\",\"text\":\"测试回答\",\"run_id\":\"run-1\"}\n\n") + _, _ = fmt.Fprint(w, "event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\",\"run_id\":\"run-1\"}\n\n") + _, _ = fmt.Fprint(w, "event: end\n\n") + default: + http.NotFound(w, r) + } + })) + defer provider.Close() + + application, err := New(context.Background(), config.Config{ + HTTPAddress: ":0", + SuperAgent: config.SuperAgentConfig{ + Enabled: true, + BaseURL: provider.URL, + OpenAPIKey: "provider-open-api-key", + ConnectTimeout: time.Second, + RecoveryMaxAttempts: 1, + RecoveryInitialBackoff: time.Millisecond, + MaxMessageBytes: 4096, + }, + Chat: config.ChatConfig{ + Enabled: true, + AuthToken: "0123456789abcdef0123456789abcdef", + SubjectID: "app-chat-test-subject", + CompatAppID: "fire-safety-app", + MaxBodyBytes: 4096, + RunTimeout: time.Minute, + SessionTTL: time.Minute, + MaxSessions: 10, + }, + }) + if err != nil { + t.Fatalf("New() error = %v", err) + } + defer application.close() + + request := httptest.NewRequest(http.MethodPost, "/api/chat", strings.NewReader(`{"message":"你好"}`)) + request.Header.Set("Authorization", "Bearer 0123456789abcdef0123456789abcdef") + request.Header.Set("Content-Type", "application/json") + request.Header.Set("Accept", "text/event-stream") + response := httptest.NewRecorder() + application.server.Handler.ServeHTTP(response, request) + + if response.Code != http.StatusOK || !strings.Contains(response.Body.String(), "event: message") || + !strings.Contains(response.Body.String(), "测试回答") || !strings.Contains(response.Body.String(), "event: done") { + t.Fatalf("chat status=%d body=%s", response.Code, response.Body.String()) + } + + compatRequest := httptest.NewRequest(http.MethodPost, "/api/v1/apps/fire-safety-app/completion", strings.NewReader(`{"input":{"prompt":"你好"},"parameters":{}}`)) + compatRequest.Header.Set("xtoken", "0123456789abcdef0123456789abcdef") + compatRequest.Header.Set("Content-Type", "application/json") + compatResponse := httptest.NewRecorder() + application.server.Handler.ServeHTTP(compatResponse, compatRequest) + if compatResponse.Code != http.StatusOK || !strings.Contains(compatResponse.Body.String(), "event: result") || + !strings.Contains(compatResponse.Body.String(), `"finish_reason":"stop"`) || !strings.Contains(compatResponse.Body.String(), "测试回答") { + t.Fatalf("compat chat status=%d body=%s", compatResponse.Code, compatResponse.Body.String()) + } +} diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..16cf10b --- /dev/null +++ b/internal/config/config.go @@ -0,0 +1,631 @@ +// Package config loads and validates process configuration. +package config + +import ( + "fmt" + "net/url" + "os" + "strconv" + "strings" + "time" + "unicode" +) + +const ( + // HTTPAddressEnv is the environment variable used to configure the listen address. + HTTPAddressEnv = "FIRE_SAFETY_HTTP_ADDR" + + SuperAgentEnabledEnv = "FIRE_SAFETY_SUPERAGENT_ENABLED" + SuperAgentBaseURLEnv = "FIRE_SAFETY_SUPERAGENT_BASE_URL" + SuperAgentOpenAPIKeyEnv = "FIRE_SAFETY_SUPERAGENT_OPEN_API_KEY" + SuperAgentConnectTimeoutEnv = "FIRE_SAFETY_SUPERAGENT_CONNECT_TIMEOUT" + SuperAgentRecoveryMaxAttemptsEnv = "FIRE_SAFETY_SUPERAGENT_RECOVERY_MAX_ATTEMPTS" + SuperAgentRecoveryInitialBackoffEnv = "FIRE_SAFETY_SUPERAGENT_RECOVERY_INITIAL_BACKOFF" + SuperAgentMaxMessageBytesEnv = "FIRE_SAFETY_SUPERAGENT_MAX_MESSAGE_BYTES" + SuperAgentProbeSubjectIDEnv = "FIRE_SAFETY_SUPERAGENT_PROBE_SUBJECT_ID" + SuperAgentProbeTimeoutEnv = "FIRE_SAFETY_SUPERAGENT_PROBE_TIMEOUT" + + ChatEnabledEnv = "FIRE_SAFETY_CHAT_ENABLED" + ChatAuthTokenEnv = "FIRE_SAFETY_CHAT_AUTH_TOKEN" + ChatSubjectIDEnv = "FIRE_SAFETY_CHAT_SUBJECT_ID" + ChatAllowedOriginsEnv = "FIRE_SAFETY_CHAT_ALLOWED_ORIGINS" + ChatMaxBodyBytesEnv = "FIRE_SAFETY_CHAT_MAX_BODY_BYTES" + ChatRunTimeoutEnv = "FIRE_SAFETY_CHAT_RUN_TIMEOUT" + ChatSessionTTLEnv = "FIRE_SAFETY_CHAT_SESSION_TTL" + ChatMaxSessionsEnv = "FIRE_SAFETY_CHAT_MAX_SESSIONS" + ChatCompatAppIDEnv = "FIRE_SAFETY_CHAT_COMPAT_APP_ID" + + MCPEnabledEnv = "FIRE_SAFETY_MCP_ENABLED" + MCPAuthTokenEnv = "FIRE_SAFETY_MCP_AUTH_TOKEN" + MCPScopeModeEnv = "FIRE_SAFETY_MCP_SCOPE_MODE" + MCPAllowedTownsEnv = "FIRE_SAFETY_MCP_ALLOWED_TOWNS" + MCPMaxBodyBytesEnv = "FIRE_SAFETY_MCP_MAX_BODY_BYTES" + MCPToolTimeoutEnv = "FIRE_SAFETY_MCP_TOOL_TIMEOUT" + + PostGISEnabledEnv = "FIRE_SAFETY_POSTGIS_ENABLED" + PostGISDSNEnv = "FIRE_SAFETY_POSTGIS_DSN" + PostGISExpectedSRIDEnv = "FIRE_SAFETY_POSTGIS_EXPECTED_SRID" + PostGISConnectTimeoutEnv = "FIRE_SAFETY_POSTGIS_CONNECT_TIMEOUT" + PostGISQueryTimeoutEnv = "FIRE_SAFETY_POSTGIS_QUERY_TIMEOUT" + PostGISMaxConnsEnv = "FIRE_SAFETY_POSTGIS_MAX_CONNS" + + defaultHTTPAddress = ":8080" + defaultSuperAgentConnectTimeout = 15 * time.Second + defaultSuperAgentRecoveryMaxAttempts = 5 + maximumSuperAgentRecoveryAttempts = 20 + defaultSuperAgentRecoveryBackoff = 250 * time.Millisecond + defaultSuperAgentMaxMessageBytes int64 = 64 * 1024 + maximumSuperAgentMessageBytes int64 = 16 * 1024 * 1024 + defaultSuperAgentProbeSubjectID = "fire-safety-ymd-connectivity-probe" + defaultSuperAgentProbeTimeout = 10 * time.Minute + defaultChatSubjectID = "fire-safety-ymd-chat-test-subject" + defaultChatMaxBodyBytes int64 = 128 * 1024 + maximumChatBodyBytes int64 = 1024 * 1024 + defaultChatRunTimeout = 10 * time.Minute + maximumChatRunTimeout = 30 * time.Minute + defaultChatSessionTTL = 30 * time.Minute + maximumChatSessionTTL = 24 * time.Hour + defaultChatMaxSessions int64 = 1000 + maximumChatMaxSessions int64 = 10_000 + maximumChatAllowedOrigins = 50 + defaultMCPMaxBodyBytes int64 = 256 * 1024 + maximumMCPBodyBytes int64 = 1024 * 1024 + defaultMCPToolTimeout = 5 * time.Second + maximumMCPToolTimeout = 30 * time.Second + defaultMCPScopeMode = MCPScopeModeTownAllowlist + defaultPostGISConnectTimeout = 5 * time.Second + defaultPostGISQueryTimeout = 3 * time.Second + maximumPostGISTimeout = 30 * time.Second + defaultPostGISMaxConns int64 = 4 + maximumPostGISMaxConns int64 = 20 + requiredMCPSpatialSRID = 4326 +) + +// MCPScopeMode defines the trusted server-side data scope applied to every MCP query. +type MCPScopeMode string + +const ( + // MCPScopeModeTownAllowlist restricts results to configured town names. + MCPScopeModeTownAllowlist MCPScopeMode = "town_allowlist" + // MCPScopeModeAll permits every town in the MCP's fixed query tables. + MCPScopeModeAll MCPScopeMode = "all" +) + +// Config contains immutable application configuration. +type Config struct { + HTTPAddress string + SuperAgent SuperAgentConfig + Chat ChatConfig + MCP MCPConfig + PostGIS PostGISConfig +} + +// SuperAgentConfig contains outbound SuperAgent Open API settings. +type SuperAgentConfig struct { + Enabled bool + BaseURL string + OpenAPIKey string + ConnectTimeout time.Duration + RecoveryMaxAttempts int + RecoveryInitialBackoff time.Duration + MaxMessageBytes int64 + ProbeSubjectID string + ProbeTimeout time.Duration +} + +// ChatConfig contains the protected user-facing chat transport and in-memory +// conversation settings. The static credential is a test-stage boundary, not +// final end-user authentication. +type ChatConfig struct { + Enabled bool + AuthToken string + SubjectID string + CompatAppID string + AllowedOrigins []string + MaxBodyBytes int64 + RunTimeout time.Duration + SessionTTL time.Duration + MaxSessions int +} + +// MCPConfig contains inbound SuperAgent MCP settings. +type MCPConfig struct { + Enabled bool + AuthToken string + ScopeMode MCPScopeMode + AllowedTowns []string + MaxBodyBytes int64 + ToolTimeout time.Duration +} + +// PostGISConfig contains PostgreSQL/PostGIS connection and query settings. +type PostGISConfig struct { + Enabled bool + DSN string + ExpectedSRID int + ConnectTimeout time.Duration + QueryTimeout time.Duration + MaxConns int32 +} + +// Load reads configuration from the environment, applies safe defaults, and validates it. +func Load() (Config, error) { + address := strings.TrimSpace(os.Getenv(HTTPAddressEnv)) + if address == "" { + address = defaultHTTPAddress + } + + enabled, err := parseBool(SuperAgentEnabledEnv, false) + if err != nil { + return Config{}, err + } + connectTimeout, err := parseDuration(SuperAgentConnectTimeoutEnv, defaultSuperAgentConnectTimeout) + if err != nil { + return Config{}, err + } + recoveryMaxAttempts, err := parseInt(SuperAgentRecoveryMaxAttemptsEnv, defaultSuperAgentRecoveryMaxAttempts) + if err != nil { + return Config{}, err + } + recoveryInitialBackoff, err := parseDuration(SuperAgentRecoveryInitialBackoffEnv, defaultSuperAgentRecoveryBackoff) + if err != nil { + return Config{}, err + } + maxMessageBytes, err := parseInt64(SuperAgentMaxMessageBytesEnv, defaultSuperAgentMaxMessageBytes) + if err != nil { + return Config{}, err + } + probeTimeout, err := parseDuration(SuperAgentProbeTimeoutEnv, defaultSuperAgentProbeTimeout) + if err != nil { + return Config{}, err + } + + superAgent := SuperAgentConfig{ + Enabled: enabled, + BaseURL: strings.TrimRight(strings.TrimSpace(os.Getenv(SuperAgentBaseURLEnv)), "/"), + OpenAPIKey: strings.TrimSpace(os.Getenv(SuperAgentOpenAPIKeyEnv)), + ConnectTimeout: connectTimeout, + RecoveryMaxAttempts: recoveryMaxAttempts, + RecoveryInitialBackoff: recoveryInitialBackoff, + MaxMessageBytes: maxMessageBytes, + ProbeSubjectID: valueOrDefault(SuperAgentProbeSubjectIDEnv, defaultSuperAgentProbeSubjectID), + ProbeTimeout: probeTimeout, + } + if err := validateSuperAgent(superAgent); err != nil { + return Config{}, err + } + + chat, err := loadChatConfig() + if err != nil { + return Config{}, err + } + + mcp, err := loadMCPConfig() + if err != nil { + return Config{}, err + } + postGIS, err := loadPostGISConfig() + if err != nil { + return Config{}, err + } + if err := validateMCPDependencies(mcp, postGIS, superAgent); err != nil { + return Config{}, err + } + if err := validateChatDependencies(chat, superAgent, mcp); err != nil { + return Config{}, err + } + + return Config{ + HTTPAddress: address, + SuperAgent: superAgent, + Chat: chat, + MCP: mcp, + PostGIS: postGIS, + }, nil +} + +func loadChatConfig() (ChatConfig, error) { + enabled, err := parseBool(ChatEnabledEnv, false) + if err != nil { + return ChatConfig{}, err + } + maxBodyBytes, err := parseInt64(ChatMaxBodyBytesEnv, defaultChatMaxBodyBytes) + if err != nil { + return ChatConfig{}, err + } + runTimeout, err := parseDuration(ChatRunTimeoutEnv, defaultChatRunTimeout) + if err != nil { + return ChatConfig{}, err + } + sessionTTL, err := parseDuration(ChatSessionTTLEnv, defaultChatSessionTTL) + if err != nil { + return ChatConfig{}, err + } + maxSessions, err := parseInt64(ChatMaxSessionsEnv, defaultChatMaxSessions) + if err != nil { + return ChatConfig{}, err + } + allowedOrigins, err := normalizeChatOrigins(parseList(os.Getenv(ChatAllowedOriginsEnv))) + if err != nil { + return ChatConfig{}, err + } + + cfg := ChatConfig{ + Enabled: enabled, + AuthToken: strings.TrimSpace(os.Getenv(ChatAuthTokenEnv)), + SubjectID: valueOrDefault(ChatSubjectIDEnv, defaultChatSubjectID), + CompatAppID: strings.TrimSpace(os.Getenv(ChatCompatAppIDEnv)), + AllowedOrigins: allowedOrigins, + MaxBodyBytes: maxBodyBytes, + RunTimeout: runTimeout, + SessionTTL: sessionTTL, + MaxSessions: int(maxSessions), + } + if cfg.MaxBodyBytes <= 0 || cfg.MaxBodyBytes > maximumChatBodyBytes { + return ChatConfig{}, fmt.Errorf("%s must be between 1 and %d", ChatMaxBodyBytesEnv, maximumChatBodyBytes) + } + if cfg.RunTimeout <= 0 || cfg.RunTimeout > maximumChatRunTimeout { + return ChatConfig{}, fmt.Errorf("%s must be greater than zero and not exceed %s", ChatRunTimeoutEnv, maximumChatRunTimeout) + } + if cfg.SessionTTL <= 0 || cfg.SessionTTL > maximumChatSessionTTL { + return ChatConfig{}, fmt.Errorf("%s must be greater than zero and not exceed %s", ChatSessionTTLEnv, maximumChatSessionTTL) + } + if maxSessions <= 0 || maxSessions > maximumChatMaxSessions { + return ChatConfig{}, fmt.Errorf("%s must be between 1 and %d", ChatMaxSessionsEnv, maximumChatMaxSessions) + } + if strings.TrimSpace(cfg.SubjectID) == "" || len(cfg.SubjectID) > 512 || containsControlCharacter(cfg.SubjectID) { + return ChatConfig{}, fmt.Errorf("%s must be a non-empty value of at most 512 bytes without control characters", ChatSubjectIDEnv) + } + if cfg.CompatAppID != "" && !validChatCompatAppID(cfg.CompatAppID) { + return ChatConfig{}, fmt.Errorf("%s must contain 1 to 128 ASCII letters, digits, underscores, or hyphens", ChatCompatAppIDEnv) + } + if cfg.Enabled && (len(cfg.AuthToken) < 32 || !validSecretHeaderValue(cfg.AuthToken)) { + return ChatConfig{}, fmt.Errorf("%s must contain at least 32 printable ASCII characters", ChatAuthTokenEnv) + } + return cfg, nil +} + +func validChatCompatAppID(value string) bool { + if value == "" || len(value) > 128 { + return false + } + for _, character := range value { + if !(character >= 'a' && character <= 'z') && + !(character >= 'A' && character <= 'Z') && + !(character >= '0' && character <= '9') && + character != '_' && character != '-' { + return false + } + } + return true +} + +func normalizeChatOrigins(values []string) ([]string, error) { + if len(values) > maximumChatAllowedOrigins { + return nil, fmt.Errorf("%s must contain no more than %d values", ChatAllowedOriginsEnv, maximumChatAllowedOrigins) + } + normalized := make([]string, 0, len(values)) + seen := make(map[string]struct{}, len(values)) + for _, value := range values { + parsed, err := url.Parse(value) + if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") || + parsed.User != nil || parsed.RawQuery != "" || parsed.ForceQuery || parsed.Fragment != "" || + (parsed.Path != "" && parsed.Path != "/") || value == "*" { + return nil, fmt.Errorf("%s must contain only exact HTTP(S) origins", ChatAllowedOriginsEnv) + } + origin := parsed.Scheme + "://" + parsed.Host + if _, exists := seen[origin]; exists { + continue + } + seen[origin] = struct{}{} + normalized = append(normalized, origin) + } + return normalized, nil +} + +func loadMCPConfig() (MCPConfig, error) { + enabled, err := parseBool(MCPEnabledEnv, false) + if err != nil { + return MCPConfig{}, err + } + scopeMode, err := parseMCPScopeMode(os.Getenv(MCPScopeModeEnv)) + if err != nil { + return MCPConfig{}, err + } + maxBodyBytes, err := parseInt64(MCPMaxBodyBytesEnv, defaultMCPMaxBodyBytes) + if err != nil { + return MCPConfig{}, err + } + toolTimeout, err := parseDuration(MCPToolTimeoutEnv, defaultMCPToolTimeout) + if err != nil { + return MCPConfig{}, err + } + + cfg := MCPConfig{ + Enabled: enabled, + AuthToken: strings.TrimSpace(os.Getenv(MCPAuthTokenEnv)), + ScopeMode: scopeMode, + AllowedTowns: parseList(os.Getenv(MCPAllowedTownsEnv)), + MaxBodyBytes: maxBodyBytes, + ToolTimeout: toolTimeout, + } + if cfg.MaxBodyBytes <= 0 || cfg.MaxBodyBytes > maximumMCPBodyBytes { + return MCPConfig{}, fmt.Errorf("%s must be between 1 and %d", MCPMaxBodyBytesEnv, maximumMCPBodyBytes) + } + if cfg.ToolTimeout <= 0 || cfg.ToolTimeout > maximumMCPToolTimeout { + return MCPConfig{}, fmt.Errorf("%s must be greater than zero and not exceed %s", MCPToolTimeoutEnv, maximumMCPToolTimeout) + } + if !cfg.Enabled { + return cfg, nil + } + if len(cfg.AuthToken) < 32 || !validSecretHeaderValue(cfg.AuthToken) { + return MCPConfig{}, fmt.Errorf("%s must contain at least 32 printable ASCII characters", MCPAuthTokenEnv) + } + switch cfg.ScopeMode { + case MCPScopeModeAll: + if len(cfg.AllowedTowns) > 0 { + return MCPConfig{}, fmt.Errorf("%s must be empty when %s is %s", MCPAllowedTownsEnv, MCPScopeModeEnv, MCPScopeModeAll) + } + case MCPScopeModeTownAllowlist: + if len(cfg.AllowedTowns) == 0 { + return MCPConfig{}, fmt.Errorf("%s is required when %s is %s", MCPAllowedTownsEnv, MCPScopeModeEnv, MCPScopeModeTownAllowlist) + } + if len(cfg.AllowedTowns) > 100 { + return MCPConfig{}, fmt.Errorf("%s must contain no more than 100 values", MCPAllowedTownsEnv) + } + for _, town := range cfg.AllowedTowns { + if len([]rune(town)) > 100 || containsControlCharacter(town) { + return MCPConfig{}, fmt.Errorf("%s contains an invalid town value", MCPAllowedTownsEnv) + } + } + } + return cfg, nil +} + +func parseMCPScopeMode(raw string) (MCPScopeMode, error) { + value := MCPScopeMode(strings.ToLower(strings.TrimSpace(raw))) + if value == "" { + return defaultMCPScopeMode, nil + } + switch value { + case MCPScopeModeTownAllowlist, MCPScopeModeAll: + return value, nil + default: + return "", fmt.Errorf("%s must be %q or %q", MCPScopeModeEnv, MCPScopeModeTownAllowlist, MCPScopeModeAll) + } +} + +func containsControlCharacter(value string) bool { + for _, character := range value { + if unicode.IsControl(character) { + return true + } + } + return false +} + +func loadPostGISConfig() (PostGISConfig, error) { + enabled, err := parseBool(PostGISEnabledEnv, false) + if err != nil { + return PostGISConfig{}, err + } + expectedSRID, err := parseInt(PostGISExpectedSRIDEnv, 0) + if err != nil { + return PostGISConfig{}, err + } + connectTimeout, err := parseDuration(PostGISConnectTimeoutEnv, defaultPostGISConnectTimeout) + if err != nil { + return PostGISConfig{}, err + } + queryTimeout, err := parseDuration(PostGISQueryTimeoutEnv, defaultPostGISQueryTimeout) + if err != nil { + return PostGISConfig{}, err + } + maxConns, err := parseInt64(PostGISMaxConnsEnv, defaultPostGISMaxConns) + if err != nil { + return PostGISConfig{}, err + } + + cfg := PostGISConfig{ + Enabled: enabled, + DSN: strings.TrimSpace(os.Getenv(PostGISDSNEnv)), + ExpectedSRID: expectedSRID, + ConnectTimeout: connectTimeout, + QueryTimeout: queryTimeout, + MaxConns: int32(maxConns), + } + if cfg.ConnectTimeout <= 0 || cfg.ConnectTimeout > maximumPostGISTimeout { + return PostGISConfig{}, fmt.Errorf("%s must be greater than zero and not exceed %s", PostGISConnectTimeoutEnv, maximumPostGISTimeout) + } + if cfg.QueryTimeout <= 0 || cfg.QueryTimeout > maximumPostGISTimeout { + return PostGISConfig{}, fmt.Errorf("%s must be greater than zero and not exceed %s", PostGISQueryTimeoutEnv, maximumPostGISTimeout) + } + if maxConns <= 0 || maxConns > maximumPostGISMaxConns { + return PostGISConfig{}, fmt.Errorf("%s must be between 1 and %d", PostGISMaxConnsEnv, maximumPostGISMaxConns) + } + if cfg.ExpectedSRID < 0 || cfg.ExpectedSRID > 999999 { + return PostGISConfig{}, fmt.Errorf("%s must be between 0 and 999999", PostGISExpectedSRIDEnv) + } + if !cfg.Enabled { + return cfg, nil + } + if cfg.DSN == "" { + return PostGISConfig{}, fmt.Errorf("%s is required when %s is true", PostGISDSNEnv, PostGISEnabledEnv) + } + if !validPostGISDSN(cfg.DSN) { + return PostGISConfig{}, fmt.Errorf("%s must be an absolute postgresql:// or postgres:// URL", PostGISDSNEnv) + } + return cfg, nil +} + +func validateMCPDependencies(mcp MCPConfig, postGIS PostGISConfig, superAgent SuperAgentConfig) error { + if !mcp.Enabled { + return nil + } + if !postGIS.Enabled { + return fmt.Errorf("%s must be true when %s is true", PostGISEnabledEnv, MCPEnabledEnv) + } + if postGIS.ExpectedSRID != requiredMCPSpatialSRID { + return fmt.Errorf("%s must be explicitly set to %d when %s is true", PostGISExpectedSRIDEnv, requiredMCPSpatialSRID, MCPEnabledEnv) + } + if superAgent.OpenAPIKey != "" && mcp.AuthToken == superAgent.OpenAPIKey { + return fmt.Errorf("%s must not reuse %s", MCPAuthTokenEnv, SuperAgentOpenAPIKeyEnv) + } + return nil +} + +func validateChatDependencies(chat ChatConfig, superAgent SuperAgentConfig, mcp MCPConfig) error { + if !chat.Enabled { + return nil + } + if !superAgent.Enabled { + return fmt.Errorf("%s must be true when %s is true", SuperAgentEnabledEnv, ChatEnabledEnv) + } + if chat.AuthToken == superAgent.OpenAPIKey { + return fmt.Errorf("%s must not reuse %s", ChatAuthTokenEnv, SuperAgentOpenAPIKeyEnv) + } + if mcp.AuthToken != "" && chat.AuthToken == mcp.AuthToken { + return fmt.Errorf("%s must not reuse %s", ChatAuthTokenEnv, MCPAuthTokenEnv) + } + return nil +} + +func validateSuperAgent(cfg SuperAgentConfig) error { + if cfg.ConnectTimeout <= 0 { + return fmt.Errorf("%s must be greater than zero", SuperAgentConnectTimeoutEnv) + } + if cfg.RecoveryMaxAttempts < 0 { + return fmt.Errorf("%s must not be negative", SuperAgentRecoveryMaxAttemptsEnv) + } + if cfg.RecoveryMaxAttempts > maximumSuperAgentRecoveryAttempts { + return fmt.Errorf("%s must not exceed %d", SuperAgentRecoveryMaxAttemptsEnv, maximumSuperAgentRecoveryAttempts) + } + if cfg.RecoveryInitialBackoff <= 0 { + return fmt.Errorf("%s must be greater than zero", SuperAgentRecoveryInitialBackoffEnv) + } + if cfg.MaxMessageBytes <= 0 { + return fmt.Errorf("%s must be greater than zero", SuperAgentMaxMessageBytesEnv) + } + if cfg.MaxMessageBytes > maximumSuperAgentMessageBytes { + return fmt.Errorf("%s must not exceed %d", SuperAgentMaxMessageBytesEnv, maximumSuperAgentMessageBytes) + } + if cfg.ProbeTimeout <= 0 { + return fmt.Errorf("%s must be greater than zero", SuperAgentProbeTimeoutEnv) + } + if !cfg.Enabled { + return nil + } + if cfg.BaseURL == "" { + return fmt.Errorf("%s is required when %s is true", SuperAgentBaseURLEnv, SuperAgentEnabledEnv) + } + if cfg.OpenAPIKey == "" { + return fmt.Errorf("%s is required when %s is true", SuperAgentOpenAPIKeyEnv, SuperAgentEnabledEnv) + } + if !validSecretHeaderValue(cfg.OpenAPIKey) { + return fmt.Errorf("%s must be a valid header token", SuperAgentOpenAPIKeyEnv) + } + if strings.TrimSpace(cfg.ProbeSubjectID) == "" || len(cfg.ProbeSubjectID) > 512 { + return fmt.Errorf("%s is required when %s is true", SuperAgentProbeSubjectIDEnv, SuperAgentEnabledEnv) + } + parsed, err := url.Parse(cfg.BaseURL) + if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { + return fmt.Errorf("%s must be an absolute HTTP(S) URL", SuperAgentBaseURLEnv) + } + if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { + return fmt.Errorf("%s must not include user information, query, or fragment", SuperAgentBaseURLEnv) + } + return nil +} + +func validSecretHeaderValue(value string) bool { + if value == "" || len(value) > 4096 { + return false + } + for _, character := range value { + if character <= 0x20 || character >= 0x7f { + return false + } + } + return true +} + +func validPostGISDSN(value string) bool { + parsed, err := url.Parse(value) + if err != nil || parsed.Host == "" { + return false + } + return parsed.Scheme == "postgres" || parsed.Scheme == "postgresql" +} + +func parseList(raw string) []string { + seen := make(map[string]struct{}) + values := make([]string, 0) + for _, item := range strings.Split(raw, ",") { + value := strings.TrimSpace(item) + if value == "" { + continue + } + if _, ok := seen[value]; ok { + continue + } + seen[value] = struct{}{} + values = append(values, value) + } + return values +} + +func parseBool(key string, fallback bool) (bool, error) { + raw := strings.TrimSpace(os.Getenv(key)) + if raw == "" { + return fallback, nil + } + value, err := strconv.ParseBool(raw) + if err != nil { + return false, fmt.Errorf("%s must be a boolean", key) + } + return value, nil +} + +func parseDuration(key string, fallback time.Duration) (time.Duration, error) { + raw := strings.TrimSpace(os.Getenv(key)) + if raw == "" { + return fallback, nil + } + value, err := time.ParseDuration(raw) + if err != nil { + return 0, fmt.Errorf("%s must be a valid duration", key) + } + return value, nil +} + +func parseInt(key string, fallback int) (int, error) { + raw := strings.TrimSpace(os.Getenv(key)) + if raw == "" { + return fallback, nil + } + value, err := strconv.Atoi(raw) + if err != nil { + return 0, fmt.Errorf("%s must be an integer", key) + } + return value, nil +} + +func parseInt64(key string, fallback int64) (int64, error) { + raw := strings.TrimSpace(os.Getenv(key)) + if raw == "" { + return fallback, nil + } + value, err := strconv.ParseInt(raw, 10, 64) + if err != nil { + return 0, fmt.Errorf("%s must be an integer", key) + } + return value, nil +} + +func valueOrDefault(key, fallback string) string { + value := strings.TrimSpace(os.Getenv(key)) + if value == "" { + return fallback + } + return value +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..b262278 --- /dev/null +++ b/internal/config/config_test.go @@ -0,0 +1,490 @@ +package config + +import ( + "strings" + "testing" + "time" +) + +func TestLoadDefaults(t *testing.T) { + clearEnvironment(t) + + cfg, err := Load() + if err != nil { + t.Fatalf("Load() error = %v", err) + } + if cfg.HTTPAddress != ":8080" { + t.Fatalf("HTTPAddress = %q, want %q", cfg.HTTPAddress, ":8080") + } + if cfg.SuperAgent.Enabled { + t.Fatal("SuperAgent.Enabled = true, want false") + } + if cfg.SuperAgent.ConnectTimeout != 15*time.Second { + t.Fatalf("ConnectTimeout = %v, want 15s", cfg.SuperAgent.ConnectTimeout) + } + if cfg.SuperAgent.RecoveryMaxAttempts != 5 { + t.Fatalf("RecoveryMaxAttempts = %d, want 5", cfg.SuperAgent.RecoveryMaxAttempts) + } + if cfg.SuperAgent.RecoveryInitialBackoff != 250*time.Millisecond { + t.Fatalf("RecoveryInitialBackoff = %v, want 250ms", cfg.SuperAgent.RecoveryInitialBackoff) + } + if cfg.SuperAgent.MaxMessageBytes != 64*1024 { + t.Fatalf("MaxMessageBytes = %d, want %d", cfg.SuperAgent.MaxMessageBytes, 64*1024) + } + if cfg.SuperAgent.ProbeSubjectID != "fire-safety-ymd-connectivity-probe" { + t.Fatalf("ProbeSubjectID = %q", cfg.SuperAgent.ProbeSubjectID) + } + if cfg.SuperAgent.ProbeTimeout != 10*time.Minute { + t.Fatalf("ProbeTimeout = %v, want 10m", cfg.SuperAgent.ProbeTimeout) + } + if cfg.Chat.Enabled { + t.Fatal("Chat.Enabled = true, want false") + } + if cfg.Chat.SubjectID != "fire-safety-ymd-chat-test-subject" || cfg.Chat.MaxBodyBytes != 128*1024 || + cfg.Chat.CompatAppID != "" || cfg.Chat.RunTimeout != 10*time.Minute || cfg.Chat.SessionTTL != 30*time.Minute || cfg.Chat.MaxSessions != 1000 { + t.Fatalf("unexpected Chat defaults: %#v", cfg.Chat) + } + if cfg.MCP.Enabled || cfg.PostGIS.Enabled { + t.Fatalf("MCP/PostGIS must be disabled by default: MCP=%v PostGIS=%v", cfg.MCP.Enabled, cfg.PostGIS.Enabled) + } + if cfg.MCP.MaxBodyBytes != 256*1024 || cfg.MCP.ToolTimeout != 5*time.Second { + t.Fatalf("unexpected MCP defaults: %#v", cfg.MCP) + } + if cfg.MCP.ScopeMode != MCPScopeModeTownAllowlist { + t.Fatalf("MCP.ScopeMode = %q, want %q", cfg.MCP.ScopeMode, MCPScopeModeTownAllowlist) + } + if cfg.PostGIS.ConnectTimeout != 5*time.Second || cfg.PostGIS.QueryTimeout != 3*time.Second || cfg.PostGIS.MaxConns != 4 { + t.Fatalf("unexpected PostGIS defaults: %#v", cfg.PostGIS) + } +} + +func TestLoadConfiguredChat(t *testing.T) { + clearEnvironment(t) + t.Setenv(SuperAgentEnabledEnv, "true") + t.Setenv(SuperAgentBaseURLEnv, "https://superagent.example.test") + t.Setenv(SuperAgentOpenAPIKeyEnv, "superagent-test-open-api-key") + t.Setenv(ChatEnabledEnv, "true") + t.Setenv(ChatAuthTokenEnv, "0123456789abcdef0123456789abcdef") + t.Setenv(ChatSubjectIDEnv, " local-chat-test-subject ") + t.Setenv(ChatCompatAppIDEnv, " fire-safety-public-app ") + t.Setenv(ChatAllowedOriginsEnv, " http://localhost:5173/,https://fire.example.test,http://localhost:5173 ") + t.Setenv(ChatMaxBodyBytesEnv, "8192") + t.Setenv(ChatRunTimeoutEnv, "2m") + t.Setenv(ChatSessionTTLEnv, "45m") + t.Setenv(ChatMaxSessionsEnv, "50") + + cfg, err := Load() + if err != nil { + t.Fatalf("Load() error = %v", err) + } + if !cfg.Chat.Enabled || cfg.Chat.AuthToken == "" || cfg.Chat.SubjectID != "local-chat-test-subject" || cfg.Chat.CompatAppID != "fire-safety-public-app" { + t.Fatalf("unexpected Chat config: %#v", cfg.Chat) + } + if got, want := strings.Join(cfg.Chat.AllowedOrigins, ","), "http://localhost:5173,https://fire.example.test"; got != want { + t.Fatalf("AllowedOrigins = %q, want %q", got, want) + } + if cfg.Chat.MaxBodyBytes != 8192 || cfg.Chat.RunTimeout != 2*time.Minute || cfg.Chat.SessionTTL != 45*time.Minute || cfg.Chat.MaxSessions != 50 { + t.Fatalf("unexpected Chat limits: %#v", cfg.Chat) + } +} + +func TestLoadConfiguredMCPAndPostGIS(t *testing.T) { + clearEnvironment(t) + t.Setenv(MCPEnabledEnv, "true") + t.Setenv(MCPAuthTokenEnv, "0123456789abcdef0123456789abcdef") + t.Setenv(MCPScopeModeEnv, " TOWN_ALLOWLIST ") + t.Setenv(MCPAllowedTownsEnv, " 莒格庄镇, 高陵镇,莒格庄镇 ") + t.Setenv(MCPMaxBodyBytesEnv, "4096") + t.Setenv(MCPToolTimeoutEnv, "7s") + t.Setenv(PostGISEnabledEnv, "true") + t.Setenv(PostGISDSNEnv, "postgresql://fire:secret@127.0.0.1:5432/fire_safety?sslmode=disable") + t.Setenv(PostGISExpectedSRIDEnv, "4326") + t.Setenv(PostGISConnectTimeoutEnv, "2s") + t.Setenv(PostGISQueryTimeoutEnv, "4s") + t.Setenv(PostGISMaxConnsEnv, "6") + + cfg, err := Load() + if err != nil { + t.Fatalf("Load() error = %v", err) + } + if !cfg.MCP.Enabled || cfg.MCP.AuthToken == "" || cfg.MCP.MaxBodyBytes != 4096 || cfg.MCP.ToolTimeout != 7*time.Second { + t.Fatalf("unexpected MCP config: %#v", cfg.MCP) + } + if cfg.MCP.ScopeMode != MCPScopeModeTownAllowlist { + t.Fatalf("MCP.ScopeMode = %q, want %q", cfg.MCP.ScopeMode, MCPScopeModeTownAllowlist) + } + if got, want := strings.Join(cfg.MCP.AllowedTowns, ","), "莒格庄镇,高陵镇"; got != want { + t.Fatalf("AllowedTowns = %q, want %q", got, want) + } + if !cfg.PostGIS.Enabled || cfg.PostGIS.ExpectedSRID != 4326 || cfg.PostGIS.ConnectTimeout != 2*time.Second || cfg.PostGIS.QueryTimeout != 4*time.Second || cfg.PostGIS.MaxConns != 6 { + t.Fatalf("unexpected PostGIS config: %#v", cfg.PostGIS) + } +} + +func TestLoadConfiguredMCPAllScopeWithoutTownEnumeration(t *testing.T) { + clearEnvironment(t) + t.Setenv(MCPEnabledEnv, "true") + t.Setenv(MCPAuthTokenEnv, "0123456789abcdef0123456789abcdef") + t.Setenv(MCPScopeModeEnv, "all") + t.Setenv(PostGISEnabledEnv, "true") + t.Setenv(PostGISDSNEnv, "postgresql://fire:secret@127.0.0.1:5432/fire_safety?sslmode=disable") + t.Setenv(PostGISExpectedSRIDEnv, "4326") + + cfg, err := Load() + if err != nil { + t.Fatalf("Load() error = %v", err) + } + if cfg.MCP.ScopeMode != MCPScopeModeAll { + t.Fatalf("MCP.ScopeMode = %q, want %q", cfg.MCP.ScopeMode, MCPScopeModeAll) + } + if len(cfg.MCP.AllowedTowns) != 0 { + t.Fatalf("AllowedTowns = %#v, want empty in all scope", cfg.MCP.AllowedTowns) + } +} + +func TestLoadRejectsAllScopeWithTownAllowlist(t *testing.T) { + clearEnvironment(t) + t.Setenv(MCPEnabledEnv, "true") + t.Setenv(MCPAuthTokenEnv, "0123456789abcdef0123456789abcdef") + t.Setenv(MCPScopeModeEnv, "all") + t.Setenv(MCPAllowedTownsEnv, "莒格庄镇") + + _, err := Load() + if err == nil || !strings.Contains(err.Error(), MCPAllowedTownsEnv) { + t.Fatalf("Load() error = %v, want ambiguous scope rejection", err) + } +} + +func TestLoadConfiguredValues(t *testing.T) { + clearEnvironment(t) + t.Setenv(HTTPAddressEnv, " 127.0.0.1:9090 ") + t.Setenv(SuperAgentEnabledEnv, "true") + t.Setenv(SuperAgentBaseURLEnv, " https://superagent.example.test/root/ ") + t.Setenv(SuperAgentOpenAPIKeyEnv, " test-open-api-key ") + t.Setenv(SuperAgentConnectTimeoutEnv, "3s") + t.Setenv(SuperAgentRecoveryMaxAttemptsEnv, "2") + t.Setenv(SuperAgentRecoveryInitialBackoffEnv, "20ms") + t.Setenv(SuperAgentMaxMessageBytesEnv, "4096") + t.Setenv(SuperAgentProbeSubjectIDEnv, " probe-user ") + t.Setenv(SuperAgentProbeTimeoutEnv, "30s") + + cfg, err := Load() + if err != nil { + t.Fatalf("Load() error = %v", err) + } + if cfg.HTTPAddress != "127.0.0.1:9090" || !cfg.SuperAgent.Enabled { + t.Fatalf("unexpected base config: %#v", cfg) + } + if cfg.SuperAgent.BaseURL != "https://superagent.example.test/root" { + t.Fatalf("BaseURL = %q", cfg.SuperAgent.BaseURL) + } + if cfg.SuperAgent.OpenAPIKey != "test-open-api-key" { + t.Fatal("OpenAPIKey was not loaded") + } + if cfg.SuperAgent.ConnectTimeout != 3*time.Second || cfg.SuperAgent.RecoveryMaxAttempts != 2 || cfg.SuperAgent.RecoveryInitialBackoff != 20*time.Millisecond { + t.Fatalf("unexpected recovery config: %#v", cfg.SuperAgent) + } + if cfg.SuperAgent.MaxMessageBytes != 4096 || cfg.SuperAgent.ProbeSubjectID != "probe-user" || cfg.SuperAgent.ProbeTimeout != 30*time.Second { + t.Fatalf("unexpected probe config: %#v", cfg.SuperAgent) + } +} + +func TestLoadRejectsInvalidValues(t *testing.T) { + tests := []struct { + name string + key string + value string + }{ + {name: "boolean", key: SuperAgentEnabledEnv, value: "sometimes"}, + {name: "connect timeout syntax", key: SuperAgentConnectTimeoutEnv, value: "soon"}, + {name: "connect timeout non-positive", key: SuperAgentConnectTimeoutEnv, value: "0s"}, + {name: "recovery attempts syntax", key: SuperAgentRecoveryMaxAttemptsEnv, value: "many"}, + {name: "recovery attempts negative", key: SuperAgentRecoveryMaxAttemptsEnv, value: "-1"}, + {name: "recovery attempts too large", key: SuperAgentRecoveryMaxAttemptsEnv, value: "21"}, + {name: "recovery backoff syntax", key: SuperAgentRecoveryInitialBackoffEnv, value: "later"}, + {name: "recovery backoff non-positive", key: SuperAgentRecoveryInitialBackoffEnv, value: "0s"}, + {name: "message bytes syntax", key: SuperAgentMaxMessageBytesEnv, value: "large"}, + {name: "message bytes non-positive", key: SuperAgentMaxMessageBytesEnv, value: "0"}, + {name: "message bytes too large", key: SuperAgentMaxMessageBytesEnv, value: "16777217"}, + {name: "probe timeout syntax", key: SuperAgentProbeTimeoutEnv, value: "forever"}, + {name: "probe timeout non-positive", key: SuperAgentProbeTimeoutEnv, value: "0s"}, + {name: "Chat boolean", key: ChatEnabledEnv, value: "sometimes"}, + {name: "Chat body syntax", key: ChatMaxBodyBytesEnv, value: "large"}, + {name: "Chat body non-positive", key: ChatMaxBodyBytesEnv, value: "0"}, + {name: "Chat body too large", key: ChatMaxBodyBytesEnv, value: "1048577"}, + {name: "Chat run timeout syntax", key: ChatRunTimeoutEnv, value: "forever"}, + {name: "Chat run timeout too large", key: ChatRunTimeoutEnv, value: "31m"}, + {name: "Chat session TTL non-positive", key: ChatSessionTTLEnv, value: "0s"}, + {name: "Chat session TTL too large", key: ChatSessionTTLEnv, value: "25h"}, + {name: "Chat session capacity syntax", key: ChatMaxSessionsEnv, value: "many"}, + {name: "Chat session capacity too large", key: ChatMaxSessionsEnv, value: "10001"}, + {name: "Chat invalid origin", key: ChatAllowedOriginsEnv, value: "https://fire.example.test/path"}, + {name: "Chat wildcard origin", key: ChatAllowedOriginsEnv, value: "*"}, + {name: "Chat invalid subject", key: ChatSubjectIDEnv, value: "subject\nvalue"}, + {name: "Chat invalid compatibility app ID", key: ChatCompatAppIDEnv, value: "invalid/app/id"}, + {name: "MCP boolean", key: MCPEnabledEnv, value: "sometimes"}, + {name: "MCP scope mode", key: MCPScopeModeEnv, value: "automatic"}, + {name: "MCP body syntax", key: MCPMaxBodyBytesEnv, value: "large"}, + {name: "MCP body non-positive", key: MCPMaxBodyBytesEnv, value: "0"}, + {name: "MCP body too large", key: MCPMaxBodyBytesEnv, value: "1048577"}, + {name: "MCP timeout syntax", key: MCPToolTimeoutEnv, value: "soon"}, + {name: "MCP timeout too large", key: MCPToolTimeoutEnv, value: "31s"}, + {name: "PostGIS boolean", key: PostGISEnabledEnv, value: "sometimes"}, + {name: "PostGIS SRID syntax", key: PostGISExpectedSRIDEnv, value: "wgs84"}, + {name: "PostGIS SRID negative", key: PostGISExpectedSRIDEnv, value: "-1"}, + {name: "PostGIS connect timeout", key: PostGISConnectTimeoutEnv, value: "31s"}, + {name: "PostGIS query timeout", key: PostGISQueryTimeoutEnv, value: "0s"}, + {name: "PostGIS max conns syntax", key: PostGISMaxConnsEnv, value: "many"}, + {name: "PostGIS max conns too large", key: PostGISMaxConnsEnv, value: "21"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + clearEnvironment(t) + t.Setenv(tt.key, tt.value) + + _, err := Load() + if err == nil || !strings.Contains(err.Error(), tt.key) { + t.Fatalf("Load() error = %v, want error mentioning %s", err, tt.key) + } + }) + } +} + +func TestLoadRequiresEnabledChatSettings(t *testing.T) { + tests := []struct { + name string + configure func(*testing.T) + wantErrKey string + }{ + { + name: "missing chat token", + configure: func(t *testing.T) { + t.Setenv(SuperAgentEnabledEnv, "true") + t.Setenv(SuperAgentBaseURLEnv, "https://superagent.example.test") + t.Setenv(SuperAgentOpenAPIKeyEnv, "superagent-open-api-key") + }, + wantErrKey: ChatAuthTokenEnv, + }, + { + name: "SuperAgent disabled", + configure: func(t *testing.T) { + t.Setenv(ChatAuthTokenEnv, "0123456789abcdef0123456789abcdef") + }, + wantErrKey: SuperAgentEnabledEnv, + }, + { + name: "reused SuperAgent key", + configure: func(t *testing.T) { + shared := "0123456789abcdef0123456789abcdef" + t.Setenv(SuperAgentEnabledEnv, "true") + t.Setenv(SuperAgentBaseURLEnv, "https://superagent.example.test") + t.Setenv(SuperAgentOpenAPIKeyEnv, shared) + t.Setenv(ChatAuthTokenEnv, shared) + }, + wantErrKey: ChatAuthTokenEnv, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + clearEnvironment(t) + t.Setenv(ChatEnabledEnv, "true") + tt.configure(t) + + _, err := Load() + if err == nil || !strings.Contains(err.Error(), tt.wantErrKey) { + t.Fatalf("Load() error = %v, want error mentioning %s", err, tt.wantErrKey) + } + }) + } +} + +func TestLoadRequiresEnabledMCPSettings(t *testing.T) { + validToken := "0123456789abcdef0123456789abcdef" + validDSN := "postgresql://fire:secret@127.0.0.1:5432/fire_safety?sslmode=disable" + tests := []struct { + name string + configure func(*testing.T) + wantErrKey string + }{ + { + name: "missing token", + configure: func(t *testing.T) { + t.Setenv(MCPAllowedTownsEnv, "莒格庄镇") + }, + wantErrKey: MCPAuthTokenEnv, + }, + { + name: "short token", + configure: func(t *testing.T) { + t.Setenv(MCPAuthTokenEnv, "short-token") + t.Setenv(MCPAllowedTownsEnv, "莒格庄镇") + }, + wantErrKey: MCPAuthTokenEnv, + }, + { + name: "missing scope", + configure: func(t *testing.T) { + t.Setenv(MCPAuthTokenEnv, validToken) + }, + wantErrKey: MCPAllowedTownsEnv, + }, + { + name: "PostGIS disabled", + configure: func(t *testing.T) { + t.Setenv(MCPAuthTokenEnv, validToken) + t.Setenv(MCPAllowedTownsEnv, "莒格庄镇") + }, + wantErrKey: PostGISEnabledEnv, + }, + { + name: "missing expected SRID", + configure: func(t *testing.T) { + t.Setenv(MCPAuthTokenEnv, validToken) + t.Setenv(MCPAllowedTownsEnv, "莒格庄镇") + t.Setenv(PostGISEnabledEnv, "true") + t.Setenv(PostGISDSNEnv, validDSN) + }, + wantErrKey: PostGISExpectedSRIDEnv, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + clearEnvironment(t) + t.Setenv(MCPEnabledEnv, "true") + tt.configure(t) + + _, err := Load() + if err == nil || !strings.Contains(err.Error(), tt.wantErrKey) { + t.Fatalf("Load() error = %v, want error mentioning %s", err, tt.wantErrKey) + } + }) + } +} + +func TestLoadRejectsReusedSuperAgentAndMCPToken(t *testing.T) { + clearEnvironment(t) + sharedToken := "0123456789abcdef0123456789abcdef" + t.Setenv(SuperAgentEnabledEnv, "true") + t.Setenv(SuperAgentBaseURLEnv, "https://superagent.example.test") + t.Setenv(SuperAgentOpenAPIKeyEnv, sharedToken) + t.Setenv(MCPEnabledEnv, "true") + t.Setenv(MCPAuthTokenEnv, sharedToken) + t.Setenv(MCPAllowedTownsEnv, "莒格庄镇") + t.Setenv(PostGISEnabledEnv, "true") + t.Setenv(PostGISDSNEnv, "postgresql://fire:secret@127.0.0.1:5432/fire_safety?sslmode=disable") + t.Setenv(PostGISExpectedSRIDEnv, "4326") + + _, err := Load() + if err == nil || !strings.Contains(err.Error(), MCPAuthTokenEnv) { + t.Fatalf("Load() error = %v, want token reuse rejection", err) + } +} + +func TestLoadRejectsReusedChatAndMCPToken(t *testing.T) { + clearEnvironment(t) + sharedToken := "0123456789abcdef0123456789abcdef" + t.Setenv(SuperAgentEnabledEnv, "true") + t.Setenv(SuperAgentBaseURLEnv, "https://superagent.example.test") + t.Setenv(SuperAgentOpenAPIKeyEnv, "different-superagent-open-api-key") + t.Setenv(ChatEnabledEnv, "true") + t.Setenv(ChatAuthTokenEnv, sharedToken) + t.Setenv(MCPEnabledEnv, "true") + t.Setenv(MCPAuthTokenEnv, sharedToken) + t.Setenv(MCPScopeModeEnv, "all") + t.Setenv(PostGISEnabledEnv, "true") + t.Setenv(PostGISDSNEnv, "postgresql://fire:secret@127.0.0.1:5432/fire_safety?sslmode=disable") + t.Setenv(PostGISExpectedSRIDEnv, "4326") + + _, err := Load() + if err == nil || !strings.Contains(err.Error(), ChatAuthTokenEnv) { + t.Fatalf("Load() error = %v, want chat/MCP token reuse rejection", err) + } +} + +func TestLoadRejectsInvalidPostGISDSNWithoutLeakingIt(t *testing.T) { + clearEnvironment(t) + secret := "do-not-leak-this-password" + t.Setenv(PostGISEnabledEnv, "true") + t.Setenv(PostGISDSNEnv, "not-a-url-"+secret) + + _, err := Load() + if err == nil || !strings.Contains(err.Error(), PostGISDSNEnv) { + t.Fatalf("Load() error = %v, want error mentioning %s", err, PostGISDSNEnv) + } + if strings.Contains(err.Error(), secret) { + t.Fatal("PostGIS configuration error leaked DSN content") + } +} + +func TestLoadRequiresEnabledSuperAgentSettings(t *testing.T) { + tests := []struct { + name string + baseURL string + apiKey string + wantErrKey string + }{ + {name: "missing base URL", apiKey: "test-open-api-key", wantErrKey: SuperAgentBaseURLEnv}, + {name: "missing key", baseURL: "https://superagent.example.test", wantErrKey: SuperAgentOpenAPIKeyEnv}, + {name: "relative URL", baseURL: "/api", apiKey: "test-open-api-key", wantErrKey: SuperAgentBaseURLEnv}, + {name: "unsupported scheme", baseURL: "ftp://superagent.example.test", apiKey: "test-open-api-key", wantErrKey: SuperAgentBaseURLEnv}, + {name: "URL credentials", baseURL: "https://user:pass@superagent.example.test", apiKey: "test-open-api-key", wantErrKey: SuperAgentBaseURLEnv}, + {name: "URL query", baseURL: "https://superagent.example.test?key=value", apiKey: "test-open-api-key", wantErrKey: SuperAgentBaseURLEnv}, + {name: "invalid key header", baseURL: "https://superagent.example.test", apiKey: "test open api key", wantErrKey: SuperAgentOpenAPIKeyEnv}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + clearEnvironment(t) + t.Setenv(SuperAgentEnabledEnv, "true") + t.Setenv(SuperAgentBaseURLEnv, tt.baseURL) + t.Setenv(SuperAgentOpenAPIKeyEnv, tt.apiKey) + + _, err := Load() + if err == nil || !strings.Contains(err.Error(), tt.wantErrKey) { + t.Fatalf("Load() error = %v, want error mentioning %s", err, tt.wantErrKey) + } + }) + } +} + +func clearEnvironment(t *testing.T) { + t.Helper() + for _, key := range []string{ + HTTPAddressEnv, + SuperAgentEnabledEnv, + SuperAgentBaseURLEnv, + SuperAgentOpenAPIKeyEnv, + SuperAgentConnectTimeoutEnv, + SuperAgentRecoveryMaxAttemptsEnv, + SuperAgentRecoveryInitialBackoffEnv, + SuperAgentMaxMessageBytesEnv, + SuperAgentProbeSubjectIDEnv, + SuperAgentProbeTimeoutEnv, + ChatEnabledEnv, + ChatAuthTokenEnv, + ChatSubjectIDEnv, + ChatAllowedOriginsEnv, + ChatMaxBodyBytesEnv, + ChatRunTimeoutEnv, + ChatSessionTTLEnv, + ChatMaxSessionsEnv, + ChatCompatAppIDEnv, + MCPEnabledEnv, + MCPAuthTokenEnv, + MCPScopeModeEnv, + MCPAllowedTownsEnv, + MCPMaxBodyBytesEnv, + MCPToolTimeoutEnv, + PostGISEnabledEnv, + PostGISDSNEnv, + PostGISExpectedSRIDEnv, + PostGISConnectTimeoutEnv, + PostGISQueryTimeoutEnv, + PostGISMaxConnsEnv, + } { + t.Setenv(key, "") + } +} diff --git a/internal/domain/doc.go b/internal/domain/doc.go new file mode 100644 index 0000000..d6d3a1f --- /dev/null +++ b/internal/domain/doc.go @@ -0,0 +1,2 @@ +// Package domain defines stable forest-fire-safety concepts and business rules. +package domain diff --git a/internal/domain/spatial.go b/internal/domain/spatial.go new file mode 100644 index 0000000..51db2a2 --- /dev/null +++ b/internal/domain/spatial.go @@ -0,0 +1,112 @@ +package domain + +// Coordinate is a WGS84 longitude/latitude pair. +type Coordinate struct { + Longitude float64 `json:"longitude"` + Latitude float64 `json:"latitude"` +} + +// SpatialScope is a trusted, server-derived data scope. +type SpatialScope struct { + // AllTowns authorizes every town in the MCP's fixed query tables, including records with no town value. + AllTowns bool + // AllowedTowns is used only when AllTowns is false. + AllowedTowns []string +} + +// NearbyQuery contains a bounded spatial search requested by a service use case. +type NearbyQuery struct { + Point Coordinate + RadiusMeters float64 + Limit int + Scope SpatialScope +} + +// PlaceSearchQuery searches existing fire-safety records for user-confirmable place candidates. +type PlaceSearchQuery struct { + PlaceName string + Limit int + Scope SpatialScope +} + +// PlaceCandidate is a recorded or representative location that still requires user confirmation. +type PlaceCandidate struct { + SourceRecordID string `json:"source_record_id"` + PlaceType string `json:"place_type"` + Name string `json:"name,omitempty"` + Town string `json:"town,omitempty"` + Village string `json:"village,omitempty"` + MatchedField string `json:"matched_field"` + MatchedText string `json:"matched_text"` + MatchKind string `json:"match_kind"` + Location Coordinate `json:"location"` + LocationKind string `json:"location_kind"` +} + +// IncidentContext is a fire-grid area covering the supplied incident point. +type IncidentContext struct { + GridID string `json:"grid_id"` + Town string `json:"town,omitempty"` + AreaLabel string `json:"area_label,omitempty"` +} + +// WaterSource is a spatially nearby water candidate from an existing source table. +type WaterSource struct { + SourceRecordID string `json:"source_record_id"` + Category string `json:"category"` + Name string `json:"name,omitempty"` + Town string `json:"town,omitempty"` + Village string `json:"village,omitempty"` + Location Coordinate `json:"location"` + DistanceMeters float64 `json:"distance_meters"` + CapacityCubicMeters *float64 `json:"capacity_cubic_meters,omitempty"` + ResourceType string `json:"resource_type,omitempty"` + ReportedStatus string `json:"reported_status,omitempty"` + SourceTimestampRaw *int64 `json:"source_timestamp_raw,omitempty"` +} + +// CommandPostCandidate is a nearby existing facility, not a confirmed command post. +type CommandPostCandidate struct { + SourceRecordID string `json:"source_record_id"` + FacilityType string `json:"facility_type"` + Name string `json:"name,omitempty"` + Town string `json:"town,omitempty"` + Village string `json:"village,omitempty"` + Location Coordinate `json:"location"` + DistanceMeters float64 `json:"distance_meters"` + ReportedStatus string `json:"reported_status,omitempty"` + ManagementUnit string `json:"management_unit,omitempty"` +} + +// AccessLine is a nearby fire-access-line candidate and its closest point. +type AccessLine struct { + SourceRecordID string `json:"source_record_id"` + Name string `json:"name,omitempty"` + Town string `json:"town,omitempty"` + DistanceMeters float64 `json:"distance_meters"` + NearestPoint Coordinate `json:"nearest_point"` + LengthMeters *float64 `json:"length_meters,omitempty"` + SourceUpdatedRaw string `json:"source_updated_raw,omitempty"` +} + +// ResponsibleUnit is the non-personal responsibility information stored on a fire grid. +type ResponsibleUnit struct { + GridID string `json:"grid_id"` + Town string `json:"town,omitempty"` + AreaLabel string `json:"area_label,omitempty"` + FireTeam string `json:"fire_team,omitempty"` + LiveLocationAvailable bool `json:"live_location_available"` + AssemblySiteAvailable bool `json:"assembly_site_available"` +} + +// RiskArea is a cemetery or forest-area enterprise polygon near the supplied point. +type RiskArea struct { + SourceRecordID string `json:"source_record_id"` + RiskType string `json:"risk_type"` + Name string `json:"name,omitempty"` + Town string `json:"town,omitempty"` + Village string `json:"village,omitempty"` + Direction string `json:"direction,omitempty"` + CoversPoint bool `json:"covers_point"` + DistanceMeters float64 `json:"distance_meters"` +} diff --git a/internal/handler/chat.go b/internal/handler/chat.go new file mode 100644 index 0000000..3f56838 --- /dev/null +++ b/internal/handler/chat.go @@ -0,0 +1,466 @@ +package handler + +import ( + "bytes" + "context" + "crypto/subtle" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "mime" + "net/http" + "net/url" + "strconv" + "strings" + "time" + + "fire-safety-ymd/internal/service" +) + +const ( + maximumChatBodyBytes = int64(1024 * 1024) + maximumChatRunTime = 30 * time.Minute + chatHeartbeatInterval = 15 * time.Second +) + +// ChatUseCase is the inbound chat capability consumed by ChatHandler. +type ChatUseCase interface { + Prepare(context.Context, service.ChatRequest) (service.ChatTurn, error) +} + +// ChatOptions contains the independently authenticated user-facing transport settings. +type ChatOptions struct { + AuthToken string + AllowedOrigins []string + MaxBodyBytes int64 + RunTimeout time.Duration + Logger *log.Logger +} + +// ChatHandler exposes a bounded SSE chat API without exposing provider credentials or sessions. +type ChatHandler struct { + chat ChatUseCase + authToken string + allowedOrigins map[string]struct{} + maxBodyBytes int64 + runTimeout time.Duration + logger *log.Logger +} + +// NewChatHandler constructs the protected chat transport. +func NewChatHandler(chat ChatUseCase, options ChatOptions) (*ChatHandler, error) { + if chat == nil { + return nil, errors.New("chat service is required") + } + if !validChatToken(options.AuthToken) { + return nil, errors.New("chat auth token must contain at least 32 printable ASCII characters") + } + if options.MaxBodyBytes <= 0 || options.MaxBodyBytes > maximumChatBodyBytes { + return nil, errors.New("chat max body bytes is outside the supported range") + } + if options.RunTimeout <= 0 || options.RunTimeout > maximumChatRunTime { + return nil, errors.New("chat run timeout is outside the supported range") + } + allowedOrigins := make(map[string]struct{}, len(options.AllowedOrigins)) + for _, origin := range options.AllowedOrigins { + normalized, ok := validChatOrigin(origin) + if !ok || normalized != origin { + return nil, errors.New("chat allowed origins must be normalized exact HTTP(S) origins") + } + allowedOrigins[origin] = struct{}{} + } + logger := options.Logger + if logger == nil { + logger = log.New(io.Discard, "", 0) + } + return &ChatHandler{ + chat: chat, + authToken: options.AuthToken, + allowedOrigins: allowedOrigins, + maxBodyBytes: options.MaxBodyBytes, + runTimeout: options.RunTimeout, + logger: logger, + }, nil +} + +// ServeHTTP handles POST chat turns and browser CORS preflight. +func (h *ChatHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + started := time.Now() + requestID := safeRequestID(r.Header.Get("X-Request-ID")) + w.Header().Set("X-Request-ID", requestID) + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("X-Content-Type-Options", "nosniff") + + if r.Method == http.MethodOptions { + h.handlePreflight(w, r, requestID, started) + return + } + if r.Method != http.MethodPost { + w.Header().Set("Allow", http.MethodPost+", "+http.MethodOptions) + h.writeError(w, http.StatusMethodNotAllowed, requestID, "CHAT_METHOD_NOT_ALLOWED", "Chat endpoint accepts POST requests only.") + h.logResult(requestID, "method_not_allowed", false, started) + return + } + if !h.authorizeOrigin(w, r.Header.Get("Origin")) { + h.writeError(w, http.StatusForbidden, requestID, "CHAT_ORIGIN_FORBIDDEN", "Chat browser origin is not allowed.") + h.logResult(requestID, "forbidden_origin", false, started) + return + } + if !h.validAuthorization(r.Header.Get("Authorization")) { + w.Header().Set("WWW-Authenticate", `Bearer realm="fire-safety-ymd-chat"`) + h.writeError(w, http.StatusUnauthorized, requestID, "CHAT_AUTH_INVALID", "Chat authentication failed.") + h.logResult(requestID, "unauthorized", false, started) + return + } + if !isJSONContentType(r.Header.Get("Content-Type")) { + h.writeError(w, http.StatusUnsupportedMediaType, requestID, "CHAT_CONTENT_TYPE_INVALID", "Content-Type must be application/json.") + h.logResult(requestID, "unsupported_media_type", false, started) + return + } + if !acceptsEventStream(r.Header.Get("Accept")) { + h.writeError(w, http.StatusNotAcceptable, requestID, "CHAT_ACCEPT_INVALID", "Accept must allow text/event-stream.") + h.logResult(requestID, "not_acceptable", false, started) + return + } + + body, tooLarge, err := readBoundedBody(r.Body, h.maxBodyBytes) + if err != nil { + h.writeError(w, http.StatusBadRequest, requestID, "CHAT_REQUEST_INVALID", "Unable to read chat request.") + h.logResult(requestID, "read_error", false, started) + return + } + if tooLarge { + h.writeError(w, http.StatusRequestEntityTooLarge, requestID, "CHAT_REQUEST_BODY_TOO_LARGE", "Chat request body exceeds the configured limit.") + h.logResult(requestID, "body_too_large", false, started) + return + } + request, err := decodeChatRequest(body) + if err != nil { + h.writeError(w, http.StatusBadRequest, requestID, "CHAT_REQUEST_INVALID", "Chat request JSON is invalid.") + h.logResult(requestID, "invalid_request", false, started) + return + } + + runCtx, cancelRun := context.WithTimeout(r.Context(), h.runTimeout) + defer cancelRun() + turn, err := h.chat.Prepare(runCtx, service.ChatRequest{ + Message: request.Message, + ConversationID: request.ConversationID, + RequestID: requestID, + }) + if err != nil { + status, code, message := publicChatError(err) + if errors.Is(err, service.ErrChatConversationBusy) { + w.Header().Set("Retry-After", "1") + } + if errors.Is(err, service.ErrChatCapacityReached) { + w.Header().Set("Retry-After", "30") + } + h.writeError(w, status, requestID, code, message) + h.logResult(requestID, code, request.ConversationID != "", started) + return + } + defer turn.Close() + + flusher, ok := w.(http.Flusher) + if !ok { + h.writeError(w, http.StatusInternalServerError, requestID, "CHAT_STREAM_UNSUPPORTED", "Chat streaming is unavailable.") + h.logResult(requestID, "stream_unsupported", turn.Reused(), started) + return + } + if deadline, ok := runCtx.Deadline(); ok { + _ = http.NewResponseController(w).SetWriteDeadline(deadline.Add(5 * time.Second)) + } + w.Header().Set("Content-Type", "text/event-stream; charset=utf-8") + w.Header().Set("X-Accel-Buffering", "no") + w.WriteHeader(http.StatusOK) + if err := writeChatSSE(w, flusher, "conversation", map[string]any{ + "conversation_id": turn.ConversationID(), + "reused": turn.Reused(), + }); err != nil { + h.logResult(requestID, "client_write_failed", turn.Reused(), started) + return + } + + streamCtx, cancelStream := context.WithCancel(runCtx) + defer cancelStream() + type streamOutcome struct { + result service.ChatResult + err error + } + progress := make(chan service.AgentTraceEvent) + outcome := make(chan streamOutcome, 1) + go func() { + result, err := turn.Stream(streamCtx, func(event service.AgentTraceEvent) { + if !exposeProgressEvent(event.Event) { + return + } + select { + case progress <- event: + case <-streamCtx.Done(): + } + }) + outcome <- streamOutcome{result: result, err: err} + }() + + heartbeat := time.NewTicker(chatHeartbeatInterval) + defer heartbeat.Stop() + contextDone := runCtx.Done() + var result service.ChatResult + var streamErr error +streamLoop: + for { + select { + case event := <-progress: + if err := writeChatSSE(w, flusher, "progress", event); err != nil { + cancelStream() + h.logResult(requestID, "client_write_failed", turn.Reused(), started) + return + } + case completed := <-outcome: + result = completed.result + streamErr = completed.err + break streamLoop + case <-heartbeat.C: + if err := writeChatHeartbeat(w, flusher); err != nil { + cancelStream() + h.logResult(requestID, "client_write_failed", turn.Reused(), started) + return + } + case <-contextDone: + cancelStream() + contextDone = nil + } + } + if streamErr != nil { + _, code, message := publicChatError(streamErr) + _ = writeChatSSE(w, flusher, "error", map[string]string{ + "code": code, + "message": message, + "conversation_id": turn.ConversationID(), + "request_id": requestID, + }) + h.logResult(requestID, code, turn.Reused(), started) + return + } + if err := writeChatSSE(w, flusher, "message", map[string]string{ + "conversation_id": result.ConversationID, + "answer": result.Answer, + }); err != nil { + h.logResult(requestID, "client_write_failed", turn.Reused(), started) + return + } + if err := writeChatSSE(w, flusher, "done", map[string]any{ + "conversation_id": result.ConversationID, + "run_id": result.RunID, + "usage": result.Usage, + }); err != nil { + h.logResult(requestID, "client_write_failed", turn.Reused(), started) + return + } + h.logResult(requestID, "success", turn.Reused(), started) +} + +type chatRequestPayload struct { + Message string + ConversationID string +} + +func decodeChatRequest(body []byte) (chatRequestPayload, error) { + var payload struct { + Message *string `json:"message"` + ConversationID json.RawMessage `json:"conversation_id"` + } + trimmed := bytes.TrimSpace(body) + if len(trimmed) == 0 || trimmed[0] != '{' { + return chatRequestPayload{}, errors.New("chat request must be an object") + } + decoder := json.NewDecoder(bytes.NewReader(trimmed)) + decoder.DisallowUnknownFields() + if err := decoder.Decode(&payload); err != nil || payload.Message == nil { + return chatRequestPayload{}, errors.New("chat request is invalid") + } + if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + return chatRequestPayload{}, errors.New("chat request contains trailing JSON") + } + request := chatRequestPayload{Message: *payload.Message} + if len(payload.ConversationID) > 0 { + if err := json.Unmarshal(payload.ConversationID, &request.ConversationID); err != nil || strings.TrimSpace(request.ConversationID) == "" { + return chatRequestPayload{}, errors.New("conversation_id must be a non-empty string when provided") + } + } + return request, nil +} + +func (h *ChatHandler) handlePreflight(w http.ResponseWriter, r *http.Request, requestID string, started time.Time) { + origin := r.Header.Get("Origin") + if origin == "" || !h.authorizeOrigin(w, origin) || r.Header.Get("Access-Control-Request-Method") != http.MethodPost || + !validPreflightHeaders(r.Header.Get("Access-Control-Request-Headers")) { + h.writeError(w, http.StatusForbidden, requestID, "CHAT_ORIGIN_FORBIDDEN", "Chat browser origin or preflight request is not allowed.") + h.logResult(requestID, "preflight_forbidden", false, started) + return + } + w.Header().Set("Access-Control-Allow-Methods", http.MethodPost) + w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type") + w.Header().Set("Access-Control-Max-Age", "600") + w.WriteHeader(http.StatusNoContent) + h.logResult(requestID, "preflight_success", false, started) +} + +func (h *ChatHandler) authorizeOrigin(w http.ResponseWriter, origin string) bool { + if origin == "" { + return true + } + if _, allowed := h.allowedOrigins[origin]; !allowed { + return false + } + w.Header().Add("Vary", "Origin") + w.Header().Set("Access-Control-Allow-Origin", origin) + w.Header().Set("Access-Control-Expose-Headers", "X-Request-ID") + return true +} + +func (h *ChatHandler) validAuthorization(value string) bool { + expected := "Bearer " + h.authToken + if len(value) != len(expected) { + return false + } + return subtle.ConstantTimeCompare([]byte(value), []byte(expected)) == 1 +} + +func (h *ChatHandler) writeError(w http.ResponseWriter, status int, requestID, code, message string) { + writeJSON(w, status, map[string]any{ + "error": map[string]string{ + "code": code, + "message": message, + }, + "request_id": requestID, + }) +} + +func (h *ChatHandler) logResult(requestID, result string, reused bool, started time.Time) { + h.logger.Printf("chat_request request_id=%s result=%s reused=%t duration_ms=%d", requestID, result, reused, time.Since(started).Milliseconds()) +} + +func writeChatSSE(w io.Writer, flusher http.Flusher, event string, payload any) error { + data, err := json.Marshal(payload) + if err != nil { + return err + } + if _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", event, data); err != nil { + return err + } + flusher.Flush() + return nil +} + +func writeChatHeartbeat(w io.Writer, flusher http.Flusher) error { + if _, err := io.WriteString(w, ": keepalive\n\n"); err != nil { + return err + } + flusher.Flush() + return nil +} + +func publicChatError(err error) (int, string, string) { + switch { + case errors.Is(err, service.ErrChatInvalidArgument): + return http.StatusBadRequest, "CHAT_REQUEST_INVALID", "Chat request is invalid." + case errors.Is(err, service.ErrChatConversationNotFound): + return http.StatusNotFound, "CHAT_CONVERSATION_NOT_FOUND", "Chat conversation was not found or has expired." + case errors.Is(err, service.ErrChatConversationBusy): + return http.StatusConflict, "CHAT_CONVERSATION_BUSY", "Chat conversation already has an active run." + case errors.Is(err, service.ErrChatCapacityReached): + return http.StatusServiceUnavailable, "CHAT_CAPACITY_REACHED", "Chat session capacity is temporarily exhausted." + case errors.Is(err, context.DeadlineExceeded): + return http.StatusGatewayTimeout, "CHAT_UPSTREAM_TIMEOUT", "Chat run did not complete within its time limit." + case errors.Is(err, context.Canceled): + return http.StatusRequestTimeout, "CHAT_REQUEST_CANCELED", "Chat request was canceled." + case errors.Is(err, service.ErrChatRunFailed): + return http.StatusBadGateway, "CHAT_RUN_FAILED", "SuperAgent reported a failed run." + case errors.Is(err, service.ErrChatUpstreamProtocol): + return http.StatusBadGateway, "CHAT_UPSTREAM_PROTOCOL_ERROR", "SuperAgent returned an incomplete or invalid response." + case errors.Is(err, service.ErrChatUpstreamUnavailable): + return http.StatusBadGateway, "CHAT_UPSTREAM_UNAVAILABLE", "SuperAgent is unavailable." + default: + return http.StatusInternalServerError, "CHAT_INTERNAL_ERROR", "Chat request failed." + } +} + +func validChatToken(value string) bool { + if len(value) < 32 || len(value) > 4096 { + return false + } + for _, character := range value { + if character <= 0x20 || character >= 0x7f { + return false + } + } + return true +} + +func validChatOrigin(value string) (string, bool) { + parsed, err := url.Parse(value) + if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") || + parsed.User != nil || parsed.RawQuery != "" || parsed.ForceQuery || parsed.Fragment != "" || + (parsed.Path != "" && parsed.Path != "/") || value == "*" { + return "", false + } + return parsed.Scheme + "://" + parsed.Host, true +} + +func acceptsEventStream(value string) bool { + if strings.TrimSpace(value) == "" { + return true + } + for _, candidate := range strings.Split(value, ",") { + mediaType, parameters, err := mime.ParseMediaType(strings.TrimSpace(candidate)) + if err != nil { + continue + } + if quality, exists := parameters["q"]; exists { + parsed, parseErr := strconv.ParseFloat(quality, 64) + if parseErr != nil || parsed <= 0 || parsed > 1 { + continue + } + } + if mediaType == "*/*" || mediaType == "text/*" || strings.EqualFold(mediaType, "text/event-stream") { + return true + } + } + return false +} + +func validPreflightHeaders(value string) bool { + for _, header := range strings.Split(value, ",") { + header = strings.ToLower(strings.TrimSpace(header)) + if header == "" { + continue + } + if header != "authorization" && header != "content-type" { + return false + } + } + return true +} + +func exposeProgressEvent(event string) bool { + event = strings.TrimSpace(event) + if !strings.HasPrefix(event, "tool.") && !strings.HasPrefix(event, "run.") { + return false + } + if len(event) > 128 { + return false + } + for _, character := range event { + if !(character >= 'a' && character <= 'z') && + !(character >= 'A' && character <= 'Z') && + !(character >= '0' && character <= '9') && + !strings.ContainsRune("._-", character) { + return false + } + } + return true +} diff --git a/internal/handler/chat_test.go b/internal/handler/chat_test.go new file mode 100644 index 0000000..5809c8c --- /dev/null +++ b/internal/handler/chat_test.go @@ -0,0 +1,273 @@ +package handler + +import ( + "bytes" + "context" + "fmt" + "log" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "fire-safety-ymd/internal/service" +) + +const testChatToken = "abcdef0123456789abcdef0123456789" + +func TestNewChatHandlerRejectsUnsafeOptions(t *testing.T) { + tests := []ChatOptions{ + {AuthToken: "short", MaxBodyBytes: 1, RunTimeout: time.Second}, + {AuthToken: testChatToken, MaxBodyBytes: maximumChatBodyBytes + 1, RunTimeout: time.Second}, + {AuthToken: testChatToken, MaxBodyBytes: 1, RunTimeout: maximumChatRunTime + time.Second}, + {AuthToken: testChatToken, AllowedOrigins: []string{"*"}, MaxBodyBytes: 1, RunTimeout: time.Second}, + {AuthToken: testChatToken, AllowedOrigins: []string{"https://fire.example.test/"}, MaxBodyBytes: 1, RunTimeout: time.Second}, + } + for _, options := range tests { + if _, err := NewChatHandler(&fakeChatUseCase{}, options); err == nil { + t.Fatalf("NewChatHandler(%#v) error = nil", options) + } + } +} + +func TestChatTransportGuardsRunBeforeService(t *testing.T) { + tests := []struct { + name string + method string + body string + configure func(*http.Request) + maxBody int64 + wantStatus int + wantCode string + }{ + {name: "wrong method", method: http.MethodGet, body: `{}`, wantStatus: http.StatusMethodNotAllowed, wantCode: "CHAT_METHOD_NOT_ALLOWED"}, + {name: "forbidden origin", method: http.MethodPost, body: `{"message":"hello"}`, configure: func(r *http.Request) { r.Header.Set("Origin", "https://evil.example") }, wantStatus: http.StatusForbidden, wantCode: "CHAT_ORIGIN_FORBIDDEN"}, + {name: "missing auth", method: http.MethodPost, body: `{"message":"hello"}`, configure: func(r *http.Request) { r.Header.Del("Authorization") }, wantStatus: http.StatusUnauthorized, wantCode: "CHAT_AUTH_INVALID"}, + {name: "wrong content type", method: http.MethodPost, body: `{"message":"hello"}`, configure: func(r *http.Request) { r.Header.Set("Content-Type", "text/plain") }, wantStatus: http.StatusUnsupportedMediaType, wantCode: "CHAT_CONTENT_TYPE_INVALID"}, + {name: "wrong accept", method: http.MethodPost, body: `{"message":"hello"}`, configure: func(r *http.Request) { r.Header.Set("Accept", "application/json") }, wantStatus: http.StatusNotAcceptable, wantCode: "CHAT_ACCEPT_INVALID"}, + {name: "event stream explicitly rejected", method: http.MethodPost, body: `{"message":"hello"}`, configure: func(r *http.Request) { r.Header.Set("Accept", "text/event-stream;q=0") }, wantStatus: http.StatusNotAcceptable, wantCode: "CHAT_ACCEPT_INVALID"}, + {name: "invalid JSON", method: http.MethodPost, body: `{"message":`, wantStatus: http.StatusBadRequest, wantCode: "CHAT_REQUEST_INVALID"}, + {name: "unknown field", method: http.MethodPost, body: `{"message":"hello","user_id":"admin"}`, wantStatus: http.StatusBadRequest, wantCode: "CHAT_REQUEST_INVALID"}, + {name: "missing message", method: http.MethodPost, body: `{}`, wantStatus: http.StatusBadRequest, wantCode: "CHAT_REQUEST_INVALID"}, + {name: "null conversation", method: http.MethodPost, body: `{"message":"hello","conversation_id":null}`, wantStatus: http.StatusBadRequest, wantCode: "CHAT_REQUEST_INVALID"}, + {name: "empty conversation", method: http.MethodPost, body: `{"message":"hello","conversation_id":""}`, wantStatus: http.StatusBadRequest, wantCode: "CHAT_REQUEST_INVALID"}, + {name: "body too large", method: http.MethodPost, body: `{"message":"hello"}`, maxBody: 4, wantStatus: http.StatusRequestEntityTooLarge, wantCode: "CHAT_REQUEST_BODY_TOO_LARGE"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + chat := &fakeChatUseCase{} + maxBody := tt.maxBody + if maxBody == 0 { + maxBody = 4096 + } + handler := newTestChatHandler(t, chat, maxBody, []string{"https://allowed.example"}, nil) + request := authenticatedChatRequest(tt.method, tt.body) + if tt.configure != nil { + tt.configure(request) + } + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != tt.wantStatus || !strings.Contains(response.Body.String(), `"code":"`+tt.wantCode+`"`) { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } + if chat.prepareCalls != 0 { + t.Fatalf("Prepare calls = %d, want 0", chat.prepareCalls) + } + }) + } +} + +func TestChatCORSPreflight(t *testing.T) { + handler := newTestChatHandler(t, &fakeChatUseCase{}, 4096, []string{"http://localhost:5173"}, nil) + + allowed := httptest.NewRequest(http.MethodOptions, "/api/chat", nil) + allowed.Header.Set("Origin", "http://localhost:5173") + allowed.Header.Set("Access-Control-Request-Method", http.MethodPost) + allowed.Header.Set("Access-Control-Request-Headers", "authorization, content-type") + allowedResponse := httptest.NewRecorder() + handler.ServeHTTP(allowedResponse, allowed) + if allowedResponse.Code != http.StatusNoContent || allowedResponse.Header().Get("Access-Control-Allow-Origin") != "http://localhost:5173" { + t.Fatalf("allowed preflight status=%d headers=%#v", allowedResponse.Code, allowedResponse.Header()) + } + + forbidden := httptest.NewRequest(http.MethodOptions, "/api/chat", nil) + forbidden.Header.Set("Origin", "http://localhost:5173") + forbidden.Header.Set("Access-Control-Request-Method", http.MethodPost) + forbidden.Header.Set("Access-Control-Request-Headers", "x-admin-role") + forbiddenResponse := httptest.NewRecorder() + handler.ServeHTTP(forbiddenResponse, forbidden) + if forbiddenResponse.Code != http.StatusForbidden || !strings.Contains(forbiddenResponse.Body.String(), "CHAT_ORIGIN_FORBIDDEN") { + t.Fatalf("forbidden preflight status=%d body=%s", forbiddenResponse.Code, forbiddenResponse.Body.String()) + } +} + +func TestChatStreamsSafeProgressAndFinalAnswer(t *testing.T) { + turn := &fakeChatTurn{ + conversationID: "conv_0123456789abcdef01234567", + events: []service.AgentTraceEvent{ + {Event: "message.delta", Status: "running"}, + {Event: "tool.started", ToolName: "fire_safety_find_nearby_water_sources", Status: "running"}, + }, + result: service.ChatResult{ + ConversationID: "conv_0123456789abcdef01234567", + RunID: "run-safe-1", + Answer: "候选水源需要现场确认。", + Usage: service.ChatTokenUsage{Input: 10, Output: 20, Total: 30}, + }, + } + chat := &fakeChatUseCase{turn: turn} + var logs bytes.Buffer + handler := newTestChatHandler(t, chat, 4096, []string{"https://allowed.example"}, log.New(&logs, "", 0)) + request := authenticatedChatRequest(http.MethodPost, `{"message":"不要记录这个问题 secret-user-text"}`) + request.Header.Set("Origin", "https://allowed.example") + request.Header.Set("X-Request-ID", "chat-request-1") + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusOK || !strings.HasPrefix(response.Header().Get("Content-Type"), "text/event-stream") { + t.Fatalf("status=%d content-type=%q body=%s", response.Code, response.Header().Get("Content-Type"), response.Body.String()) + } + body := response.Body.String() + for _, expected := range []string{"event: conversation", "event: progress", "event: message", "event: done", "候选水源需要现场确认。", `"reused":false`} { + if !strings.Contains(body, expected) { + t.Fatalf("SSE body missing %q: %s", expected, body) + } + } + if strings.Contains(body, "message.delta") || strings.Contains(body, "provider-session") { + t.Fatalf("SSE body exposed suppressed provider details: %s", body) + } + if !(strings.Index(body, "event: conversation") < strings.Index(body, "event: progress") && + strings.Index(body, "event: progress") < strings.Index(body, "event: message") && + strings.Index(body, "event: message") < strings.Index(body, "event: done")) { + t.Fatalf("unexpected SSE order: %s", body) + } + if chat.lastRequest.Message != "不要记录这个问题 secret-user-text" || chat.lastRequest.RequestID != "chat-request-1" { + t.Fatalf("service request = %#v", chat.lastRequest) + } + if !turn.closed { + t.Fatal("turn was not closed") + } + if strings.Contains(logs.String(), "secret-user-text") || strings.Contains(logs.String(), "候选水源") { + t.Fatalf("logs contain chat content: %s", logs.String()) + } +} + +func TestChatStreamFailureDoesNotExposePartialAnswer(t *testing.T) { + turn := &fakeChatTurn{ + conversationID: "conv_0123456789abcdef01234567", + events: []service.AgentTraceEvent{{Event: "tool.failed", ToolName: "safe_tool", Status: "failed"}}, + err: fmt.Errorf("provider secret response: %w", service.ErrChatUpstreamProtocol), + } + handler := newTestChatHandler(t, &fakeChatUseCase{turn: turn}, 4096, nil, nil) + request := authenticatedChatRequest(http.MethodPost, `{"message":"hello"}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + body := response.Body.String() + if response.Code != http.StatusOK || !strings.Contains(body, "event: error") || !strings.Contains(body, "CHAT_UPSTREAM_PROTOCOL_ERROR") { + t.Fatalf("status=%d body=%s", response.Code, body) + } + if strings.Contains(body, "event: message") || strings.Contains(body, "event: done") || strings.Contains(body, "provider secret response") { + t.Fatalf("failure stream exposed final/secret data: %s", body) + } +} + +func TestChatMapsPreparationErrorsBeforeStartingSSE(t *testing.T) { + tests := []struct { + err error + wantStatus int + wantCode string + }{ + {err: service.ErrChatInvalidArgument, wantStatus: http.StatusBadRequest, wantCode: "CHAT_REQUEST_INVALID"}, + {err: service.ErrChatConversationNotFound, wantStatus: http.StatusNotFound, wantCode: "CHAT_CONVERSATION_NOT_FOUND"}, + {err: service.ErrChatConversationBusy, wantStatus: http.StatusConflict, wantCode: "CHAT_CONVERSATION_BUSY"}, + {err: service.ErrChatCapacityReached, wantStatus: http.StatusServiceUnavailable, wantCode: "CHAT_CAPACITY_REACHED"}, + {err: context.DeadlineExceeded, wantStatus: http.StatusGatewayTimeout, wantCode: "CHAT_UPSTREAM_TIMEOUT"}, + {err: service.ErrChatUpstreamUnavailable, wantStatus: http.StatusBadGateway, wantCode: "CHAT_UPSTREAM_UNAVAILABLE"}, + } + for _, tt := range tests { + t.Run(tt.wantCode, func(t *testing.T) { + handler := newTestChatHandler(t, &fakeChatUseCase{err: tt.err}, 4096, nil, nil) + request := authenticatedChatRequest(http.MethodPost, `{"message":"hello"}`) + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != tt.wantStatus || !strings.Contains(response.Body.String(), tt.wantCode) || strings.Contains(response.Header().Get("Content-Type"), "text/event-stream") { + t.Fatalf("status=%d content-type=%q body=%s", response.Code, response.Header().Get("Content-Type"), response.Body.String()) + } + }) + } +} + +func newTestChatHandler(t *testing.T, chat ChatUseCase, maxBody int64, origins []string, logger *log.Logger) *ChatHandler { + t.Helper() + handler, err := NewChatHandler(chat, ChatOptions{ + AuthToken: testChatToken, + AllowedOrigins: origins, + MaxBodyBytes: maxBody, + RunTimeout: time.Minute, + Logger: logger, + }) + if err != nil { + t.Fatalf("NewChatHandler() error = %v", err) + } + return handler +} + +func authenticatedChatRequest(method, body string) *http.Request { + request := httptest.NewRequest(method, "/api/chat", strings.NewReader(body)) + request.Header.Set("Authorization", "Bearer "+testChatToken) + request.Header.Set("Content-Type", "application/json") + request.Header.Set("Accept", "text/event-stream") + return request +} + +type fakeChatUseCase struct { + turn service.ChatTurn + err error + prepareCalls int + lastRequest service.ChatRequest +} + +func (f *fakeChatUseCase) Prepare(_ context.Context, request service.ChatRequest) (service.ChatTurn, error) { + f.prepareCalls++ + f.lastRequest = request + if f.err != nil { + return nil, f.err + } + if f.turn == nil { + return &fakeChatTurn{conversationID: "conv_0123456789abcdef01234567", result: service.ChatResult{ + ConversationID: "conv_0123456789abcdef01234567", + Answer: "ok", + }}, nil + } + return f.turn, nil +} + +type fakeChatTurn struct { + conversationID string + reused bool + events []service.AgentTraceEvent + result service.ChatResult + err error + closed bool +} + +func (f *fakeChatTurn) ConversationID() string { return f.conversationID } + +func (f *fakeChatTurn) Reused() bool { return f.reused } + +func (f *fakeChatTurn) Stream(_ context.Context, trace func(service.AgentTraceEvent)) (service.ChatResult, error) { + for _, event := range f.events { + if trace != nil { + trace(event) + } + } + return f.result, f.err +} + +func (f *fakeChatTurn) Close() { f.closed = true } diff --git a/internal/handler/dashscope_chat.go b/internal/handler/dashscope_chat.go new file mode 100644 index 0000000..920744e --- /dev/null +++ b/internal/handler/dashscope_chat.go @@ -0,0 +1,429 @@ +package handler + +import ( + "bytes" + "context" + "crypto/subtle" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "net/http" + "strings" + "time" + + "fire-safety-ymd/internal/service" +) + +// DashScopeChatOptions configures the public DashScope-compatible chat +// transport. Its token authenticates callers of this service and is never a +// provider credential. +type DashScopeChatOptions struct { + AppID string + AuthToken string + AllowedOrigins []string + MaxBodyBytes int64 + RunTimeout time.Duration + Logger *log.Logger +} + +// DashScopeChatHandler adapts the provider-neutral ChatUseCase to the narrow +// DashScope application completion contract used by existing callers. +type DashScopeChatHandler struct { + chat ChatUseCase + appID string + authToken string + allowedOrigins map[string]struct{} + maxBodyBytes int64 + runTimeout time.Duration + logger *log.Logger +} + +// NewDashScopeChatHandler constructs the protected compatibility transport. +func NewDashScopeChatHandler(chat ChatUseCase, options DashScopeChatOptions) (*DashScopeChatHandler, error) { + if chat == nil { + return nil, errors.New("chat service is required") + } + if !validDashScopeAppID(options.AppID) { + return nil, errors.New("DashScope-compatible app ID is invalid") + } + if !validChatToken(options.AuthToken) { + return nil, errors.New("DashScope-compatible auth token must contain at least 32 printable ASCII characters") + } + if options.MaxBodyBytes <= 0 || options.MaxBodyBytes > maximumChatBodyBytes { + return nil, errors.New("DashScope-compatible max body bytes is outside the supported range") + } + if options.RunTimeout <= 0 || options.RunTimeout > maximumChatRunTime { + return nil, errors.New("DashScope-compatible run timeout is outside the supported range") + } + allowedOrigins := make(map[string]struct{}, len(options.AllowedOrigins)) + for _, origin := range options.AllowedOrigins { + normalized, ok := validChatOrigin(origin) + if !ok || normalized != origin { + return nil, errors.New("DashScope-compatible allowed origins must be normalized exact HTTP(S) origins") + } + allowedOrigins[origin] = struct{}{} + } + logger := options.Logger + if logger == nil { + logger = log.New(io.Discard, "", 0) + } + return &DashScopeChatHandler{ + chat: chat, + appID: options.AppID, + authToken: options.AuthToken, + allowedOrigins: allowedOrigins, + maxBodyBytes: options.MaxBodyBytes, + runTimeout: options.RunTimeout, + logger: logger, + }, nil +} + +// ServeHTTP handles compatible application completion requests and CORS +// preflight. It buffers the upstream answer until strict ChatUseCase success. +func (h *DashScopeChatHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + started := time.Now() + requestID := safeRequestID(r.Header.Get("X-Request-ID")) + w.Header().Set("X-Request-ID", requestID) + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("X-Content-Type-Options", "nosniff") + + if r.PathValue("app_id") != h.appID { + h.writeError(w, http.StatusNotFound, requestID, "CHAT_APP_NOT_FOUND", "Chat application was not found.") + h.logResult(requestID, "app_not_found", false, started) + return + } + if r.Method == http.MethodOptions { + h.handlePreflight(w, r, requestID, started) + return + } + if r.Method != http.MethodPost { + w.Header().Set("Allow", http.MethodPost+", "+http.MethodOptions) + h.writeError(w, http.StatusMethodNotAllowed, requestID, "CHAT_METHOD_NOT_ALLOWED", "Chat endpoint accepts POST requests only.") + h.logResult(requestID, "method_not_allowed", false, started) + return + } + if !h.authorizeOrigin(w, r.Header.Get("Origin")) { + h.writeError(w, http.StatusForbidden, requestID, "CHAT_ORIGIN_FORBIDDEN", "Chat browser origin is not allowed.") + h.logResult(requestID, "forbidden_origin", false, started) + return + } + if !h.validToken(r.Header.Get("xtoken")) { + w.Header().Set("WWW-Authenticate", `XToken realm="fire-safety-ymd-chat"`) + h.writeError(w, http.StatusUnauthorized, requestID, "CHAT_AUTH_INVALID", "Chat authentication failed.") + h.logResult(requestID, "unauthorized", false, started) + return + } + if !isJSONContentType(r.Header.Get("Content-Type")) { + h.writeError(w, http.StatusUnsupportedMediaType, requestID, "CHAT_CONTENT_TYPE_INVALID", "Content-Type must be application/json.") + h.logResult(requestID, "unsupported_media_type", false, started) + return + } + + body, tooLarge, err := readBoundedBody(r.Body, h.maxBodyBytes) + if err != nil { + h.writeError(w, http.StatusBadRequest, requestID, "CHAT_REQUEST_INVALID", "Unable to read chat request.") + h.logResult(requestID, "read_error", false, started) + return + } + if tooLarge { + h.writeError(w, http.StatusRequestEntityTooLarge, requestID, "CHAT_REQUEST_BODY_TOO_LARGE", "Chat request body exceeds the configured limit.") + h.logResult(requestID, "body_too_large", false, started) + return + } + request, err := decodeDashScopeChatRequest(body) + if err != nil { + h.writeError(w, http.StatusBadRequest, requestID, "CHAT_REQUEST_INVALID", "Chat request JSON is invalid.") + h.logResult(requestID, "invalid_request", false, started) + return + } + + runCtx, cancelRun := context.WithTimeout(r.Context(), h.runTimeout) + defer cancelRun() + turn, err := h.chat.Prepare(runCtx, service.ChatRequest{ + Message: request.Prompt, + ConversationID: request.SessionID, + RequestID: requestID, + }) + if err != nil { + status, code, message := publicChatError(err) + if errors.Is(err, service.ErrChatConversationBusy) { + w.Header().Set("Retry-After", "1") + } + if errors.Is(err, service.ErrChatCapacityReached) { + w.Header().Set("Retry-After", "30") + } + h.writeError(w, status, requestID, code, message) + h.logResult(requestID, code, request.SessionID != "", started) + return + } + defer turn.Close() + + flusher, ok := w.(http.Flusher) + if !ok { + h.writeError(w, http.StatusInternalServerError, requestID, "CHAT_STREAM_UNSUPPORTED", "Chat streaming is unavailable.") + h.logResult(requestID, "stream_unsupported", turn.Reused(), started) + return + } + if deadline, ok := runCtx.Deadline(); ok { + _ = http.NewResponseController(w).SetWriteDeadline(deadline.Add(5 * time.Second)) + } + w.Header().Set("Content-Type", "text/event-stream; charset=utf-8") + w.Header().Set("X-Accel-Buffering", "no") + w.WriteHeader(http.StatusOK) + if err := writeDashScopeSSE(w, flusher, 1, "result", dashScopeResultPayload{ + Output: dashScopeOutput{SessionID: turn.ConversationID(), FinishReason: "null"}, + Usage: dashScopeUsage{}, + RequestID: requestID, + }); err != nil { + h.logResult(requestID, "client_write_failed", turn.Reused(), started) + return + } + + streamCtx, cancelStream := context.WithCancel(runCtx) + defer cancelStream() + type streamOutcome struct { + result service.ChatResult + err error + } + outcome := make(chan streamOutcome, 1) + go func() { + result, streamErr := turn.Stream(streamCtx, nil) + outcome <- streamOutcome{result: result, err: streamErr} + }() + + heartbeat := time.NewTicker(chatHeartbeatInterval) + defer heartbeat.Stop() + contextDone := runCtx.Done() + var result service.ChatResult + var streamErr error +streamLoop: + for { + select { + case completed := <-outcome: + result = completed.result + streamErr = completed.err + break streamLoop + case <-heartbeat.C: + if err := writeChatHeartbeat(w, flusher); err != nil { + cancelStream() + h.logResult(requestID, "client_write_failed", turn.Reused(), started) + return + } + case <-contextDone: + cancelStream() + contextDone = nil + } + } + if streamErr != nil { + _, code, message := publicChatError(streamErr) + _ = writeDashScopeSSE(w, flusher, 2, "error", dashScopeStreamError{ + Code: code, + Message: message, + RequestID: requestID, + SessionID: turn.ConversationID(), + }) + h.logResult(requestID, code, turn.Reused(), started) + return + } + if err := writeDashScopeSSE(w, flusher, 2, "result", dashScopeResultPayload{ + Output: dashScopeOutput{ + SessionID: result.ConversationID, + FinishReason: "stop", + Text: result.Answer, + }, + Usage: dashScopeUsage{Models: []dashScopeModelUsage{{ + InputTokens: result.Usage.Input, + OutputTokens: result.Usage.Output, + ModelID: result.ModelID, + }}}, + RequestID: requestID, + }); err != nil { + h.logResult(requestID, "client_write_failed", turn.Reused(), started) + return + } + h.logResult(requestID, "success", turn.Reused(), started) +} + +type dashScopeChatRequest struct { + Prompt string + SessionID string +} + +func decodeDashScopeChatRequest(body []byte) (dashScopeChatRequest, error) { + trimmed := bytes.TrimSpace(body) + if len(trimmed) == 0 || trimmed[0] != '{' { + return dashScopeChatRequest{}, errors.New("request must be an object") + } + var envelope struct { + Input json.RawMessage `json:"input"` + Parameters json.RawMessage `json:"parameters"` + Debug json.RawMessage `json:"debug"` + } + decoder := json.NewDecoder(bytes.NewReader(trimmed)) + decoder.DisallowUnknownFields() + if err := decoder.Decode(&envelope); err != nil { + return dashScopeChatRequest{}, errors.New("request is invalid") + } + if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + return dashScopeChatRequest{}, errors.New("request contains trailing JSON") + } + + input, err := decodeDashScopeInput(envelope.Input) + if err != nil { + return dashScopeChatRequest{}, err + } + if len(envelope.Parameters) > 0 { + var parameters struct { + IncrementalOutput json.RawMessage `json:"incremental_output"` + } + if err := decodeStrictObject(envelope.Parameters, ¶meters); err != nil { + return dashScopeChatRequest{}, errors.New("parameters must be a supported object") + } + if len(parameters.IncrementalOutput) > 0 { + var incrementalOutput bool + if err := json.Unmarshal(parameters.IncrementalOutput, &incrementalOutput); err != nil || bytes.Equal(bytes.TrimSpace(parameters.IncrementalOutput), []byte("null")) { + return dashScopeChatRequest{}, errors.New("parameters.incremental_output must be a boolean") + } + } + } + if len(envelope.Debug) > 0 { + var debug map[string]json.RawMessage + if err := decodeStrictObject(envelope.Debug, &debug); err != nil || len(debug) != 0 { + return dashScopeChatRequest{}, errors.New("debug must be an empty object") + } + } + return input, nil +} + +func decodeDashScopeInput(raw json.RawMessage) (dashScopeChatRequest, error) { + var input struct { + Prompt *string `json:"prompt"` + SessionID json.RawMessage `json:"session_id"` + } + if err := decodeStrictObject(raw, &input); err != nil || input.Prompt == nil || strings.TrimSpace(*input.Prompt) == "" { + return dashScopeChatRequest{}, errors.New("input.prompt is required") + } + request := dashScopeChatRequest{Prompt: *input.Prompt} + if len(input.SessionID) > 0 { + if err := json.Unmarshal(input.SessionID, &request.SessionID); err != nil || strings.TrimSpace(request.SessionID) == "" { + return dashScopeChatRequest{}, errors.New("input.session_id must be a non-empty string when provided") + } + } + return request, nil +} + +type dashScopeResultPayload struct { + Output dashScopeOutput `json:"output"` + Usage dashScopeUsage `json:"usage"` + RequestID string `json:"request_id"` +} + +type dashScopeOutput struct { + SessionID string `json:"session_id"` + FinishReason string `json:"finish_reason"` + Text string `json:"text,omitempty"` +} + +type dashScopeUsage struct { + Models []dashScopeModelUsage `json:"models,omitempty"` +} + +type dashScopeModelUsage struct { + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + ModelID string `json:"model_id,omitempty"` +} + +type dashScopeStreamError struct { + Code string `json:"code"` + Message string `json:"message"` + RequestID string `json:"request_id"` + SessionID string `json:"session_id"` +} + +func (h *DashScopeChatHandler) handlePreflight(w http.ResponseWriter, r *http.Request, requestID string, started time.Time) { + origin := r.Header.Get("Origin") + if origin == "" || !h.authorizeOrigin(w, origin) || r.Header.Get("Access-Control-Request-Method") != http.MethodPost || + !validDashScopePreflightHeaders(r.Header.Get("Access-Control-Request-Headers")) { + h.writeError(w, http.StatusForbidden, requestID, "CHAT_ORIGIN_FORBIDDEN", "Chat browser origin or preflight request is not allowed.") + h.logResult(requestID, "preflight_forbidden", false, started) + return + } + w.Header().Set("Access-Control-Allow-Methods", http.MethodPost) + w.Header().Set("Access-Control-Allow-Headers", "Content-Type, xtoken, X-DashScope-SSE, X-Request-ID") + w.Header().Set("Access-Control-Max-Age", "600") + w.WriteHeader(http.StatusNoContent) + h.logResult(requestID, "preflight_success", false, started) +} + +func (h *DashScopeChatHandler) authorizeOrigin(w http.ResponseWriter, origin string) bool { + if origin == "" { + return true + } + if _, allowed := h.allowedOrigins[origin]; !allowed { + return false + } + w.Header().Add("Vary", "Origin") + w.Header().Set("Access-Control-Allow-Origin", origin) + w.Header().Set("Access-Control-Expose-Headers", "X-Request-ID") + return true +} + +func (h *DashScopeChatHandler) validToken(value string) bool { + if len(value) != len(h.authToken) { + return false + } + return subtle.ConstantTimeCompare([]byte(value), []byte(h.authToken)) == 1 +} + +func (h *DashScopeChatHandler) writeError(w http.ResponseWriter, status int, requestID, code, message string) { + writeJSON(w, status, map[string]string{ + "code": code, + "message": message, + "request_id": requestID, + }) +} + +func (h *DashScopeChatHandler) logResult(requestID, result string, reused bool, started time.Time) { + h.logger.Printf("dashscope_chat_request request_id=%s result=%s reused=%t duration_ms=%d", requestID, result, reused, time.Since(started).Milliseconds()) +} + +func writeDashScopeSSE(w io.Writer, flusher http.Flusher, id int, event string, payload any) error { + data, err := json.Marshal(payload) + if err != nil { + return err + } + if _, err := fmt.Fprintf(w, "id: %d\nevent: %s\n:HTTP_STATUS/200\ndata: %s\n\n", id, event, data); err != nil { + return err + } + flusher.Flush() + return nil +} + +func validDashScopePreflightHeaders(value string) bool { + for _, header := range strings.Split(value, ",") { + switch strings.ToLower(strings.TrimSpace(header)) { + case "", "content-type", "xtoken", "x-dashscope-sse", "x-request-id": + default: + return false + } + } + return true +} + +func validDashScopeAppID(value string) bool { + if value == "" || len(value) > 128 || strings.TrimSpace(value) != value { + return false + } + for _, character := range value { + if !(character >= 'a' && character <= 'z') && + !(character >= 'A' && character <= 'Z') && + !(character >= '0' && character <= '9') && + !strings.ContainsRune("_-", character) { + return false + } + } + return true +} diff --git a/internal/handler/dashscope_chat_test.go b/internal/handler/dashscope_chat_test.go new file mode 100644 index 0000000..e5541aa --- /dev/null +++ b/internal/handler/dashscope_chat_test.go @@ -0,0 +1,330 @@ +package handler + +import ( + "bytes" + "encoding/json" + "fmt" + "log" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "fire-safety-ymd/internal/service" +) + +const testDashScopeAppID = "fire-safety-test-app" + +func TestNewDashScopeChatHandlerRejectsUnsafeOptions(t *testing.T) { + tests := []DashScopeChatOptions{ + {AppID: "", AuthToken: testChatToken, MaxBodyBytes: 1, RunTimeout: time.Second}, + {AppID: "app/other", AuthToken: testChatToken, MaxBodyBytes: 1, RunTimeout: time.Second}, + {AppID: "app.other", AuthToken: testChatToken, MaxBodyBytes: 1, RunTimeout: time.Second}, + {AppID: testDashScopeAppID, AuthToken: "short", MaxBodyBytes: 1, RunTimeout: time.Second}, + {AppID: testDashScopeAppID, AuthToken: testChatToken, MaxBodyBytes: maximumChatBodyBytes + 1, RunTimeout: time.Second}, + {AppID: testDashScopeAppID, AuthToken: testChatToken, MaxBodyBytes: 1, RunTimeout: maximumChatRunTime + time.Second}, + {AppID: testDashScopeAppID, AuthToken: testChatToken, AllowedOrigins: []string{"*"}, MaxBodyBytes: 1, RunTimeout: time.Second}, + } + for _, options := range tests { + if _, err := NewDashScopeChatHandler(&fakeChatUseCase{}, options); err == nil { + t.Fatalf("NewDashScopeChatHandler(%#v) error = nil", options) + } + } + if _, err := NewDashScopeChatHandler(nil, DashScopeChatOptions{ + AppID: testDashScopeAppID, AuthToken: testChatToken, MaxBodyBytes: 1, RunTimeout: time.Second, + }); err == nil { + t.Fatal("NewDashScopeChatHandler(nil) error = nil") + } +} + +func TestDashScopeChatStreamsCompatibleResultAfterStrictSuccess(t *testing.T) { + turn := &fakeChatTurn{ + conversationID: "conv_0123456789abcdef01234567", + events: []service.AgentTraceEvent{ + {Event: "message.delta", Status: "running"}, + {Event: "tool.started", ToolName: "sensitive_tool", Status: "running"}, + }, + result: service.ChatResult{ + ConversationID: "conv_0123456789abcdef01234567", + RunID: "provider-run-must-not-leak", + ModelID: "qwen-plus-latest", + Answer: "候选水源需要现场确认。", + Usage: service.ChatTokenUsage{Input: 12, Output: 34, Total: 46}, + }, + } + chat := &fakeChatUseCase{turn: turn} + var logs bytes.Buffer + handler := newTestDashScopeChatHandler(t, chat, []string{"https://allowed.example"}, log.New(&logs, "", 0)) + request := authenticatedDashScopeRequest(http.MethodPost, `{"input":{"prompt":"不要记录 secret prompt"},"parameters":{"incremental_output":true},"debug":{}}`) + request.Header.Set("Origin", "https://allowed.example") + request.Header.Set("X-Request-ID", "compat-request-1") + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusOK || !strings.HasPrefix(response.Header().Get("Content-Type"), "text/event-stream") { + t.Fatalf("status=%d content-type=%q body=%s", response.Code, response.Header().Get("Content-Type"), response.Body.String()) + } + if response.Header().Get("X-Accel-Buffering") != "no" || response.Header().Get("Cache-Control") != "no-store" { + t.Fatalf("stream headers = %#v", response.Header()) + } + blocks := dashScopeSSEBlocks(response.Body.String()) + if len(blocks) != 2 { + t.Fatalf("SSE block count=%d body=%s", len(blocks), response.Body.String()) + } + for index, expectedID := range []string{"id: 1", "id: 2"} { + if !strings.Contains(blocks[index], expectedID) || !strings.Contains(blocks[index], "event: result") || !strings.Contains(blocks[index], ":HTTP_STATUS/200") { + t.Fatalf("block %d is incompatible: %s", index, blocks[index]) + } + } + initial := decodeDashScopeSSEData(t, blocks[0]) + if initial.Output.SessionID != turn.conversationID || initial.Output.FinishReason != "null" || initial.Output.Text != "" || len(initial.Usage.Models) != 0 || initial.RequestID != "compat-request-1" { + t.Fatalf("initial payload = %#v", initial) + } + final := decodeDashScopeSSEData(t, blocks[1]) + if final.Output.SessionID != turn.conversationID || final.Output.FinishReason != "stop" || final.Output.Text != "候选水源需要现场确认。" || final.RequestID != "compat-request-1" { + t.Fatalf("final payload = %#v", final) + } + if len(final.Usage.Models) != 1 || final.Usage.Models[0].InputTokens != 12 || final.Usage.Models[0].OutputTokens != 34 || final.Usage.Models[0].ModelID != "qwen-plus-latest" { + t.Fatalf("final usage = %#v", final.Usage) + } + body := response.Body.String() + for _, forbidden := range []string{"message.delta", "sensitive_tool", "provider-run-must-not-leak"} { + if strings.Contains(body, forbidden) { + t.Fatalf("SSE exposed %q: %s", forbidden, body) + } + } + if chat.lastRequest.Message != "不要记录 secret prompt" || chat.lastRequest.ConversationID != "" || chat.lastRequest.RequestID != "compat-request-1" { + t.Fatalf("service request = %#v", chat.lastRequest) + } + if !turn.closed { + t.Fatal("turn was not closed") + } + for _, sensitive := range []string{"secret prompt", "候选水源", turn.conversationID, testChatToken} { + if strings.Contains(logs.String(), sensitive) { + t.Fatalf("logs contain sensitive value %q: %s", sensitive, logs.String()) + } + } +} + +func TestDashScopeChatMapsSessionIDToLocalConversation(t *testing.T) { + turn := &fakeChatTurn{ + conversationID: "conv_0123456789abcdef01234567", + reused: true, + result: service.ChatResult{ + ConversationID: "conv_0123456789abcdef01234567", + Answer: "ok", + }, + } + chat := &fakeChatUseCase{turn: turn} + handler := newTestDashScopeChatHandler(t, chat, nil, nil) + request := authenticatedDashScopeRequest(http.MethodPost, `{"input":{"prompt":"follow up","session_id":"conv_0123456789abcdef01234567"},"parameters":{}}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } + if chat.lastRequest.ConversationID != "conv_0123456789abcdef01234567" { + t.Fatalf("conversation ID = %q", chat.lastRequest.ConversationID) + } +} + +func TestDashScopeChatTransportGuardsBeforeService(t *testing.T) { + tests := []struct { + name string + method string + body string + configure func(*http.Request) + maxBody int64 + wantStatus int + wantCode string + }{ + {name: "wrong app", method: http.MethodPost, body: `{"input":{"prompt":"hello"}}`, configure: func(r *http.Request) { r.SetPathValue("app_id", "another-app") }, wantStatus: http.StatusNotFound, wantCode: "CHAT_APP_NOT_FOUND"}, + {name: "wrong method", method: http.MethodGet, body: `{}`, wantStatus: http.StatusMethodNotAllowed, wantCode: "CHAT_METHOD_NOT_ALLOWED"}, + {name: "forbidden origin", method: http.MethodPost, body: `{"input":{"prompt":"hello"}}`, configure: func(r *http.Request) { r.Header.Set("Origin", "https://evil.example") }, wantStatus: http.StatusForbidden, wantCode: "CHAT_ORIGIN_FORBIDDEN"}, + {name: "missing xtoken", method: http.MethodPost, body: `{"input":{"prompt":"hello"}}`, configure: func(r *http.Request) { r.Header.Del("xtoken") }, wantStatus: http.StatusUnauthorized, wantCode: "CHAT_AUTH_INVALID"}, + {name: "bearer is not xtoken", method: http.MethodPost, body: `{"input":{"prompt":"hello"}}`, configure: func(r *http.Request) { r.Header.Del("xtoken"); r.Header.Set("Authorization", "Bearer "+testChatToken) }, wantStatus: http.StatusUnauthorized, wantCode: "CHAT_AUTH_INVALID"}, + {name: "wrong content type", method: http.MethodPost, body: `{"input":{"prompt":"hello"}}`, configure: func(r *http.Request) { r.Header.Set("Content-Type", "text/plain") }, wantStatus: http.StatusUnsupportedMediaType, wantCode: "CHAT_CONTENT_TYPE_INVALID"}, + {name: "body too large", method: http.MethodPost, body: `{"input":{"prompt":"hello"}}`, maxBody: 4, wantStatus: http.StatusRequestEntityTooLarge, wantCode: "CHAT_REQUEST_BODY_TOO_LARGE"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + chat := &fakeChatUseCase{} + maxBody := tt.maxBody + if maxBody == 0 { + maxBody = 4096 + } + handler := newTestDashScopeChatHandlerWithLimit(t, chat, maxBody, []string{"https://allowed.example"}, nil) + request := authenticatedDashScopeRequest(tt.method, tt.body) + if tt.configure != nil { + tt.configure(request) + } + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != tt.wantStatus || !strings.Contains(response.Body.String(), `"code":"`+tt.wantCode+`"`) { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } + if chat.prepareCalls != 0 { + t.Fatalf("Prepare calls = %d, want 0", chat.prepareCalls) + } + }) + } +} + +func TestDashScopeChatCORSPreflight(t *testing.T) { + handler := newTestDashScopeChatHandler(t, &fakeChatUseCase{}, []string{"http://localhost:5173"}, nil) + + allowed := authenticatedDashScopeRequest(http.MethodOptions, "") + allowed.Header.Set("Origin", "http://localhost:5173") + allowed.Header.Set("Access-Control-Request-Method", http.MethodPost) + allowed.Header.Set("Access-Control-Request-Headers", "content-type, xtoken, x-dashscope-sse, x-request-id") + allowedResponse := httptest.NewRecorder() + handler.ServeHTTP(allowedResponse, allowed) + if allowedResponse.Code != http.StatusNoContent || allowedResponse.Header().Get("Access-Control-Allow-Origin") != "http://localhost:5173" { + t.Fatalf("allowed preflight status=%d headers=%#v body=%s", allowedResponse.Code, allowedResponse.Header(), allowedResponse.Body.String()) + } + + forbidden := authenticatedDashScopeRequest(http.MethodOptions, "") + forbidden.Header.Set("Origin", "http://localhost:5173") + forbidden.Header.Set("Access-Control-Request-Method", http.MethodPost) + forbidden.Header.Set("Access-Control-Request-Headers", "authorization") + forbiddenResponse := httptest.NewRecorder() + handler.ServeHTTP(forbiddenResponse, forbidden) + if forbiddenResponse.Code != http.StatusForbidden || !strings.Contains(forbiddenResponse.Body.String(), "CHAT_ORIGIN_FORBIDDEN") { + t.Fatalf("forbidden preflight status=%d body=%s", forbiddenResponse.Code, forbiddenResponse.Body.String()) + } +} + +func TestDashScopeChatRejectsNonCompatibleJSON(t *testing.T) { + tests := []string{ + ``, + `[]`, + `{}`, + `{"input":null}`, + `{"input":{}}`, + `{"input":{"prompt":null}}`, + `{"input":{"prompt":" "}}`, + `{"input":{"prompt":"hello","session_id":null}}`, + `{"input":{"prompt":"hello","user_id":"admin"}}`, + `{"input":{"prompt":"hello"},"parameters":null}`, + `{"input":{"prompt":"hello"},"parameters":{"incremental_output":null}}`, + `{"input":{"prompt":"hello"},"parameters":{"temperature":1}}`, + `{"input":{"prompt":"hello"},"parameters":{"incremental_output":"yes"}}`, + `{"input":{"prompt":"hello"},"debug":null}`, + `{"input":{"prompt":"hello"},"debug":{"trace":true}}`, + `{"input":{"prompt":"hello"},"metadata":{"role":"admin"}}`, + `{"input":{"prompt":"hello"}} {}`, + } + for _, body := range tests { + t.Run(fmt.Sprintf("%q", body), func(t *testing.T) { + chat := &fakeChatUseCase{} + handler := newTestDashScopeChatHandler(t, chat, nil, nil) + request := authenticatedDashScopeRequest(http.MethodPost, body) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusBadRequest || !strings.Contains(response.Body.String(), `"code":"CHAT_REQUEST_INVALID"`) { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } + if chat.prepareCalls != 0 { + t.Fatalf("Prepare calls = %d, want 0", chat.prepareCalls) + } + }) + } +} + +func TestDashScopeChatStreamFailureDoesNotExposePartialAnswerOrStop(t *testing.T) { + turn := &fakeChatTurn{ + conversationID: "conv_0123456789abcdef01234567", + events: []service.AgentTraceEvent{ + {Event: "message.delta", ToolName: "secret tool", Status: "partial answer that must not leak"}, + }, + err: fmt.Errorf("provider secret response: %w", service.ErrChatUpstreamProtocol), + } + handler := newTestDashScopeChatHandler(t, &fakeChatUseCase{turn: turn}, nil, nil) + request := authenticatedDashScopeRequest(http.MethodPost, `{"input":{"prompt":"hello"}}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + body := response.Body.String() + if response.Code != http.StatusOK || !strings.Contains(body, "event: error") || !strings.Contains(body, "CHAT_UPSTREAM_PROTOCOL_ERROR") { + t.Fatalf("status=%d body=%s", response.Code, body) + } + for _, forbidden := range []string{`"finish_reason":"stop"`, "partial answer", "provider secret response", "secret tool"} { + if strings.Contains(body, forbidden) { + t.Fatalf("failure stream exposed %q: %s", forbidden, body) + } + } +} + +func TestDashScopeChatPreparationBusyReturnsRetryAfter(t *testing.T) { + handler := newTestDashScopeChatHandler(t, &fakeChatUseCase{err: service.ErrChatConversationBusy}, nil, nil) + request := authenticatedDashScopeRequest(http.MethodPost, `{"input":{"prompt":"hello","session_id":"conv_0123456789abcdef01234567"}}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusConflict || response.Header().Get("Retry-After") != "1" || !strings.Contains(response.Body.String(), `"code":"CHAT_CONVERSATION_BUSY"`) { + t.Fatalf("status=%d headers=%#v body=%s", response.Code, response.Header(), response.Body.String()) + } +} + +func newTestDashScopeChatHandler(t *testing.T, chat ChatUseCase, origins []string, logger *log.Logger) *DashScopeChatHandler { + t.Helper() + return newTestDashScopeChatHandlerWithLimit(t, chat, 4096, origins, logger) +} + +func newTestDashScopeChatHandlerWithLimit(t *testing.T, chat ChatUseCase, maxBody int64, origins []string, logger *log.Logger) *DashScopeChatHandler { + t.Helper() + handler, err := NewDashScopeChatHandler(chat, DashScopeChatOptions{ + AppID: testDashScopeAppID, + AuthToken: testChatToken, + AllowedOrigins: origins, + MaxBodyBytes: maxBody, + RunTimeout: time.Minute, + Logger: logger, + }) + if err != nil { + t.Fatalf("NewDashScopeChatHandler() error = %v", err) + } + return handler +} + +func authenticatedDashScopeRequest(method, body string) *http.Request { + request := httptest.NewRequest(method, "/api/v1/apps/"+testDashScopeAppID+"/completion", strings.NewReader(body)) + request.SetPathValue("app_id", testDashScopeAppID) + request.Header.Set("xtoken", testChatToken) + request.Header.Set("Content-Type", "application/json") + return request +} + +func dashScopeSSEBlocks(body string) []string { + trimmed := strings.TrimSpace(body) + if trimmed == "" { + return nil + } + return strings.Split(trimmed, "\n\n") +} + +func decodeDashScopeSSEData(t *testing.T, block string) dashScopeResultPayload { + t.Helper() + for _, line := range strings.Split(block, "\n") { + if !strings.HasPrefix(line, "data: ") { + continue + } + var payload dashScopeResultPayload + if err := json.Unmarshal([]byte(strings.TrimPrefix(line, "data: ")), &payload); err != nil { + t.Fatalf("decode SSE data: %v", err) + } + return payload + } + t.Fatalf("SSE block has no data: %s", block) + return dashScopeResultPayload{} +} diff --git a/internal/handler/health.go b/internal/handler/health.go new file mode 100644 index 0000000..9728eaa --- /dev/null +++ b/internal/handler/health.go @@ -0,0 +1,43 @@ +// Package handler implements inbound HTTP contracts. +package handler + +import ( + "encoding/json" + "net/http" +) + +type healthResponse struct { + Status string `json:"status"` +} + +// RouterOptions contains optional handlers wired by the application layer. +type RouterOptions struct { + Chat http.Handler + DashScopeChat http.Handler + MCP http.Handler +} + +// NewRouter returns the HTTP handler for the application. +func NewRouter(options ...RouterOptions) http.Handler { + mux := http.NewServeMux() + mux.HandleFunc("GET /health", health) + if len(options) > 0 { + if options[0].Chat != nil { + mux.Handle("/api/chat", options[0].Chat) + } + if options[0].DashScopeChat != nil { + mux.Handle("/api/v1/apps/{app_id}/completion", options[0].DashScopeChat) + } + if options[0].MCP != nil { + mux.Handle("/mcp", options[0].MCP) + } + } + + return mux +} + +func health(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.WriteHeader(http.StatusOK) + _ = json.NewEncoder(w).Encode(healthResponse{Status: "ok"}) +} diff --git a/internal/handler/health_test.go b/internal/handler/health_test.go new file mode 100644 index 0000000..67d2a67 --- /dev/null +++ b/internal/handler/health_test.go @@ -0,0 +1,83 @@ +package handler + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" +) + +func TestHealth(t *testing.T) { + request := httptest.NewRequest(http.MethodGet, "/health", nil) + response := httptest.NewRecorder() + + NewRouter().ServeHTTP(response, request) + + if response.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", response.Code, http.StatusOK) + } + if got := response.Header().Get("Content-Type"); got != "application/json; charset=utf-8" { + t.Fatalf("Content-Type = %q, want application/json; charset=utf-8", got) + } + + var body healthResponse + if err := json.NewDecoder(response.Body).Decode(&body); err != nil { + t.Fatalf("decode response: %v", err) + } + if body.Status != "ok" { + t.Fatalf("status body = %q, want ok", body.Status) + } +} + +func TestHealthRejectsOtherMethods(t *testing.T) { + request := httptest.NewRequest(http.MethodPost, "/health", nil) + response := httptest.NewRecorder() + + NewRouter().ServeHTTP(response, request) + + if response.Code != http.StatusMethodNotAllowed { + t.Fatalf("status = %d, want %d", response.Code, http.StatusMethodNotAllowed) + } +} + +func TestRouterDoesNotExposeMCPWithoutConfiguredHandler(t *testing.T) { + request := httptest.NewRequest(http.MethodPost, "/mcp", nil) + response := httptest.NewRecorder() + + NewRouter().ServeHTTP(response, request) + + if response.Code != http.StatusNotFound { + t.Fatalf("status = %d, want %d", response.Code, http.StatusNotFound) + } +} + +func TestRouterRegistersConfiguredMCPHandler(t *testing.T) { + request := httptest.NewRequest(http.MethodPost, "/mcp", nil) + response := httptest.NewRecorder() + mcp := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusAccepted) + }) + + NewRouter(RouterOptions{MCP: mcp}).ServeHTTP(response, request) + + if response.Code != http.StatusAccepted { + t.Fatalf("status = %d, want %d", response.Code, http.StatusAccepted) + } +} + +func TestRouterRegistersConfiguredDashScopeChatHandler(t *testing.T) { + request := httptest.NewRequest(http.MethodPost, "/api/v1/apps/test-app/completion", nil) + response := httptest.NewRecorder() + chat := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.PathValue("app_id") != "test-app" { + t.Fatalf("app_id = %q", r.PathValue("app_id")) + } + w.WriteHeader(http.StatusAccepted) + }) + + NewRouter(RouterOptions{DashScopeChat: chat}).ServeHTTP(response, request) + + if response.Code != http.StatusAccepted { + t.Fatalf("status = %d, want %d", response.Code, http.StatusAccepted) + } +} diff --git a/internal/handler/mcp.go b/internal/handler/mcp.go new file mode 100644 index 0000000..ee62358 --- /dev/null +++ b/internal/handler/mcp.go @@ -0,0 +1,854 @@ +package handler + +import ( + "bytes" + "context" + "crypto/rand" + "crypto/subtle" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "mime" + "net/http" + "strings" + "time" + + "fire-safety-ymd/internal/domain" + "fire-safety-ymd/internal/service" +) + +const ( + mcpProtocolVersion = "2025-06-18" + mcpServerName = "fire-safety-ymd-spatial-readonly" + mcpServerVersion = "0.2.0" + maximumMCPBody = int64(1024 * 1024) + maximumToolTimeout = 30 * time.Second + + toolSearchPlaceCandidates = "fire_safety_search_place_candidates" + toolResolveIncidentContext = "fire_safety_resolve_incident_context" + toolFindNearbyWaterSources = "fire_safety_find_nearby_water_sources" + toolFindCommandPostCandidates = "fire_safety_find_command_post_candidates" + toolListNearbyAccessLines = "fire_safety_list_nearby_access_lines" + toolGetResponsibleUnits = "fire_safety_get_responsible_units" + toolFindNearbyRiskAreas = "fire_safety_find_nearby_risk_areas" +) + +// MCPSpatialService is the use-case surface exposed through the MCP handler. +type MCPSpatialService interface { + SearchPlaceCandidates(context.Context, string, int) (service.QueryResult[[]domain.PlaceCandidate], error) + ResolveIncidentContext(context.Context, domain.Coordinate) (service.QueryResult[[]domain.IncidentContext], error) + FindNearbyWaterSources(context.Context, domain.Coordinate, float64, int) (service.QueryResult[[]domain.WaterSource], error) + FindCommandPostCandidates(context.Context, domain.Coordinate, float64, int) (service.QueryResult[[]domain.CommandPostCandidate], error) + ListNearbyAccessLines(context.Context, domain.Coordinate, float64, int) (service.QueryResult[[]domain.AccessLine], error) + GetResponsibleUnits(context.Context, domain.Coordinate) (service.QueryResult[[]domain.ResponsibleUnit], error) + FindNearbyRiskAreas(context.Context, domain.Coordinate, float64, int) (service.QueryResult[[]domain.RiskArea], error) +} + +// MCPOptions contains the independently authenticated MCP transport settings. +type MCPOptions struct { + AuthToken string + MaxBodyBytes int64 + ToolTimeout time.Duration + Logger *log.Logger +} + +// MCPHandler implements the stateless JSON response subset of MCP Streamable HTTP. +type MCPHandler struct { + service MCPSpatialService + authToken string + maxBodyBytes int64 + toolTimeout time.Duration + logger *log.Logger +} + +// NewMCPHandler constructs a protected, read-only MCP HTTP handler. +func NewMCPHandler(spatialService MCPSpatialService, options MCPOptions) (*MCPHandler, error) { + if spatialService == nil { + return nil, errors.New("MCP spatial service is required") + } + if !validMCPToken(options.AuthToken) { + return nil, errors.New("MCP auth token must contain at least 32 printable ASCII characters") + } + if options.MaxBodyBytes <= 0 || options.MaxBodyBytes > maximumMCPBody { + return nil, errors.New("MCP max body bytes is outside the supported range") + } + if options.ToolTimeout <= 0 || options.ToolTimeout > maximumToolTimeout { + return nil, errors.New("MCP tool timeout is outside the supported range") + } + logger := options.Logger + if logger == nil { + logger = log.New(io.Discard, "", 0) + } + return &MCPHandler{ + service: spatialService, + authToken: options.AuthToken, + maxBodyBytes: options.MaxBodyBytes, + toolTimeout: options.ToolTimeout, + logger: logger, + }, nil +} + +func (h *MCPHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + started := time.Now() + requestID := safeRequestID(r.Header.Get("X-Request-ID")) + w.Header().Set("X-Request-ID", requestID) + w.Header().Set("Cache-Control", "no-store") + + if r.Method != http.MethodPost { + w.Header().Set("Allow", http.MethodPost) + w.WriteHeader(http.StatusMethodNotAllowed) + return + } + if r.Header.Get("Origin") != "" { + h.writeRPCError(w, http.StatusForbidden, nil, -32003, "MCP_ORIGIN_FORBIDDEN", "MCP endpoint does not accept browser-origin requests.") + h.logResult(requestID, "transport", "forbidden_origin", started) + return + } + if !h.validAuthorization(r.Header.Get("Authorization")) { + w.Header().Set("WWW-Authenticate", `Bearer realm="fire-safety-ymd-mcp"`) + h.writeRPCError(w, http.StatusUnauthorized, nil, -32001, "MCP_AUTH_INVALID", "MCP authentication failed.") + h.logResult(requestID, "transport", "unauthorized", started) + return + } + if !isJSONContentType(r.Header.Get("Content-Type")) { + h.writeRPCError(w, http.StatusUnsupportedMediaType, nil, -32600, "MCP_CONTENT_TYPE_INVALID", "Content-Type must be application/json.") + h.logResult(requestID, "transport", "unsupported_media_type", started) + return + } + if version := strings.TrimSpace(r.Header.Get("MCP-Protocol-Version")); version != "" && version != mcpProtocolVersion { + h.writeRPCError(w, http.StatusBadRequest, nil, -32602, "MCP_PROTOCOL_VERSION_UNSUPPORTED", "Unsupported MCP protocol version.") + h.logResult(requestID, "transport", "unsupported_protocol", started) + return + } + + body, tooLarge, err := readBoundedBody(r.Body, h.maxBodyBytes) + if err != nil { + h.writeRPCError(w, http.StatusBadRequest, nil, -32700, "MCP_REQUEST_INVALID", "Unable to read MCP request.") + h.logResult(requestID, "transport", "read_error", started) + return + } + if tooLarge { + h.writeRPCError(w, http.StatusRequestEntityTooLarge, nil, -32002, "MCP_REQUEST_BODY_TOO_LARGE", "MCP request body exceeds the configured limit.") + h.logResult(requestID, "transport", "body_too_large", started) + return + } + + request, rpcFailure := decodeRPCRequest(body) + if rpcFailure != nil { + h.writeRPCError(w, rpcFailure.httpStatus, nil, rpcFailure.rpcCode, rpcFailure.code, rpcFailure.message) + h.logResult(requestID, "transport", rpcFailure.code, started) + return + } + if len(request.ID) == 0 && request.Method != "notifications/initialized" { + w.WriteHeader(http.StatusAccepted) + h.logResult(requestID, "notification", "ignored", started) + return + } + + switch request.Method { + case "initialize": + if h.handleInitialize(w, request) { + h.logResult(requestID, "initialize", "success", started) + } else { + h.logResult(requestID, "initialize", "unsupported_protocol", started) + } + case "notifications/initialized": + w.WriteHeader(http.StatusAccepted) + h.logResult(requestID, "notifications/initialized", "accepted", started) + case "tools/list": + h.writeRPCResult(w, request.ID, map[string]any{"tools": mcpToolDefinitions()}) + h.logResult(requestID, "tools/list", "success", started) + case "tools/call": + h.handleToolCall(w, r, request, requestID, started) + default: + h.writeRPCError(w, http.StatusOK, request.ID, -32601, "MCP_METHOD_NOT_FOUND", "MCP method not found.") + h.logResult(requestID, "unknown", "method_not_found", started) + } +} + +func (h *MCPHandler) handleInitialize(w http.ResponseWriter, request rpcRequest) bool { + var params struct { + ProtocolVersion string `json:"protocolVersion"` + } + if len(request.Params) == 0 || json.Unmarshal(request.Params, ¶ms) != nil || params.ProtocolVersion != mcpProtocolVersion { + h.writeRPCError(w, http.StatusOK, request.ID, -32602, "MCP_PROTOCOL_VERSION_UNSUPPORTED", "Client must support MCP protocol version 2025-06-18.") + return false + } + h.writeRPCResult(w, request.ID, map[string]any{ + "protocolVersion": mcpProtocolVersion, + "capabilities": map[string]any{ + "tools": map[string]any{"listChanged": false}, + }, + "serverInfo": map[string]any{ + "name": mcpServerName, + "title": "fire-safety-ymd spatial read-only MCP", + "version": mcpServerVersion, + }, + "instructions": "Read-only planning support. Source records with invalid geometries are excluded, so results may be incomplete. Place-name matches are candidates that require user confirmation before coordinate-based analysis. Treat all records as potentially stale, verify resource availability and field safety, and never present access-line candidates as confirmed routes or responsibility records as live team locations.", + }) + return true +} + +func (h *MCPHandler) handleToolCall(w http.ResponseWriter, r *http.Request, request rpcRequest, requestID string, started time.Time) { + var params struct { + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` + } + if len(request.Params) == 0 || json.Unmarshal(request.Params, ¶ms) != nil || strings.TrimSpace(params.Name) == "" { + h.writeRPCError(w, http.StatusOK, request.ID, -32602, "MCP_TOOL_PARAMS_INVALID", "MCP tool call parameters are invalid.") + h.logResult(requestID, "tools/call", "invalid_params", started) + return + } + if len(params.Arguments) == 0 || bytes.Equal(params.Arguments, []byte("null")) { + params.Arguments = json.RawMessage(`{}`) + } + + toolCtx, cancel := context.WithTimeout(r.Context(), h.toolTimeout) + defer cancel() + output, err := h.callTool(toolCtx, params.Name, params.Arguments) + if err != nil { + code, message := publicToolError(err) + h.writeToolResult(w, request.ID, toolErrorPayload(code, message), true) + h.logResult(requestID, safeToolName(params.Name), code, started) + return + } + h.writeToolResult(w, request.ID, output, false) + h.logResult(requestID, safeToolName(params.Name), "success", started) +} + +func (h *MCPHandler) callTool(ctx context.Context, name string, arguments json.RawMessage) (any, error) { + switch name { + case toolSearchPlaceCandidates: + placeName, limit, err := decodePlaceSearchArguments(arguments) + if err != nil { + return nil, err + } + return h.service.SearchPlaceCandidates(ctx, placeName, limit) + case toolResolveIncidentContext: + point, err := decodePointArguments(arguments) + if err != nil { + return nil, err + } + return h.service.ResolveIncidentContext(ctx, point) + case toolFindNearbyWaterSources: + point, radius, limit, err := decodeNearbyArguments(arguments, 30_000) + if err != nil { + return nil, err + } + return h.service.FindNearbyWaterSources(ctx, point, radius, limit) + case toolFindCommandPostCandidates: + point, radius, limit, err := decodeNearbyArguments(arguments, 20_000) + if err != nil { + return nil, err + } + return h.service.FindCommandPostCandidates(ctx, point, radius, limit) + case toolListNearbyAccessLines: + point, radius, limit, err := decodeNearbyArguments(arguments, 10_000) + if err != nil { + return nil, err + } + return h.service.ListNearbyAccessLines(ctx, point, radius, limit) + case toolGetResponsibleUnits: + point, err := decodePointArguments(arguments) + if err != nil { + return nil, err + } + return h.service.GetResponsibleUnits(ctx, point) + case toolFindNearbyRiskAreas: + point, radius, limit, err := decodeNearbyArguments(arguments, 10_000) + if err != nil { + return nil, err + } + return h.service.FindNearbyRiskAreas(ctx, point, radius, limit) + default: + return nil, errToolNotFound + } +} + +func decodePlaceSearchArguments(raw json.RawMessage) (string, int, error) { + var arguments struct { + PlaceName *string `json:"place_name"` + Limit *int `json:"limit"` + } + if err := decodeStrictObject(raw, &arguments); err != nil || arguments.PlaceName == nil { + return "", 0, fmt.Errorf("%w: place_name is required and unknown fields are not allowed", service.ErrInvalidArgument) + } + var limit int + if arguments.Limit != nil { + limit = *arguments.Limit + if limit < 1 || limit > 20 { + return "", 0, fmt.Errorf("%w: limit must be between 1 and 20", service.ErrInvalidArgument) + } + } + return *arguments.PlaceName, limit, nil +} + +func (h *MCPHandler) writeToolResult(w http.ResponseWriter, id json.RawMessage, output any, isError bool) { + textContent, err := json.Marshal(output) + if err != nil { + output = toolErrorPayload("INTERNAL_ERROR", "MCP tool result could not be encoded.") + textContent, _ = json.Marshal(output) + isError = true + } + h.writeRPCResult(w, id, map[string]any{ + "content": []map[string]string{{"type": "text", "text": string(textContent)}}, + "structuredContent": output, + "isError": isError, + }) +} + +func (h *MCPHandler) writeRPCResult(w http.ResponseWriter, id json.RawMessage, result any) { + writeJSON(w, http.StatusOK, rpcResponse{ + JSONRPC: "2.0", + ID: normalizedID(id), + Result: result, + }) +} + +func (h *MCPHandler) writeRPCError(w http.ResponseWriter, status int, id json.RawMessage, rpcCode int, code, message string) { + writeJSON(w, status, rpcResponse{ + JSONRPC: "2.0", + ID: normalizedID(id), + Error: &rpcError{ + Code: rpcCode, + Message: message, + Data: map[string]string{"code": code}, + }, + }) +} + +func (h *MCPHandler) validAuthorization(value string) bool { + expected := "Bearer " + h.authToken + if len(value) != len(expected) { + return false + } + return subtle.ConstantTimeCompare([]byte(value), []byte(expected)) == 1 +} + +func validMCPToken(value string) bool { + if len(value) < 32 || len(value) > 4096 { + return false + } + for _, character := range value { + if character <= 0x20 || character >= 0x7f { + return false + } + } + return true +} + +func (h *MCPHandler) logResult(requestID, operation, result string, started time.Time) { + h.logger.Printf("mcp_request request_id=%s operation=%s result=%s duration_ms=%d", requestID, operation, result, time.Since(started).Milliseconds()) +} + +type rpcRequest struct { + JSONRPC string `json:"jsonrpc"` + ID json.RawMessage `json:"id"` + Method string `json:"method"` + Params json.RawMessage `json:"params"` +} + +type rpcResponse struct { + JSONRPC string `json:"jsonrpc"` + ID json.RawMessage `json:"id"` + Result any `json:"result,omitempty"` + Error *rpcError `json:"error,omitempty"` +} + +type rpcError struct { + Code int `json:"code"` + Message string `json:"message"` + Data map[string]string `json:"data,omitempty"` +} + +type rpcFailure struct { + httpStatus int + rpcCode int + code string + message string +} + +func decodeRPCRequest(body []byte) (rpcRequest, *rpcFailure) { + if !json.Valid(body) { + return rpcRequest{}, &rpcFailure{httpStatus: http.StatusBadRequest, rpcCode: -32700, code: "MCP_REQUEST_INVALID", message: "MCP request JSON is invalid."} + } + trimmed := bytes.TrimSpace(body) + if len(trimmed) == 0 || trimmed[0] != '{' { + return rpcRequest{}, &rpcFailure{httpStatus: http.StatusBadRequest, rpcCode: -32600, code: "MCP_REQUEST_INVALID", message: "MCP request must be one JSON-RPC object."} + } + var request rpcRequest + if err := json.Unmarshal(trimmed, &request); err != nil || request.JSONRPC != "2.0" || strings.TrimSpace(request.Method) == "" { + return rpcRequest{}, &rpcFailure{httpStatus: http.StatusBadRequest, rpcCode: -32600, code: "MCP_REQUEST_INVALID", message: "MCP request is not a valid JSON-RPC 2.0 request."} + } + if len(request.ID) > 0 && !validRPCID(request.ID) { + return rpcRequest{}, &rpcFailure{httpStatus: http.StatusBadRequest, rpcCode: -32600, code: "MCP_REQUEST_INVALID", message: "MCP request has an invalid JSON-RPC id."} + } + return request, nil +} + +func validRPCID(id json.RawMessage) bool { + trimmed := bytes.TrimSpace(id) + if bytes.Equal(trimmed, []byte("null")) { + return true + } + if len(trimmed) == 0 { + return false + } + if trimmed[0] == '"' { + var value string + return json.Unmarshal(trimmed, &value) == nil + } + var value json.Number + decoder := json.NewDecoder(bytes.NewReader(trimmed)) + decoder.UseNumber() + return decoder.Decode(&value) == nil +} + +func decodePointArguments(raw json.RawMessage) (domain.Coordinate, error) { + var arguments struct { + Longitude *float64 `json:"longitude"` + Latitude *float64 `json:"latitude"` + } + if err := decodeStrictObject(raw, &arguments); err != nil || arguments.Longitude == nil || arguments.Latitude == nil { + return domain.Coordinate{}, fmt.Errorf("%w: longitude and latitude are required", service.ErrInvalidArgument) + } + return domain.Coordinate{Longitude: *arguments.Longitude, Latitude: *arguments.Latitude}, nil +} + +func decodeNearbyArguments(raw json.RawMessage, maximumRadius float64) (domain.Coordinate, float64, int, error) { + var arguments struct { + Longitude *float64 `json:"longitude"` + Latitude *float64 `json:"latitude"` + RadiusMeters *float64 `json:"radius_meters"` + Limit *int `json:"limit"` + } + if err := decodeStrictObject(raw, &arguments); err != nil || arguments.Longitude == nil || arguments.Latitude == nil { + return domain.Coordinate{}, 0, 0, fmt.Errorf("%w: longitude and latitude are required and unknown fields are not allowed", service.ErrInvalidArgument) + } + var radius float64 + if arguments.RadiusMeters != nil { + radius = *arguments.RadiusMeters + if radius < 100 || radius > maximumRadius { + return domain.Coordinate{}, 0, 0, fmt.Errorf("%w: radius_meters is outside the tool limit", service.ErrInvalidArgument) + } + } + var limit int + if arguments.Limit != nil { + limit = *arguments.Limit + if limit < 1 || limit > 20 { + return domain.Coordinate{}, 0, 0, fmt.Errorf("%w: limit must be between 1 and 20", service.ErrInvalidArgument) + } + } + return domain.Coordinate{Longitude: *arguments.Longitude, Latitude: *arguments.Latitude}, radius, limit, nil +} + +func decodeStrictObject(raw json.RawMessage, target any) error { + trimmed := bytes.TrimSpace(raw) + if len(trimmed) == 0 || trimmed[0] != '{' { + return errors.New("arguments must be an object") + } + decoder := json.NewDecoder(bytes.NewReader(trimmed)) + decoder.DisallowUnknownFields() + if err := decoder.Decode(target); err != nil { + return err + } + if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + return errors.New("arguments contain trailing JSON") + } + return nil +} + +var errToolNotFound = errors.New("MCP tool not found") + +func publicToolError(err error) (string, string) { + switch { + case errors.Is(err, errToolNotFound): + return "TOOL_NOT_FOUND", "MCP tool not found." + case errors.Is(err, service.ErrInvalidArgument): + return "INVALID_ARGUMENT", "MCP tool arguments are invalid." + case errors.Is(err, service.ErrQueryTimeout), errors.Is(err, context.DeadlineExceeded), errors.Is(err, context.Canceled): + return "QUERY_TIMEOUT", "Spatial query did not complete within its time limit." + case errors.Is(err, service.ErrDataSourceUnavailable): + return "DATA_SOURCE_UNAVAILABLE", "Spatial data source is unavailable." + default: + return "INTERNAL_ERROR", "MCP tool failed." + } +} + +func toolErrorPayload(code, message string) map[string]any { + return map[string]any{ + "status": "error", + "data": []any{}, + "metadata": map[string]any{ + "generated_at": time.Now().UTC().Format(time.RFC3339Nano), + "data_sources": []string{}, + "spatial_reference": "EPSG:4326", + "result_count": 0, + }, + "warnings": []string{}, + "error": map[string]string{ + "code": code, + "message": message, + }, + } +} + +func writeJSON(w http.ResponseWriter, status int, value any) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(value) +} + +func normalizedID(id json.RawMessage) json.RawMessage { + if len(id) == 0 { + return json.RawMessage("null") + } + return id +} + +func readBoundedBody(reader io.Reader, maximum int64) ([]byte, bool, error) { + body, err := io.ReadAll(io.LimitReader(reader, maximum+1)) + if err != nil { + return nil, false, err + } + if int64(len(body)) > maximum { + return nil, true, nil + } + return body, false, nil +} + +func isJSONContentType(value string) bool { + mediaType, _, err := mime.ParseMediaType(value) + return err == nil && mediaType == "application/json" +} + +func safeRequestID(candidate string) string { + candidate = strings.TrimSpace(candidate) + if candidate != "" && len(candidate) <= 128 { + valid := true + for _, character := range candidate { + if !(character >= 'a' && character <= 'z') && + !(character >= 'A' && character <= 'Z') && + !(character >= '0' && character <= '9') && + !strings.ContainsRune("._:-", character) { + valid = false + break + } + } + if valid { + return candidate + } + } + random := make([]byte, 16) + if _, err := rand.Read(random); err == nil { + return hex.EncodeToString(random) + } + return "request-id-unavailable" +} + +func safeToolName(name string) string { + for _, known := range []string{ + toolSearchPlaceCandidates, + toolResolveIncidentContext, + toolFindNearbyWaterSources, + toolFindCommandPostCandidates, + toolListNearbyAccessLines, + toolGetResponsibleUnits, + toolFindNearbyRiskAreas, + } { + if name == known { + return known + } + } + return "unknown_tool" +} + +type mcpToolDefinition struct { + Name string `json:"name"` + Title string `json:"title"` + Description string `json:"description"` + InputSchema map[string]any `json:"inputSchema"` + OutputSchema map[string]any `json:"outputSchema"` + Annotations map[string]any `json:"annotations"` +} + +func mcpToolDefinitions() []mcpToolDefinition { + return []mcpToolDefinition{ + { + Name: toolSearchPlaceCandidates, + Title: "按地名搜索位置候选", + Description: "在现有森林防火记录的名称、镇街和村庄字段中搜索地名,返回有界候选和 WGS84 坐标。候选必须由用户确认;线面记录只返回代表点,不能直接当作演练点。", + InputSchema: placeSearchInputSchema(), + OutputSchema: outputEnvelopeSchema(placeCandidateItemSchema()), + Annotations: readOnlyAnnotations(), + }, + { + Name: toolResolveIncidentContext, + Title: "定位演练点所属防火网格", + Description: "根据 WGS84 坐标查询覆盖该点的防火网格和镇街上下文。只返回区域信息,不返回负责人或电话。", + InputSchema: pointInputSchema(), + OutputSchema: outputEnvelopeSchema(incidentContextItemSchema()), + Annotations: readOnlyAnnotations(), + }, + { + Name: toolFindNearbyWaterSources, + Title: "查询附近候选水源", + Description: "查询演练点附近的水源地和蓄水池并按距离排序。记录存在不代表当前可用,必须现场确认水量、取水条件和道路可达性。", + InputSchema: nearbyInputSchema(30_000, 10_000), + OutputSchema: outputEnvelopeSchema(waterSourceItemSchema()), + Annotations: readOnlyAnnotations(), + }, + { + Name: toolFindCommandPostCandidates, + Title: "查询指挥部候选设施", + Description: "查询演练点附近的防火检查站和瞭望哨候选点。只代表空间候选,不能直接确定为指挥部,需现场核验安全、通信、容量和可达性。", + InputSchema: nearbyInputSchema(20_000, 10_000), + OutputSchema: outputEnvelopeSchema(commandPostItemSchema()), + Annotations: readOnlyAnnotations(), + }, + { + Name: toolListNearbyAccessLines, + Title: "查询附近防火通道候选", + Description: "查询演练点附近已绘制的防火通道及最近接入点。不是路径规划工具,不代表道路当前可通行。", + InputSchema: nearbyInputSchema(10_000, 5_000), + OutputSchema: outputEnvelopeSchema(accessLineItemSchema()), + Annotations: readOnlyAnnotations(), + }, + { + Name: toolGetResponsibleUnits, + Title: "查询责任防火队伍", + Description: "根据坐标查询防火网格中记录的责任中队名称。不提供人员联系方式,也不表示队伍实时位置、战备状态或正式集结点。", + InputSchema: pointInputSchema(), + OutputSchema: outputEnvelopeSchema(responsibleUnitItemSchema()), + Annotations: readOnlyAnnotations(), + }, + { + Name: toolFindNearbyRiskAreas, + Title: "查询周边风险区域", + Description: "查询演练点周边的墓地坟区和林区工矿企业范围。结果用于提示进一步核验,不代表实时危险程度。", + InputSchema: nearbyInputSchema(10_000, 3_000), + OutputSchema: outputEnvelopeSchema(riskAreaItemSchema()), + Annotations: readOnlyAnnotations(), + }, + } +} + +func placeSearchInputSchema() map[string]any { + return map[string]any{ + "type": "object", + "additionalProperties": false, + "properties": map[string]any{ + "place_name": map[string]any{ + "type": "string", + "minLength": 2, + "maxLength": 100, + "description": "Town, village, grid, facility, water source, access line, or risk-area name", + }, + "limit": map[string]any{ + "type": "integer", + "minimum": 1, + "maximum": 20, + "default": 10, + "description": "Maximum number of candidates", + }, + }, + "required": []string{"place_name"}, + } +} + +func readOnlyAnnotations() map[string]any { + return map[string]any{ + "readOnlyHint": true, + "destructiveHint": false, + "idempotentHint": true, + "openWorldHint": false, + } +} + +func pointInputSchema() map[string]any { + return map[string]any{ + "type": "object", + "additionalProperties": false, + "properties": map[string]any{ + "longitude": numberSchema("WGS84 longitude", -180, 180), + "latitude": numberSchema("WGS84 latitude", -90, 90), + }, + "required": []string{"longitude", "latitude"}, + } +} + +func nearbyInputSchema(maximumRadius, defaultRadius float64) map[string]any { + properties := pointInputSchema()["properties"].(map[string]any) + properties["radius_meters"] = map[string]any{ + "type": "number", + "minimum": 100, + "maximum": maximumRadius, + "default": defaultRadius, + "description": "Search radius in meters", + } + properties["limit"] = map[string]any{ + "type": "integer", + "minimum": 1, + "maximum": 20, + "default": 10, + "description": "Maximum number of results", + } + return map[string]any{ + "type": "object", + "additionalProperties": false, + "properties": properties, + "required": []string{"longitude", "latitude"}, + } +} + +func numberSchema(description string, minimum, maximum float64) map[string]any { + return map[string]any{ + "type": "number", + "minimum": minimum, + "maximum": maximum, + "description": description, + } +} + +func stringSchema() map[string]any { return map[string]any{"type": "string", "maxLength": 4096} } +func boolSchema() map[string]any { return map[string]any{"type": "boolean"} } + +func coordinateSchema() map[string]any { + return map[string]any{ + "type": "object", + "properties": map[string]any{ + "longitude": numberSchema("WGS84 longitude", -180, 180), + "latitude": numberSchema("WGS84 latitude", -90, 90), + }, + "required": []string{"longitude", "latitude"}, + } +} + +func outputEnvelopeSchema(itemSchema map[string]any) map[string]any { + return map[string]any{ + "type": "object", + "properties": map[string]any{ + "status": map[string]any{"type": "string", "enum": []string{"ok", "no_results", "error"}}, + "data": map[string]any{"type": "array", "items": itemSchema}, + "metadata": map[string]any{ + "type": "object", + "properties": map[string]any{ + "generated_at": stringSchema(), + "data_sources": map[string]any{"type": "array", "items": stringSchema()}, + "spatial_reference": stringSchema(), + "result_count": map[string]any{"type": "integer", "minimum": 0}, + "search_radius_meters": map[string]any{"type": "number", "minimum": 0}, + }, + "required": []string{"generated_at", "data_sources", "spatial_reference", "result_count"}, + }, + "warnings": map[string]any{"type": "array", "items": stringSchema()}, + "error": map[string]any{ + "type": "object", + "properties": map[string]any{"code": stringSchema(), "message": stringSchema()}, + "required": []string{"code", "message"}, + }, + }, + "required": []string{"status", "data", "metadata", "warnings"}, + } +} + +func incidentContextItemSchema() map[string]any { + return objectItemSchema(map[string]any{ + "grid_id": stringSchema(), + "town": stringSchema(), + "area_label": stringSchema(), + }, []string{"grid_id"}) +} + +func placeCandidateItemSchema() map[string]any { + return objectItemSchema(map[string]any{ + "source_record_id": stringSchema(), + "place_type": stringSchema(), + "name": stringSchema(), + "town": stringSchema(), + "village": stringSchema(), + "matched_field": map[string]any{"type": "string", "enum": []string{"name", "town", "village"}}, + "matched_text": stringSchema(), + "match_kind": map[string]any{"type": "string", "enum": []string{"exact", "partial"}}, + "location": coordinateSchema(), + "location_kind": map[string]any{"type": "string", "enum": []string{"recorded_point", "representative_point"}}, + }, []string{"source_record_id", "place_type", "matched_field", "matched_text", "match_kind", "location", "location_kind"}) +} + +func waterSourceItemSchema() map[string]any { + return objectItemSchema(map[string]any{ + "source_record_id": stringSchema(), + "category": stringSchema(), + "name": stringSchema(), + "town": stringSchema(), + "village": stringSchema(), + "location": coordinateSchema(), + "distance_meters": map[string]any{"type": "number", "minimum": 0}, + "capacity_cubic_meters": map[string]any{"type": "number"}, + "resource_type": stringSchema(), + "reported_status": stringSchema(), + "source_timestamp_raw": map[string]any{"type": "integer"}, + }, []string{"source_record_id", "category", "location", "distance_meters"}) +} + +func commandPostItemSchema() map[string]any { + return objectItemSchema(map[string]any{ + "source_record_id": stringSchema(), + "facility_type": stringSchema(), + "name": stringSchema(), + "town": stringSchema(), + "village": stringSchema(), + "location": coordinateSchema(), + "distance_meters": map[string]any{"type": "number", "minimum": 0}, + "reported_status": stringSchema(), + "management_unit": stringSchema(), + }, []string{"source_record_id", "facility_type", "location", "distance_meters"}) +} + +func accessLineItemSchema() map[string]any { + return objectItemSchema(map[string]any{ + "source_record_id": stringSchema(), + "name": stringSchema(), + "town": stringSchema(), + "distance_meters": map[string]any{"type": "number", "minimum": 0}, + "nearest_point": coordinateSchema(), + "length_meters": map[string]any{"type": "number"}, + "source_updated_raw": stringSchema(), + }, []string{"source_record_id", "distance_meters", "nearest_point"}) +} + +func responsibleUnitItemSchema() map[string]any { + return objectItemSchema(map[string]any{ + "grid_id": stringSchema(), + "town": stringSchema(), + "area_label": stringSchema(), + "fire_team": stringSchema(), + "live_location_available": boolSchema(), + "assembly_site_available": boolSchema(), + }, []string{"grid_id", "live_location_available", "assembly_site_available"}) +} + +func riskAreaItemSchema() map[string]any { + return objectItemSchema(map[string]any{ + "source_record_id": stringSchema(), + "risk_type": stringSchema(), + "name": stringSchema(), + "town": stringSchema(), + "village": stringSchema(), + "direction": stringSchema(), + "covers_point": boolSchema(), + "distance_meters": map[string]any{"type": "number", "minimum": 0}, + }, []string{"source_record_id", "risk_type", "covers_point", "distance_meters"}) +} + +func objectItemSchema(properties map[string]any, required []string) map[string]any { + return map[string]any{ + "type": "object", + "additionalProperties": false, + "properties": properties, + "required": required, + } +} diff --git a/internal/handler/mcp_test.go b/internal/handler/mcp_test.go new file mode 100644 index 0000000..c96bab9 --- /dev/null +++ b/internal/handler/mcp_test.go @@ -0,0 +1,479 @@ +package handler + +import ( + "context" + "encoding/json" + "errors" + "io" + "log" + "net/http" + "net/http/httptest" + "reflect" + "strings" + "testing" + "time" + + "fire-safety-ymd/internal/domain" + "fire-safety-ymd/internal/service" +) + +const testMCPToken = "0123456789abcdef0123456789abcdef" + +func TestNewMCPHandlerRejectsUnsafeOptions(t *testing.T) { + tests := []struct { + name string + options MCPOptions + }{ + {name: "short token", options: MCPOptions{AuthToken: "short", MaxBodyBytes: 1, ToolTimeout: time.Second}}, + {name: "header control token", options: MCPOptions{AuthToken: strings.Repeat("a", 31) + "\n", MaxBodyBytes: 1, ToolTimeout: time.Second}}, + {name: "oversized body setting", options: MCPOptions{AuthToken: testMCPToken, MaxBodyBytes: maximumMCPBody + 1, ToolTimeout: time.Second}}, + {name: "oversized timeout", options: MCPOptions{AuthToken: testMCPToken, MaxBodyBytes: 1, ToolTimeout: maximumToolTimeout + time.Second}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if _, err := NewMCPHandler(&fakeMCPSpatialService{}, tt.options); err == nil { + t.Fatal("NewMCPHandler() error = nil") + } + }) + } +} + +func TestMCPTransportGuards(t *testing.T) { + tests := []struct { + name string + body string + configure func(*http.Request) + wantStatus int + wantErrorCode string + }{ + { + name: "missing auth", + body: `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`, + configure: func(request *http.Request) { + request.Header.Del("Authorization") + }, + wantStatus: http.StatusUnauthorized, + wantErrorCode: "MCP_AUTH_INVALID", + }, + { + name: "wrong auth", + body: `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`, + configure: func(request *http.Request) { + request.Header.Set("Authorization", "Bearer wrong") + }, + wantStatus: http.StatusUnauthorized, + wantErrorCode: "MCP_AUTH_INVALID", + }, + { + name: "browser origin", + body: `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`, + configure: func(request *http.Request) { + request.Header.Set("Origin", "https://untrusted.example") + }, + wantStatus: http.StatusForbidden, + wantErrorCode: "MCP_ORIGIN_FORBIDDEN", + }, + { + name: "wrong content type", + body: `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`, + configure: func(request *http.Request) { + request.Header.Set("Content-Type", "text/plain") + }, + wantStatus: http.StatusUnsupportedMediaType, + wantErrorCode: "MCP_CONTENT_TYPE_INVALID", + }, + { + name: "unsupported protocol header", + body: `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`, + configure: func(request *http.Request) { + request.Header.Set("MCP-Protocol-Version", "1999-01-01") + }, + wantStatus: http.StatusBadRequest, + wantErrorCode: "MCP_PROTOCOL_VERSION_UNSUPPORTED", + }, + { + name: "invalid JSON", + body: `{"jsonrpc":`, + wantStatus: http.StatusBadRequest, + wantErrorCode: "MCP_REQUEST_INVALID", + }, + { + name: "batch rejected", + body: `[{"jsonrpc":"2.0","id":1,"method":"tools/list"}]`, + wantStatus: http.StatusBadRequest, + wantErrorCode: "MCP_REQUEST_INVALID", + }, + { + name: "invalid id", + body: `{"jsonrpc":"2.0","id":{},"method":"tools/list"}`, + wantStatus: http.StatusBadRequest, + wantErrorCode: "MCP_REQUEST_INVALID", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + handler := newTestMCPHandler(t, &fakeMCPSpatialService{}, 4096, time.Second) + request := authenticatedMCPRequest(tt.body) + if tt.configure != nil { + tt.configure(request) + } + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != tt.wantStatus { + t.Fatalf("status = %d, want %d; body=%s", response.Code, tt.wantStatus, response.Body.String()) + } + if got := response.Header().Get("X-Request-ID"); got == "" { + t.Fatal("X-Request-ID header is missing") + } + if !strings.Contains(response.Body.String(), `"code":"`+tt.wantErrorCode+`"`) { + t.Fatalf("response = %s, want code %s", response.Body.String(), tt.wantErrorCode) + } + }) + } +} + +func TestMCPRejectsOversizedBody(t *testing.T) { + handler := newTestMCPHandler(t, &fakeMCPSpatialService{}, 32, time.Second) + request := authenticatedMCPRequest(`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusRequestEntityTooLarge || !strings.Contains(response.Body.String(), "MCP_REQUEST_BODY_TOO_LARGE") { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } +} + +func TestMCPInitializeAndInitializedNotification(t *testing.T) { + handler := newTestMCPHandler(t, &fakeMCPSpatialService{}, 4096, time.Second) + + initialize := authenticatedMCPRequest(`{ + "jsonrpc":"2.0", + "id":"init-1", + "method":"initialize", + "params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"test","version":"1"}} + }`) + initializeResponse := httptest.NewRecorder() + handler.ServeHTTP(initializeResponse, initialize) + if initializeResponse.Code != http.StatusOK { + t.Fatalf("initialize status=%d body=%s", initializeResponse.Code, initializeResponse.Body.String()) + } + if !strings.Contains(initializeResponse.Body.String(), `"protocolVersion":"2025-06-18"`) || + !strings.Contains(initializeResponse.Body.String(), `"tools":{"listChanged":false}`) { + t.Fatalf("unexpected initialize response: %s", initializeResponse.Body.String()) + } + + initialized := authenticatedMCPRequest(`{"jsonrpc":"2.0","method":"notifications/initialized"}`) + initializedResponse := httptest.NewRecorder() + handler.ServeHTTP(initializedResponse, initialized) + if initializedResponse.Code != http.StatusAccepted || initializedResponse.Body.Len() != 0 { + t.Fatalf("initialized status=%d body=%q", initializedResponse.Code, initializedResponse.Body.String()) + } +} + +func TestMCPRejectsUnsupportedInitializeVersion(t *testing.T) { + handler := newTestMCPHandler(t, &fakeMCPSpatialService{}, 4096, time.Second) + request := authenticatedMCPRequest(`{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25"}}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusOK || !strings.Contains(response.Body.String(), "MCP_PROTOCOL_VERSION_UNSUPPORTED") { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } +} + +func TestMCPListsSevenBoundedReadOnlyTools(t *testing.T) { + handler := newTestMCPHandler(t, &fakeMCPSpatialService{}, 4096, time.Second) + request := authenticatedMCPRequest(`{"jsonrpc":"2.0","id":2,"method":"tools/list"}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusOK { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } + var decoded struct { + Result struct { + Tools []struct { + Name string `json:"name"` + InputSchema map[string]any `json:"inputSchema"` + Annotations map[string]any `json:"annotations"` + } `json:"tools"` + } `json:"result"` + } + if err := json.Unmarshal(response.Body.Bytes(), &decoded); err != nil { + t.Fatalf("decode response: %v", err) + } + if len(decoded.Result.Tools) != 7 { + t.Fatalf("tool count=%d, want 7", len(decoded.Result.Tools)) + } + for _, tool := range decoded.Result.Tools { + if tool.Annotations["readOnlyHint"] != true || tool.Annotations["destructiveHint"] != false || tool.Annotations["openWorldHint"] != false { + t.Fatalf("tool %s annotations=%#v", tool.Name, tool.Annotations) + } + properties, ok := tool.InputSchema["properties"].(map[string]any) + if !ok { + t.Fatalf("tool %s lacks properties schema: %#v", tool.Name, tool.InputSchema) + } + if tool.Name == toolSearchPlaceCandidates { + if properties["place_name"] == nil || properties["longitude"] != nil || properties["latitude"] != nil { + t.Fatalf("place search schema = %#v", tool.InputSchema) + } + } else if properties["longitude"] == nil || properties["latitude"] == nil { + t.Fatalf("tool %s lacks bounded coordinate schema: %#v", tool.Name, tool.InputSchema) + } + for _, forbidden := range []string{"user_id", "tenant_id", "role", "scope_mode", "allowed_towns", "sql"} { + if _, exists := properties[forbidden]; exists { + t.Fatalf("tool %s exposes forbidden authorization/query field %s", tool.Name, forbidden) + } + } + } + for _, sensitiveColumn := range []string{"bpld", "bplddh", "csjdh", "fhzddclxfs", "zbrylxfs"} { + if strings.Contains(response.Body.String(), sensitiveColumn) { + t.Fatalf("tools/list exposes sensitive source column %s", sensitiveColumn) + } + } +} + +func TestMCPDispatchesEverySpatialTool(t *testing.T) { + tests := []struct { + name string + arguments string + }{ + {name: toolSearchPlaceCandidates, arguments: `{"place_name":"观水镇","limit":2}`}, + {name: toolResolveIncidentContext, arguments: `{"longitude":121.7,"latitude":37.2}`}, + {name: toolFindNearbyWaterSources, arguments: `{"longitude":121.7,"latitude":37.2,"radius_meters":1000,"limit":2}`}, + {name: toolFindCommandPostCandidates, arguments: `{"longitude":121.7,"latitude":37.2}`}, + {name: toolListNearbyAccessLines, arguments: `{"longitude":121.7,"latitude":37.2}`}, + {name: toolGetResponsibleUnits, arguments: `{"longitude":121.7,"latitude":37.2}`}, + {name: toolFindNearbyRiskAreas, arguments: `{"longitude":121.7,"latitude":37.2}`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + spatial := &fakeMCPSpatialService{} + handler := newTestMCPHandler(t, spatial, 4096, time.Second) + request := authenticatedMCPRequest(`{"jsonrpc":"2.0","id":3,"method":"tools/call","params":{"name":"` + tt.name + `","arguments":` + tt.arguments + `}}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusOK || spatial.called != tt.name { + t.Fatalf("status=%d called=%q body=%s", response.Code, spatial.called, response.Body.String()) + } + var decoded struct { + Result struct { + Content []struct { + Text string `json:"text"` + } `json:"content"` + StructuredContent map[string]any `json:"structuredContent"` + IsError bool `json:"isError"` + } `json:"result"` + } + if err := json.Unmarshal(response.Body.Bytes(), &decoded); err != nil { + t.Fatalf("decode response: %v", err) + } + if decoded.Result.IsError || len(decoded.Result.Content) != 1 { + t.Fatalf("unexpected tool result: %s", response.Body.String()) + } + var textContent map[string]any + if err := json.Unmarshal([]byte(decoded.Result.Content[0].Text), &textContent); err != nil { + t.Fatalf("text content is not mirrored JSON: %v", err) + } + if !reflect.DeepEqual(textContent, decoded.Result.StructuredContent) { + t.Fatalf("text content and structuredContent differ: text=%#v structured=%#v", textContent, decoded.Result.StructuredContent) + } + }) + } +} + +func TestMCPToolErrorsAreStableAndSanitized(t *testing.T) { + tests := []struct { + name string + spatial *fakeMCPSpatialService + tool string + arguments string + timeout time.Duration + wantCode string + forbiddenText string + }{ + { + name: "missing place name", + spatial: &fakeMCPSpatialService{}, + tool: toolSearchPlaceCandidates, + arguments: `{"limit":2}`, + timeout: time.Second, + wantCode: "INVALID_ARGUMENT", + }, + { + name: "unknown argument", + spatial: &fakeMCPSpatialService{}, + tool: toolResolveIncidentContext, + arguments: `{"longitude":121,"latitude":37,"user_id":"admin"}`, + timeout: time.Second, + wantCode: "INVALID_ARGUMENT", + }, + { + name: "unknown tool", + spatial: &fakeMCPSpatialService{}, + tool: "execute_sql", + arguments: `{}`, + timeout: time.Second, + wantCode: "TOOL_NOT_FOUND", + }, + { + name: "repository failure", + spatial: &fakeMCPSpatialService{failure: errors.New("postgres://secret-user:secret-password@internal/db")}, + tool: toolFindNearbyWaterSources, + arguments: `{"longitude":121,"latitude":37}`, + timeout: time.Second, + wantCode: "DATA_SOURCE_UNAVAILABLE", + forbiddenText: "secret-password", + }, + { + name: "tool timeout", + spatial: &fakeMCPSpatialService{waitForCancellation: true}, + tool: toolFindNearbyWaterSources, + arguments: `{"longitude":121,"latitude":37}`, + timeout: 5 * time.Millisecond, + wantCode: "QUERY_TIMEOUT", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + handler := newTestMCPHandler(t, tt.spatial, 4096, tt.timeout) + request := authenticatedMCPRequest(`{"jsonrpc":"2.0","id":4,"method":"tools/call","params":{"name":"` + tt.tool + `","arguments":` + tt.arguments + `}}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusOK || !strings.Contains(response.Body.String(), `"isError":true`) || !strings.Contains(response.Body.String(), `"code":"`+tt.wantCode+`"`) { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } + if tt.forbiddenText != "" && strings.Contains(response.Body.String(), tt.forbiddenText) { + t.Fatalf("response leaked internal error: %s", response.Body.String()) + } + }) + } +} + +func TestMCPDoesNotExecuteIDLessToolNotification(t *testing.T) { + spatial := &fakeMCPSpatialService{} + handler := newTestMCPHandler(t, spatial, 4096, time.Second) + request := authenticatedMCPRequest(`{"jsonrpc":"2.0","method":"tools/call","params":{"name":"fire_safety_resolve_incident_context","arguments":{"longitude":121,"latitude":37}}}`) + response := httptest.NewRecorder() + + handler.ServeHTTP(response, request) + + if response.Code != http.StatusAccepted || response.Body.Len() != 0 || spatial.called != "" { + t.Fatalf("status=%d called=%q body=%q", response.Code, spatial.called, response.Body.String()) + } +} + +func newTestMCPHandler(t *testing.T, spatial MCPSpatialService, maxBodyBytes int64, timeout time.Duration) *MCPHandler { + t.Helper() + handler, err := NewMCPHandler(spatial, MCPOptions{ + AuthToken: testMCPToken, + MaxBodyBytes: maxBodyBytes, + ToolTimeout: timeout, + Logger: log.New(io.Discard, "", 0), + }) + if err != nil { + t.Fatalf("NewMCPHandler() error = %v", err) + } + return handler +} + +func authenticatedMCPRequest(body string) *http.Request { + request := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body)) + request.Header.Set("Authorization", "Bearer "+testMCPToken) + request.Header.Set("Content-Type", "application/json") + request.Header.Set("MCP-Protocol-Version", mcpProtocolVersion) + return request +} + +type fakeMCPSpatialService struct { + called string + failure error + waitForCancellation bool +} + +func (f *fakeMCPSpatialService) prepare(ctx context.Context, tool string) error { + f.called = tool + if f.waitForCancellation { + <-ctx.Done() + return ctx.Err() + } + if f.failure != nil { + return errors.Join(service.ErrDataSourceUnavailable, f.failure) + } + return nil +} + +func (f *fakeMCPSpatialService) SearchPlaceCandidates(ctx context.Context, _ string, _ int) (service.QueryResult[[]domain.PlaceCandidate], error) { + if err := f.prepare(ctx, toolSearchPlaceCandidates); err != nil { + return service.QueryResult[[]domain.PlaceCandidate]{}, err + } + return emptyMCPResult([]domain.PlaceCandidate{}), nil +} + +func (f *fakeMCPSpatialService) ResolveIncidentContext(ctx context.Context, _ domain.Coordinate) (service.QueryResult[[]domain.IncidentContext], error) { + if err := f.prepare(ctx, toolResolveIncidentContext); err != nil { + return service.QueryResult[[]domain.IncidentContext]{}, err + } + return emptyMCPResult([]domain.IncidentContext{}), nil +} + +func (f *fakeMCPSpatialService) FindNearbyWaterSources(ctx context.Context, _ domain.Coordinate, _ float64, _ int) (service.QueryResult[[]domain.WaterSource], error) { + if err := f.prepare(ctx, toolFindNearbyWaterSources); err != nil { + return service.QueryResult[[]domain.WaterSource]{}, err + } + return emptyMCPResult([]domain.WaterSource{}), nil +} + +func (f *fakeMCPSpatialService) FindCommandPostCandidates(ctx context.Context, _ domain.Coordinate, _ float64, _ int) (service.QueryResult[[]domain.CommandPostCandidate], error) { + if err := f.prepare(ctx, toolFindCommandPostCandidates); err != nil { + return service.QueryResult[[]domain.CommandPostCandidate]{}, err + } + return emptyMCPResult([]domain.CommandPostCandidate{}), nil +} + +func (f *fakeMCPSpatialService) ListNearbyAccessLines(ctx context.Context, _ domain.Coordinate, _ float64, _ int) (service.QueryResult[[]domain.AccessLine], error) { + if err := f.prepare(ctx, toolListNearbyAccessLines); err != nil { + return service.QueryResult[[]domain.AccessLine]{}, err + } + return emptyMCPResult([]domain.AccessLine{}), nil +} + +func (f *fakeMCPSpatialService) GetResponsibleUnits(ctx context.Context, _ domain.Coordinate) (service.QueryResult[[]domain.ResponsibleUnit], error) { + if err := f.prepare(ctx, toolGetResponsibleUnits); err != nil { + return service.QueryResult[[]domain.ResponsibleUnit]{}, err + } + return emptyMCPResult([]domain.ResponsibleUnit{}), nil +} + +func (f *fakeMCPSpatialService) FindNearbyRiskAreas(ctx context.Context, _ domain.Coordinate, _ float64, _ int) (service.QueryResult[[]domain.RiskArea], error) { + if err := f.prepare(ctx, toolFindNearbyRiskAreas); err != nil { + return service.QueryResult[[]domain.RiskArea]{}, err + } + return emptyMCPResult([]domain.RiskArea{}), nil +} + +func emptyMCPResult[T any](data []T) service.QueryResult[[]T] { + return service.QueryResult[[]T]{ + Status: "no_results", + Data: data, + Metadata: service.ResultMetadata{ + GeneratedAt: "2026-09-04T00:00:00Z", + DataSources: []string{}, + SpatialReference: "EPSG:4326", + ResultCount: 0, + }, + Warnings: []string{}, + } +} diff --git a/internal/integration/superagent/chat.go b/internal/integration/superagent/chat.go new file mode 100644 index 0000000..3a4f323 --- /dev/null +++ b/internal/integration/superagent/chat.go @@ -0,0 +1,100 @@ +package superagent + +import ( + "context" + "errors" + "fmt" + + "fire-safety-ymd/internal/service" +) + +// ChatAgentAdapter projects the SuperAgent-specific client into the +// provider-neutral service boundary. +type ChatAgentAdapter struct { + client Client +} + +var _ service.ChatAgent = (*ChatAgentAdapter)(nil) + +// NewChatAgentAdapter constructs the chat service adapter. +func NewChatAgentAdapter(client Client) (*ChatAgentAdapter, error) { + if client == nil { + return nil, errors.New("superagent client is required") + } + return &ChatAgentAdapter{client: client}, nil +} + +// CreateSession maps only the server-controlled fields required by SuperAgent. +func (a *ChatAgentAdapter) CreateSession(ctx context.Context, request service.AgentCreateSessionRequest) (service.AgentSession, error) { + session, err := a.client.CreateSession(ctx, CreateSessionRequest{ + ExternalSubjectID: request.ExternalSubjectID, + IdempotencyKey: request.IdempotencyKey, + RequestID: request.RequestID, + Metadata: request.Metadata, + }) + if err != nil { + return service.AgentSession{}, mapChatError("create session", err) + } + return service.AgentSession{ID: session.ID}, nil +} + +// StreamMessage removes provider identifiers, text deltas, and raw trace data +// from progress events. The final answer is returned only when the underlying +// strict SuperAgent client reports success. +func (a *ChatAgentAdapter) StreamMessage( + ctx context.Context, + request service.AgentMessageRequest, + traceHandler func(service.AgentTraceEvent), +) (service.AgentMessageResult, error) { + result, err := a.client.StreamMessage(ctx, StreamMessageRequest{ + SessionID: request.SessionID, + Message: request.Message, + IdempotencyKey: request.IdempotencyKey, + RequestID: request.RequestID, + Metadata: request.Metadata, + }, func(event TraceEvent) { + if traceHandler == nil { + return + } + traceHandler(service.AgentTraceEvent{ + Event: safeValue(event.Event), + ToolName: safeValue(event.ToolName), + Status: safeValue(event.Status), + }) + }) + if err != nil { + return service.AgentMessageResult{}, mapChatError("stream message", err) + } + return service.AgentMessageResult{ + RunID: result.RunID, + ModelID: safeValue(result.ModelName), + Answer: result.Answer, + Usage: service.ChatTokenUsage{ + Input: result.Usage.Input, + Output: result.Usage.Output, + Total: result.Usage.Total, + }, + }, nil +} + +func mapChatError(operation string, err error) error { + if err == nil { + return nil + } + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return err + } + var public error + switch { + case errors.Is(err, ErrRunFailed): + public = service.ErrChatRunFailed + case errors.Is(err, ErrProtocol), errors.Is(err, ErrStreamIncomplete), errors.Is(err, ErrStreamRead), + errors.Is(err, ErrRecoveryExhausted), errors.Is(err, ErrInvalidRequest): + public = service.ErrChatUpstreamProtocol + case errors.Is(err, ErrDisabled), errors.Is(err, ErrHTTPStatus): + public = service.ErrChatUpstreamUnavailable + default: + public = service.ErrChatUpstreamUnavailable + } + return fmt.Errorf("superagent %s: %w", operation, public) +} diff --git a/internal/integration/superagent/chat_test.go b/internal/integration/superagent/chat_test.go new file mode 100644 index 0000000..6e3bf3b --- /dev/null +++ b/internal/integration/superagent/chat_test.go @@ -0,0 +1,115 @@ +package superagent + +import ( + "context" + "errors" + "testing" + + "fire-safety-ymd/internal/service" +) + +func TestChatAgentAdapterMapsRequestsAndRemovesSensitiveTraceFields(t *testing.T) { + client := &fakeProviderClient{ + session: Session{ID: "provider-session-1"}, + result: Result{ + RunID: "run-1", + ModelName: "qwen-plus-latest", + Answer: "safe final answer", + Usage: TokenUsage{Input: 1, Output: 2, Total: 3}, + }, + trace: TraceEvent{ + Event: "tool.started", + RunID: "provider-run-id", + MessageID: "provider-message-id", + ToolCallID: "provider-tool-call-id", + ToolName: "fire_safety_test", + Text: "raw text must not cross the adapter", + Status: "running", + Timestamp: "2026-09-05T00:00:00Z", + }, + } + adapter, err := NewChatAgentAdapter(client) + if err != nil { + t.Fatalf("NewChatAgentAdapter() error = %v", err) + } + + session, err := adapter.CreateSession(context.Background(), service.AgentCreateSessionRequest{ + ExternalSubjectID: "test-subject", + IdempotencyKey: "session-key", + RequestID: "request-1", + Metadata: map[string]any{"source": "test"}, + }) + if err != nil || session.ID != "provider-session-1" { + t.Fatalf("CreateSession() = %#v, %v", session, err) + } + var trace service.AgentTraceEvent + result, err := adapter.StreamMessage(context.Background(), service.AgentMessageRequest{ + SessionID: session.ID, + Message: "question", + IdempotencyKey: "turn-key", + RequestID: "request-1", + }, func(event service.AgentTraceEvent) { trace = event }) + if err != nil { + t.Fatalf("StreamMessage() error = %v", err) + } + if trace != (service.AgentTraceEvent{Event: "tool.started", ToolName: "fire_safety_test", Status: "running"}) { + t.Fatalf("projected trace = %#v", trace) + } + if result.Answer != "safe final answer" || result.RunID != "run-1" || result.ModelID != "qwen-plus-latest" || result.Usage.Total != 3 { + t.Fatalf("result = %#v", result) + } + if client.createRequest.ExternalSubjectID != "test-subject" || client.messageRequest.SessionID != "provider-session-1" || client.messageRequest.Message != "question" { + t.Fatalf("provider requests = %#v / %#v", client.createRequest, client.messageRequest) + } +} + +func TestMapChatErrorUsesStableServiceCategories(t *testing.T) { + tests := []struct { + input error + want error + }{ + {input: ErrDisabled, want: service.ErrChatUpstreamUnavailable}, + {input: &HTTPStatusError{StatusCode: 503, ProviderCode: "unavailable"}, want: service.ErrChatUpstreamUnavailable}, + {input: ErrProtocol, want: service.ErrChatUpstreamProtocol}, + {input: ErrStreamIncomplete, want: service.ErrChatUpstreamProtocol}, + {input: ErrRecoveryExhausted, want: service.ErrChatUpstreamProtocol}, + {input: &RunError{Code: "failed"}, want: service.ErrChatRunFailed}, + {input: context.DeadlineExceeded, want: context.DeadlineExceeded}, + {input: context.Canceled, want: context.Canceled}, + } + for _, tt := range tests { + mapped := mapChatError("test", tt.input) + if !errors.Is(mapped, tt.want) { + t.Fatalf("mapChatError(%v) = %v, want %v", tt.input, mapped, tt.want) + } + } +} + +func TestNewChatAgentAdapterRejectsNilClient(t *testing.T) { + if _, err := NewChatAgentAdapter(nil); err == nil { + t.Fatal("NewChatAgentAdapter(nil) error = nil") + } +} + +type fakeProviderClient struct { + session Session + result Result + trace TraceEvent + createErr error + streamErr error + createRequest CreateSessionRequest + messageRequest StreamMessageRequest +} + +func (f *fakeProviderClient) CreateSession(_ context.Context, request CreateSessionRequest) (Session, error) { + f.createRequest = request + return f.session, f.createErr +} + +func (f *fakeProviderClient) StreamMessage(_ context.Context, request StreamMessageRequest, trace TraceHandler) (Result, error) { + f.messageRequest = request + if trace != nil { + trace(f.trace) + } + return f.result, f.streamErr +} diff --git a/internal/integration/superagent/client.go b/internal/integration/superagent/client.go new file mode 100644 index 0000000..535b59d --- /dev/null +++ b/internal/integration/superagent/client.go @@ -0,0 +1,661 @@ +package superagent + +import ( + "bytes" + "context" + "crypto/rand" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "mime" + "net" + "net/http" + "net/url" + "strings" + "time" +) + +const ( + defaultConnectTimeout = 15 * time.Second + defaultRecoveryInitialBackoff = 250 * time.Millisecond + defaultMaxMessageBytes int64 = 64 * 1024 + maximumMessageBytes int64 = 16 * 1024 * 1024 + maxControlResponseBytes int64 = 1024 * 1024 + maxRequestEnvelopeBytes = 256 * 1024 + maxIdentifierBytes = 512 + maximumRecoveryAttempts = 20 + maxRecoveryBackoff = 30 * time.Second +) + +// HTTPClient implements Client with the SuperAgent Open API HTTP protocol. +type HTTPClient struct { + enabled bool + baseURL *url.URL + apiKey string + httpClient *http.Client + recoveryMaxAttempts int + recoveryInitialBackoff time.Duration + maxMessageBytes int64 +} + +var _ Client = (*HTTPClient)(nil) + +// NewHTTPClient validates the provider boundary and constructs a client. It +// deliberately has no overall HTTP timeout because an SSE run can be long-lived; +// callers must provide a bounded context. +func NewHTTPClient(cfg Config) (*HTTPClient, error) { + if cfg.ConnectTimeout == 0 { + cfg.ConnectTimeout = defaultConnectTimeout + } + if cfg.RecoveryInitialBackoff == 0 { + cfg.RecoveryInitialBackoff = defaultRecoveryInitialBackoff + } + if cfg.MaxMessageBytes == 0 { + cfg.MaxMessageBytes = defaultMaxMessageBytes + } + if cfg.ConnectTimeout < 0 || cfg.RecoveryInitialBackoff < 0 || cfg.MaxMessageBytes < 0 || cfg.RecoveryMaxAttempts < 0 { + return nil, fmt.Errorf("%w: timing, size, and attempt limits must not be negative", ErrInvalidConfig) + } + if cfg.MaxMessageBytes > maximumMessageBytes { + return nil, fmt.Errorf("%w: message byte limit exceeds the supported maximum", ErrInvalidConfig) + } + if cfg.RecoveryMaxAttempts > maximumRecoveryAttempts { + return nil, fmt.Errorf("%w: recovery attempt limit exceeds the supported maximum", ErrInvalidConfig) + } + if !cfg.Enabled { + return &HTTPClient{ + enabled: false, + recoveryMaxAttempts: cfg.RecoveryMaxAttempts, + recoveryInitialBackoff: cfg.RecoveryInitialBackoff, + maxMessageBytes: cfg.MaxMessageBytes, + }, nil + } + + baseURL, err := validateBaseURL(cfg.BaseURL) + if err != nil { + return nil, err + } + if !validAPIKey(cfg.APIKey) { + return nil, fmt.Errorf("%w: API key is required and must be a valid header token", ErrInvalidConfig) + } + + transport := http.DefaultTransport.(*http.Transport).Clone() + transport.DialContext = (&net.Dialer{ + Timeout: cfg.ConnectTimeout, + KeepAlive: 30 * time.Second, + }).DialContext + + return &HTTPClient{ + enabled: true, + baseURL: baseURL, + apiKey: cfg.APIKey, + httpClient: &http.Client{ + Transport: transport, + CheckRedirect: func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + }, + }, + recoveryMaxAttempts: cfg.RecoveryMaxAttempts, + recoveryInitialBackoff: cfg.RecoveryInitialBackoff, + maxMessageBytes: cfg.MaxMessageBytes, + }, nil +} + +// CreateSession creates a provider session without sending a message. +func (c *HTTPClient) CreateSession(ctx context.Context, request CreateSessionRequest) (Session, error) { + if !c.enabled { + return Session{}, ErrDisabled + } + if err := validateCreateSessionRequest(request); err != nil { + return Session{}, err + } + + body, err := json.Marshal(struct { + ExternalSubjectID string `json:"external_subject_id"` + IdempotencyKey string `json:"idempotency_key"` + Metadata map[string]any `json:"metadata,omitempty"` + }{ + ExternalSubjectID: request.ExternalSubjectID, + IdempotencyKey: request.IdempotencyKey, + Metadata: request.Metadata, + }) + if err != nil { + return Session{}, fmt.Errorf("%w: metadata is not JSON encodable", ErrInvalidRequest) + } + if len(body) > maxRequestEnvelopeBytes { + return Session{}, fmt.Errorf("%w: create-session payload exceeds size limit", ErrInvalidRequest) + } + + httpRequest, err := c.newRequest( + ctx, + http.MethodPost, + c.endpoint("/api/open/agent-sessions"), + bytes.NewReader(body), + "application/json", + request.RequestID, + "", + ) + if err != nil { + return Session{}, err + } + httpRequest.Header.Set("Content-Type", "application/json") + + response, err := c.httpClient.Do(httpRequest) + if err != nil { + if ctx.Err() != nil { + return Session{}, ctx.Err() + } + return Session{}, fmt.Errorf("create superagent session: %w", err) + } + defer response.Body.Close() + if err := requireHTTPSuccess(response); err != nil { + return Session{}, err + } + + payload, err := readJSONObject(response.Body) + if err != nil { + return Session{}, err + } + sessionID := stringValue(payload, "session_id") + if sessionID == "" { + sessionID = stringValue(payload, "id") + } + if !validResourceID(sessionID) { + return Session{}, fmt.Errorf("%w: create-session response is missing a valid session ID", ErrProtocol) + } + return Session{ID: sessionID, Status: safeValue(stringValue(payload, "status"))}, nil +} + +// StreamMessage sends one message and returns only after strict stream success. +func (c *HTTPClient) StreamMessage( + ctx context.Context, + request StreamMessageRequest, + traceHandler TraceHandler, +) (Result, error) { + if !c.enabled { + return Result{}, ErrDisabled + } + if err := c.validateStreamMessageRequest(request); err != nil { + return Result{}, err + } + + body, err := json.Marshal(struct { + Message string `json:"message"` + IdempotencyKey string `json:"idempotency_key"` + Metadata map[string]any `json:"metadata,omitempty"` + }{ + Message: request.Message, + IdempotencyKey: request.IdempotencyKey, + Metadata: request.Metadata, + }) + if err != nil { + return Result{}, fmt.Errorf("%w: metadata is not JSON encodable", ErrInvalidRequest) + } + if int64(len(body)) > c.maxMessageBytes+maxRequestEnvelopeBytes { + return Result{}, fmt.Errorf("%w: message payload exceeds size limit", ErrInvalidRequest) + } + + streamURL := c.endpoint("/api/open/agent-sessions/" + request.SessionID + "/messages/stream") + streamURL.RawQuery = "include_trace=true" + httpRequest, err := c.newRequest( + ctx, + http.MethodPost, + streamURL, + bytes.NewReader(body), + "text/event-stream", + request.RequestID, + "", + ) + if err != nil { + return Result{}, err + } + httpRequest.Header.Set("Content-Type", "application/json") + + response, err := c.httpClient.Do(httpRequest) + if err != nil { + if ctx.Err() != nil { + return Result{}, ctx.Err() + } + return Result{}, fmt.Errorf("stream superagent message: %w", err) + } + if err := requireHTTPSuccess(response); err != nil { + response.Body.Close() + return Result{}, err + } + if err := requireEventStream(response); err != nil { + response.Body.Close() + return Result{}, err + } + + state := newStreamState(request.SessionID) + if contentLocation := strings.TrimSpace(response.Header.Get("Content-Location")); contentLocation != "" { + runURL, locationErr := c.resolveRunURL(httpRequest.URL, request.SessionID, contentLocation) + if locationErr != nil { + response.Body.Close() + return Result{}, locationErr + } + state.runURL = runURL.String() + state.runID = runIDFromURL(runURL) + } + + consumeErr := state.consume(response.Body, traceHandler) + closeErr := response.Body.Close() + if consumeErr == nil && closeErr != nil { + consumeErr = ErrStreamRead + } + if consumeErr == nil { + return state.result() + } + if ctx.Err() != nil { + return Result{}, ctx.Err() + } + if !errorsIsRecoverableStream(consumeErr) { + return Result{}, consumeErr + } + if state.runURL == "" && validResourceID(state.runID) { + state.runURL = c.endpoint("/api/open/agent-sessions/" + request.SessionID + "/runs/" + state.runID).String() + } + if state.runURL == "" { + return Result{}, consumeErr + } + if err := c.recover(ctx, request.RequestID, state, traceHandler); err != nil { + return Result{}, err + } + return state.result() +} + +func (c *HTTPClient) recover( + ctx context.Context, + requestID string, + state *streamState, + traceHandler TraceHandler, +) error { + runURL, err := url.Parse(state.runURL) + if err != nil || !c.sameOrigin(runURL) || runURL.User != nil || runURL.Fragment != "" { + return fmt.Errorf("%w: invalid recovery URL", ErrProtocol) + } + backoff := c.recoveryInitialBackoff + var lastErr error + for attempt := 0; attempt < c.recoveryMaxAttempts; attempt++ { + if err := waitForRecovery(ctx, backoff); err != nil { + return err + } + status, statusErr := c.queryRunStatus(ctx, runURL, requestID) + if statusErr != nil { + if ctx.Err() != nil { + return ctx.Err() + } + if !errors.Is(statusErr, ErrStreamRead) && !isRetryableHTTPError(statusErr) { + return statusErr + } + lastErr = statusErr + backoff = nextBackoff(backoff) + continue + } + if isFailedRunStatus(status) { + return &RunError{Status: status} + } + + eventsErr := c.subscribeRunEvents(ctx, runURL, requestID, state, traceHandler) + if eventsErr == nil && state.endSeen { + return nil + } + if ctx.Err() != nil { + return ctx.Err() + } + if eventsErr != nil && !errorsIsRecoverableStream(eventsErr) && !isRetryableHTTPError(eventsErr) { + return eventsErr + } + lastErr = eventsErr + if lastErr == nil { + lastErr = ErrStreamIncomplete + } + backoff = nextBackoff(backoff) + } + if lastErr == nil { + lastErr = ErrStreamIncomplete + } + return fmt.Errorf("%w: %w", ErrRecoveryExhausted, lastErr) +} + +func (c *HTTPClient) queryRunStatus(ctx context.Context, runURL *url.URL, requestID string) (string, error) { + request, err := c.newRequest(ctx, http.MethodGet, runURL, nil, "application/json", requestID, "") + if err != nil { + return "", err + } + response, err := c.httpClient.Do(request) + if err != nil { + if ctx.Err() != nil { + return "", ctx.Err() + } + return "", fmt.Errorf("%w: query run request failed", ErrStreamRead) + } + defer response.Body.Close() + if err := requireHTTPSuccess(response); err != nil { + return "", err + } + payload, err := readJSONObject(response.Body) + if err != nil { + return "", err + } + status := safeValue(stringValue(payload, "status")) + if status == "" { + status = "running" + } + return strings.ToLower(status), nil +} + +func (c *HTTPClient) subscribeRunEvents( + ctx context.Context, + runURL *url.URL, + requestID string, + state *streamState, + traceHandler TraceHandler, +) error { + eventsURL := *runURL + eventsURL.Path = strings.TrimRight(eventsURL.Path, "/") + "/events" + eventsURL.RawPath = "" + request, err := c.newRequest( + ctx, + http.MethodGet, + &eventsURL, + nil, + "text/event-stream", + requestID, + state.lastEventID, + ) + if err != nil { + return err + } + response, err := c.httpClient.Do(request) + if err != nil { + if ctx.Err() != nil { + return ctx.Err() + } + return fmt.Errorf("%w: subscribe request failed", ErrStreamRead) + } + defer response.Body.Close() + if err := requireHTTPSuccess(response); err != nil { + return err + } + if err := requireEventStream(response); err != nil { + return err + } + return state.consume(response.Body, traceHandler) +} + +func (c *HTTPClient) newRequest( + ctx context.Context, + method string, + target *url.URL, + body io.Reader, + accept string, + requestID string, + lastEventID string, +) (*http.Request, error) { + if target == nil || !c.sameOrigin(target) || target.User != nil || target.Fragment != "" { + return nil, fmt.Errorf("%w: outbound URL is outside the configured origin", ErrInvalidRequest) + } + if !validOptionalHeader(requestID) || !validOptionalHeader(lastEventID) { + return nil, fmt.Errorf("%w: request contains an invalid header value", ErrInvalidRequest) + } + csrfToken, err := newCSRFToken() + if err != nil { + return nil, fmt.Errorf("%w: could not create request token", ErrProtocol) + } + request, err := http.NewRequestWithContext(ctx, method, target.String(), body) + if err != nil { + return nil, fmt.Errorf("%w: could not construct HTTP request", ErrInvalidRequest) + } + request.Header.Set("Authorization", "Bearer "+c.apiKey) + request.Header.Set("Cache-Control", "no-cache") + request.Header.Set("Accept", accept) + request.Header.Set("X-CSRF-Token", csrfToken) + request.AddCookie(&http.Cookie{Name: "csrf_token", Value: csrfToken}) + if requestID != "" { + request.Header.Set("X-Request-ID", requestID) + } + if lastEventID != "" { + request.Header.Set("Last-Event-ID", lastEventID) + } + return request, nil +} + +func (c *HTTPClient) endpoint(path string) *url.URL { + target := *c.baseURL + target.Path = strings.TrimRight(target.Path, "/") + path + target.RawPath = "" + target.RawQuery = "" + target.Fragment = "" + return &target +} + +func (c *HTTPClient) resolveRunURL(requestURL *url.URL, sessionID, contentLocation string) (*url.URL, error) { + reference, err := url.Parse(contentLocation) + if err != nil || reference.User != nil || reference.Fragment != "" { + return nil, fmt.Errorf("%w: invalid Content-Location", ErrProtocol) + } + resolved := requestURL.ResolveReference(reference) + if !c.sameOrigin(resolved) { + return nil, fmt.Errorf("%w: cross-origin Content-Location", ErrProtocol) + } + runID := runIDFromURL(resolved) + if !validResourceID(runID) { + return nil, fmt.Errorf("%w: Content-Location does not identify a valid run", ErrProtocol) + } + expectedSuffix := "/agent-sessions/" + sessionID + "/runs/" + runID + if !strings.HasSuffix(strings.TrimRight(resolved.Path, "/"), expectedSuffix) { + return nil, fmt.Errorf("%w: Content-Location identifies a different session", ErrProtocol) + } + return resolved, nil +} + +func (c *HTTPClient) sameOrigin(target *url.URL) bool { + if c.baseURL == nil || target == nil { + return false + } + return strings.EqualFold(c.baseURL.Scheme, target.Scheme) && + strings.EqualFold(c.baseURL.Hostname(), target.Hostname()) && + effectivePort(c.baseURL) == effectivePort(target) +} + +func validateBaseURL(raw string) (*url.URL, error) { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { + return nil, fmt.Errorf("%w: base URL must be an absolute HTTP(S) URL", ErrInvalidConfig) + } + if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { + return nil, fmt.Errorf("%w: base URL must not contain user information, query, or fragment", ErrInvalidConfig) + } + parsed.Path = strings.TrimRight(parsed.Path, "/") + parsed.RawPath = "" + return parsed, nil +} + +func validateCreateSessionRequest(request CreateSessionRequest) error { + if strings.TrimSpace(request.ExternalSubjectID) == "" || len(request.ExternalSubjectID) > maxIdentifierBytes { + return fmt.Errorf("%w: external subject ID is required and bounded", ErrInvalidRequest) + } + if strings.TrimSpace(request.IdempotencyKey) == "" || len(request.IdempotencyKey) > maxIdentifierBytes { + return fmt.Errorf("%w: idempotency key is required and bounded", ErrInvalidRequest) + } + if !validOptionalHeader(request.RequestID) { + return fmt.Errorf("%w: request ID is not a valid header value", ErrInvalidRequest) + } + return nil +} + +func (c *HTTPClient) validateStreamMessageRequest(request StreamMessageRequest) error { + if !validResourceID(request.SessionID) { + return fmt.Errorf("%w: session ID is invalid", ErrInvalidRequest) + } + if strings.TrimSpace(request.Message) == "" { + return fmt.Errorf("%w: message is required", ErrInvalidRequest) + } + if int64(len([]byte(request.Message))) > c.maxMessageBytes { + return fmt.Errorf("%w: message exceeds configured byte limit", ErrInvalidRequest) + } + if strings.TrimSpace(request.IdempotencyKey) == "" || len(request.IdempotencyKey) > maxIdentifierBytes { + return fmt.Errorf("%w: idempotency key is required and bounded", ErrInvalidRequest) + } + if !validOptionalHeader(request.RequestID) { + return fmt.Errorf("%w: request ID is not a valid header value", ErrInvalidRequest) + } + return nil +} + +func requireHTTPSuccess(response *http.Response) error { + if response.StatusCode >= http.StatusOK && response.StatusCode < http.StatusMultipleChoices { + return nil + } + body, _ := readBounded(response.Body, maxControlResponseBytes) + return &HTTPStatusError{StatusCode: response.StatusCode, ProviderCode: providerCode(body)} +} + +func requireEventStream(response *http.Response) error { + mediaType, _, err := mime.ParseMediaType(response.Header.Get("Content-Type")) + if err != nil || !strings.EqualFold(mediaType, "text/event-stream") { + return fmt.Errorf("%w: expected text/event-stream response", ErrProtocol) + } + return nil +} + +func readJSONObject(reader io.Reader) (map[string]any, error) { + body, err := readBounded(reader, maxControlResponseBytes) + if err != nil { + return nil, err + } + return decodeJSONObject(string(body)) +} + +func readBounded(reader io.Reader, limit int64) ([]byte, error) { + body, err := io.ReadAll(io.LimitReader(reader, limit+1)) + if err != nil { + return nil, fmt.Errorf("%w: could not read control response", ErrProtocol) + } + if int64(len(body)) > limit { + return nil, fmt.Errorf("%w: control response exceeds size limit", ErrProtocol) + } + return body, nil +} + +func providerCode(body []byte) string { + var payload map[string]any + if err := json.Unmarshal(body, &payload); err != nil { + return "" + } + code := stringValue(payload, "code") + if code == "" { + if nested, ok := payload["error"].(map[string]any); ok { + code = stringValue(nested, "code") + } + } + return safeValue(code) +} + +func runIDFromURL(runURL *url.URL) string { + if runURL == nil { + return "" + } + path := strings.TrimRight(runURL.Path, "/") + marker := strings.LastIndex(path, "/runs/") + if marker < 0 { + return "" + } + runID := path[marker+len("/runs/"):] + if strings.Contains(runID, "/") { + return "" + } + return runID +} + +func newCSRFToken() (string, error) { + bytes := make([]byte, 32) + if _, err := rand.Read(bytes); err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(bytes), nil +} + +func validResourceID(value string) bool { + return safeProviderValue.MatchString(value) +} + +func validOptionalHeader(value string) bool { + if len(value) > maxIdentifierBytes { + return false + } + for _, character := range value { + if character < 0x20 || character == 0x7f { + return false + } + } + return true +} + +func validAPIKey(value string) bool { + if value == "" || len(value) > 4096 { + return false + } + for _, character := range value { + if character <= 0x20 || character >= 0x7f { + return false + } + } + return true +} + +func effectivePort(value *url.URL) string { + if port := value.Port(); port != "" { + return port + } + if strings.EqualFold(value.Scheme, "https") { + return "443" + } + return "80" +} + +func isFailedRunStatus(status string) bool { + switch strings.ToLower(status) { + case "error", "failed", "timeout", "interrupted", "cancelled", "canceled": + return true + default: + return false + } +} + +func errorsIsRecoverableStream(err error) bool { + return errors.Is(err, ErrStreamIncomplete) || errors.Is(err, ErrStreamRead) +} + +func isRetryableHTTPError(err error) bool { + var statusError *HTTPStatusError + if !errors.As(err, &statusError) { + return false + } + return statusError.StatusCode == http.StatusRequestTimeout || + statusError.StatusCode == http.StatusConflict || + statusError.StatusCode == http.StatusTooEarly || + statusError.StatusCode == http.StatusTooManyRequests || + statusError.StatusCode >= http.StatusInternalServerError +} + +func waitForRecovery(ctx context.Context, duration time.Duration) error { + timer := time.NewTimer(duration) + defer timer.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return nil + } +} + +func nextBackoff(current time.Duration) time.Duration { + if current >= maxRecoveryBackoff/2 { + return maxRecoveryBackoff + } + return current * 2 +} diff --git a/internal/integration/superagent/client_test.go b/internal/integration/superagent/client_test.go new file mode 100644 index 0000000..0041ee7 --- /dev/null +++ b/internal/integration/superagent/client_test.go @@ -0,0 +1,451 @@ +package superagent + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" +) + +const testAPIKey = "test-open-api-key" + +func TestHTTPClientCreatesSessionAndStreamsMessage(t *testing.T) { + var csrfTokens []string + var tokensMu sync.Mutex + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + csrf := requireProviderHeaders(t, request, "request-1") + tokensMu.Lock() + csrfTokens = append(csrfTokens, csrf) + tokensMu.Unlock() + + switch { + case request.Method == http.MethodPost && request.URL.Path == "/api/open/agent-sessions": + if request.Header.Get("Accept") != "application/json" || request.Header.Get("Content-Type") != "application/json" { + t.Errorf("unexpected create content headers: %#v", request.Header) + } + var payload map[string]any + if err := json.NewDecoder(request.Body).Decode(&payload); err != nil { + t.Errorf("decode create body: %v", err) + writer.WriteHeader(http.StatusBadRequest) + return + } + if payload["external_subject_id"] != "probe-subject" || payload["idempotency_key"] != "session-key" { + t.Errorf("unexpected create body: %#v", payload) + } + writer.Header().Set("Content-Type", "application/json") + fmt.Fprint(writer, `{"session_id":"session-1","status":"active"}`) + case request.Method == http.MethodPost && request.URL.Path == "/api/open/agent-sessions/session-1/messages/stream": + if request.URL.Query().Get("include_trace") != "true" || request.Header.Get("Accept") != "text/event-stream" { + t.Errorf("unexpected stream request: %s %#v", request.URL.String(), request.Header) + } + var payload map[string]any + if err := json.NewDecoder(request.Body).Decode(&payload); err != nil { + t.Errorf("decode stream body: %v", err) + writer.WriteHeader(http.StatusBadRequest) + return + } + if payload["message"] != "safe probe" || payload["idempotency_key"] != "message-key" { + t.Errorf("unexpected stream body: %#v", payload) + } + writer.Header().Set("Content-Type", "text/event-stream; charset=utf-8") + writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1") + fmt.Fprint(writer, successfulSSE("run-1", "connectivity OK")) + default: + t.Errorf("unexpected request: %s %s", request.Method, request.URL.String()) + writer.WriteHeader(http.StatusNotFound) + } + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 2, time.Millisecond) + session, err := client.CreateSession(context.Background(), CreateSessionRequest{ + ExternalSubjectID: "probe-subject", + IdempotencyKey: "session-key", + RequestID: "request-1", + Metadata: map[string]any{"purpose": "test"}, + }) + if err != nil { + t.Fatalf("CreateSession() error = %v", err) + } + if session != (Session{ID: "session-1", Status: "active"}) { + t.Fatalf("unexpected session: %#v", session) + } + + var traces []TraceEvent + result, err := client.StreamMessage(context.Background(), StreamMessageRequest{ + SessionID: session.ID, + Message: "safe probe", + IdempotencyKey: "message-key", + RequestID: "request-1", + Metadata: map[string]any{"purpose": "test"}, + }, func(event TraceEvent) { + traces = append(traces, event) + }) + if err != nil { + t.Fatalf("StreamMessage() error = %v", err) + } + if result.Answer != "connectivity OK" || result.SessionID != "session-1" || result.RunID != "run-1" { + t.Fatalf("unexpected result: %#v", result) + } + if len(traces) != 2 { + t.Fatalf("trace count = %d, want 2", len(traces)) + } + + tokensMu.Lock() + defer tokensMu.Unlock() + if len(csrfTokens) != 2 || csrfTokens[0] == csrfTokens[1] { + t.Fatalf("CSRF tokens must be non-empty and unique per request: %#v", csrfTokens) + } +} + +func TestHTTPClientRecoversWithoutRepostingMessage(t *testing.T) { + var messagePosts atomic.Int32 + var runQueries atomic.Int32 + var eventQueries atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + requireProviderHeaders(t, request, "request-recovery") + switch { + case request.Method == http.MethodPost && request.URL.Path == "/api/open/agent-sessions/session-1/messages/stream": + messagePosts.Add(1) + writer.Header().Set("Content-Type", "text/event-stream") + writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1") + fmt.Fprint(writer, "event: trace\nid: event-1\ndata: {\"event\":\"message.delta\",\"run_id\":\"run-1\",\"text\":\"partial\"}\n\n") + case request.Method == http.MethodGet && request.URL.Path == "/api/open/agent-sessions/session-1/runs/run-1": + runQueries.Add(1) + writer.Header().Set("Content-Type", "application/json") + fmt.Fprint(writer, `{"status":"running"}`) + case request.Method == http.MethodGet && request.URL.Path == "/api/open/agent-sessions/session-1/runs/run-1/events": + eventQueries.Add(1) + if request.Header.Get("Last-Event-ID") != "event-1" { + t.Errorf("Last-Event-ID = %q, want event-1", request.Header.Get("Last-Event-ID")) + } + writer.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(writer, "event: trace\nid: event-1\ndata: {\"event\":\"message.delta\",\"text\":\"duplicate\"}\n\n") + fmt.Fprint(writer, "event: trace\nid: event-2\ndata: {\"event\":\"message.final\",\"text\":\"recovered answer\"}\n\n") + fmt.Fprint(writer, "event: trace\nid: event-3\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n") + fmt.Fprint(writer, "event: end\nid: event-4\n\n") + default: + t.Errorf("unexpected request: %s %s", request.Method, request.URL.Path) + writer.WriteHeader(http.StatusNotFound) + } + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 2, time.Millisecond) + result, err := client.StreamMessage(context.Background(), StreamMessageRequest{ + SessionID: "session-1", + Message: "safe probe", + IdempotencyKey: "message-key", + RequestID: "request-recovery", + }, nil) + if err != nil { + t.Fatalf("StreamMessage() error = %v", err) + } + if result.Answer != "recovered answer" || result.LastEventID != "event-4" { + t.Fatalf("unexpected recovered result: %#v", result) + } + if messagePosts.Load() != 1 || runQueries.Load() != 1 || eventQueries.Load() != 1 { + t.Fatalf("request counts: message=%d run=%d events=%d", messagePosts.Load(), runQueries.Load(), eventQueries.Load()) + } +} + +func TestHTTPClientDerivesRecoveryURLFromMetadataRunID(t *testing.T) { + var messagePosts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch { + case request.Method == http.MethodPost && request.URL.Path == "/api/open/agent-sessions/session-1/messages/stream": + messagePosts.Add(1) + writer.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(writer, "event: metadata\nid: meta-1\ndata: {\"run_id\":\"run-derived\"}\n\n") + case request.Method == http.MethodGet && request.URL.Path == "/api/open/agent-sessions/session-1/runs/run-derived": + writer.Header().Set("Content-Type", "application/json") + fmt.Fprint(writer, `{"status":"running"}`) + case request.Method == http.MethodGet && request.URL.Path == "/api/open/agent-sessions/session-1/runs/run-derived/events": + if request.Header.Get("Last-Event-ID") != "meta-1" { + t.Errorf("Last-Event-ID = %q, want meta-1", request.Header.Get("Last-Event-ID")) + } + writer.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(writer, successfulSSE("run-derived", "derived recovery")) + default: + t.Errorf("unexpected request: %s %s", request.Method, request.URL.Path) + writer.WriteHeader(http.StatusNotFound) + } + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 1, time.Millisecond) + result, err := client.StreamMessage(context.Background(), validMessageRequest(), nil) + if err != nil { + t.Fatalf("StreamMessage() error = %v", err) + } + if result.Answer != "derived recovery" || result.RunID != "run-derived" || messagePosts.Load() != 1 { + t.Fatalf("unexpected recovery result: %#v, posts=%d", result, messagePosts.Load()) + } +} + +func TestHTTPClientStopsRecoveryOnFailedRunStatus(t *testing.T) { + var eventQueries atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch { + case request.Method == http.MethodPost: + writer.Header().Set("Content-Type", "text/event-stream") + writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1") + fmt.Fprint(writer, "event: trace\ndata: {\"event\":\"message.delta\",\"text\":\"partial\"}\n\n") + case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/run-1"): + writer.Header().Set("Content-Type", "application/json") + fmt.Fprint(writer, `{"status":"timeout"}`) + case request.Method == http.MethodGet: + eventQueries.Add(1) + writer.WriteHeader(http.StatusInternalServerError) + } + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 2, time.Millisecond) + _, err := client.StreamMessage(context.Background(), validMessageRequest(), nil) + if !errors.Is(err, ErrRunFailed) { + t.Fatalf("StreamMessage() error = %v, want ErrRunFailed", err) + } + if eventQueries.Load() != 0 { + t.Fatalf("events endpoint queried %d times after failed run", eventQueries.Load()) + } +} + +func TestHTTPClientReturnsRecoveryExhaustedAfterBoundedAttempts(t *testing.T) { + var messagePosts atomic.Int32 + var runQueries atomic.Int32 + var eventQueries atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch { + case request.Method == http.MethodPost: + messagePosts.Add(1) + writer.Header().Set("Content-Type", "text/event-stream") + writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1") + fmt.Fprint(writer, "event: trace\nid: initial-1\ndata: {\"event\":\"message.delta\",\"text\":\"partial\"}\n\n") + case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/run-1"): + runQueries.Add(1) + writer.Header().Set("Content-Type", "application/json") + fmt.Fprint(writer, `{"status":"running"}`) + case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/events"): + eventQueries.Add(1) + writer.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(writer, ": still running\n\n") + } + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 2, time.Millisecond) + result, err := client.StreamMessage(context.Background(), validMessageRequest(), nil) + if !errors.Is(err, ErrRecoveryExhausted) { + t.Fatalf("StreamMessage() error = %v, want ErrRecoveryExhausted", err) + } + if result.Answer != "" || messagePosts.Load() != 1 || runQueries.Load() != 2 || eventQueries.Load() != 2 { + t.Fatalf("unbounded or partial recovery: result=%#v messages=%d runs=%d events=%d", result, messagePosts.Load(), runQueries.Load(), eventQueries.Load()) + } +} + +func TestHTTPClientRejectsCrossOriginContentLocation(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "text/event-stream") + writer.Header().Set("Content-Location", "https://attacker.example/api/open/agent-sessions/session-1/runs/run-1") + fmt.Fprint(writer, successfulSSE("run-1", "must not be accepted")) + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 1, time.Millisecond) + _, err := client.StreamMessage(context.Background(), validMessageRequest(), nil) + if !errors.Is(err, ErrProtocol) || !strings.Contains(err.Error(), "cross-origin") { + t.Fatalf("StreamMessage() error = %v, want cross-origin protocol error", err) + } +} + +func TestHTTPClientNeverReturnsPartialAnswerWithoutRecoveryLocation(t *testing.T) { + var messagePosts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + messagePosts.Add(1) + writer.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(writer, "event: trace\ndata: {\"event\":\"message.delta\",\"text\":\"partial\"}\n\n") + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 2, time.Millisecond) + result, err := client.StreamMessage(context.Background(), validMessageRequest(), nil) + if !errors.Is(err, ErrStreamIncomplete) { + t.Fatalf("StreamMessage() error = %v, want ErrStreamIncomplete", err) + } + if result.Answer != "" || messagePosts.Load() != 1 { + t.Fatalf("partial response escaped or message was retried: result=%#v posts=%d", result, messagePosts.Load()) + } +} + +func TestHTTPClientRedactsHTTPErrorBodyAndKey(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "application/json") + writer.WriteHeader(http.StatusTooManyRequests) + fmt.Fprint(writer, `{"code":"rate_limited","message":"provider-private-value"}`) + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 0, time.Millisecond) + _, err := client.CreateSession(context.Background(), CreateSessionRequest{ + ExternalSubjectID: "subject-1", + IdempotencyKey: "session-key", + }) + if !errors.Is(err, ErrHTTPStatus) { + t.Fatalf("CreateSession() error = %v, want ErrHTTPStatus", err) + } + if !strings.Contains(err.Error(), "rate_limited") || strings.Contains(err.Error(), "provider-private-value") || strings.Contains(err.Error(), testAPIKey) { + t.Fatalf("unsafe or incomplete HTTP error: %v", err) + } +} + +func TestHTTPClientLimitsControlResponseBody(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "application/json") + fmt.Fprint(writer, `{"session_id":"`) + fmt.Fprint(writer, strings.Repeat("x", int(maxControlResponseBytes))) + fmt.Fprint(writer, `"}`) + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 0, time.Millisecond) + _, err := client.CreateSession(context.Background(), CreateSessionRequest{ + ExternalSubjectID: "subject-1", + IdempotencyKey: "session-key", + }) + if !errors.Is(err, ErrProtocol) || !strings.Contains(err.Error(), "size limit") { + t.Fatalf("CreateSession() error = %v, want bounded protocol error", err) + } +} + +func TestHTTPClientDisabledDoesNotUseNetwork(t *testing.T) { + client, err := NewHTTPClient(Config{Enabled: false}) + if err != nil { + t.Fatalf("NewHTTPClient() error = %v", err) + } + if _, err := client.CreateSession(context.Background(), CreateSessionRequest{}); !errors.Is(err, ErrDisabled) { + t.Fatalf("CreateSession() error = %v, want ErrDisabled", err) + } + if _, err := client.StreamMessage(context.Background(), StreamMessageRequest{}, nil); !errors.Is(err, ErrDisabled) { + t.Fatalf("StreamMessage() error = %v, want ErrDisabled", err) + } +} + +func TestNewHTTPClientRejectsUnsafeAPIKey(t *testing.T) { + _, err := NewHTTPClient(Config{ + Enabled: true, + BaseURL: "https://superagent.example.test", + APIKey: "unsafe key", + }) + if !errors.Is(err, ErrInvalidConfig) || strings.Contains(err.Error(), "unsafe key") { + t.Fatalf("NewHTTPClient() error = %v, want redacted config error", err) + } +} + +func TestHTTPClientHonorsContextDuringRecoveryBackoff(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + if request.Method != http.MethodPost { + t.Errorf("unexpected recovery request after timeout: %s", request.Method) + } + writer.Header().Set("Content-Type", "text/event-stream") + writer.Header().Set("Content-Location", "/api/open/agent-sessions/session-1/runs/run-1") + fmt.Fprint(writer, "event: trace\nid: event-1\ndata: {\"event\":\"message.delta\",\"text\":\"partial\"}\n\n") + })) + defer server.Close() + + client := newTestHTTPClient(t, server.URL, 2, 100*time.Millisecond) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) + defer cancel() + _, err := client.StreamMessage(ctx, validMessageRequest(), nil) + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("StreamMessage() error = %v, want context deadline", err) + } +} + +func TestHTTPClientValidatesMessageLimitAndEventContentType(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { + writer.Header().Set("Content-Type", "application/json") + fmt.Fprint(writer, `{}`) + })) + defer server.Close() + + client, err := NewHTTPClient(Config{ + Enabled: true, + BaseURL: server.URL, + APIKey: testAPIKey, + ConnectTimeout: time.Second, + RecoveryInitialBackoff: time.Millisecond, + MaxMessageBytes: 4, + }) + if err != nil { + t.Fatalf("NewHTTPClient() error = %v", err) + } + tooLarge := validMessageRequest() + tooLarge.Message = "12345" + if _, err := client.StreamMessage(context.Background(), tooLarge, nil); !errors.Is(err, ErrInvalidRequest) { + t.Fatalf("large message error = %v, want ErrInvalidRequest", err) + } + valid := validMessageRequest() + valid.Message = "1234" + if _, err := client.StreamMessage(context.Background(), valid, nil); !errors.Is(err, ErrProtocol) { + t.Fatalf("wrong content-type error = %v, want ErrProtocol", err) + } +} + +func newTestHTTPClient(t *testing.T, baseURL string, recoveryAttempts int, backoff time.Duration) *HTTPClient { + t.Helper() + client, err := NewHTTPClient(Config{ + Enabled: true, + BaseURL: baseURL, + APIKey: testAPIKey, + ConnectTimeout: time.Second, + RecoveryMaxAttempts: recoveryAttempts, + RecoveryInitialBackoff: backoff, + MaxMessageBytes: 1024, + }) + if err != nil { + t.Fatalf("NewHTTPClient() error = %v", err) + } + return client +} + +func requireProviderHeaders(t *testing.T, request *http.Request, requestID string) string { + t.Helper() + if request.Header.Get("Authorization") != "Bearer "+testAPIKey { + t.Errorf("unexpected Authorization header") + } + if request.Header.Get("Cache-Control") != "no-cache" { + t.Errorf("Cache-Control = %q", request.Header.Get("Cache-Control")) + } + if request.Header.Get("X-Request-ID") != requestID { + t.Errorf("X-Request-ID = %q, want %q", request.Header.Get("X-Request-ID"), requestID) + } + csrf := request.Header.Get("X-CSRF-Token") + cookie, err := request.Cookie("csrf_token") + if err != nil || csrf == "" || cookie.Value != csrf { + t.Errorf("invalid CSRF double-submit values: header-present=%t cookie-error=%v", csrf != "", err) + } + return csrf +} + +func validMessageRequest() StreamMessageRequest { + return StreamMessageRequest{ + SessionID: "session-1", + Message: "safe", + IdempotencyKey: "message-key", + RequestID: "request-1", + } +} + +func successfulSSE(runID, answer string) string { + return "event: trace\nid: event-1\ndata: {\"event\":\"message.final\",\"run_id\":\"" + runID + "\",\"text\":\"" + answer + "\"}\n\n" + + "event: trace\nid: event-2\ndata: {\"event\":\"run.completed\",\"run_id\":\"" + runID + "\",\"status\":\"success\"}\n\n" + + "event: end\nid: event-3\n\n" +} diff --git a/internal/integration/superagent/sse.go b/internal/integration/superagent/sse.go new file mode 100644 index 0000000..a4f6f8a --- /dev/null +++ b/internal/integration/superagent/sse.go @@ -0,0 +1,512 @@ +package superagent + +import ( + "bufio" + "encoding/json" + "fmt" + "io" + "regexp" + "strconv" + "strings" +) + +const ( + maxSSELineBytes = 1024 * 1024 + maxSSEEventBytes = 1024 * 1024 + maxSSEStreamBytes = 32 * 1024 * 1024 + maxAnswerBytes = 4 * 1024 * 1024 + maxSSEEvents = 10_000 + maxTraceEvents = 1_000 + maxTraceText = 2_048 +) + +var ( + safeProviderValue = regexp.MustCompile(`^[A-Za-z0-9._-]{1,128}$`) + traceJSONSecret = regexp.MustCompile(`(?i)("(?:api[_-]?key|token|secret|password|authorization|cookie|csrf[_-]?token)"\s*:\s*")[^"]*(")`) + traceHeaderSecret = regexp.MustCompile(`(?i)((?:authorization|cookie|x-csrf-token)\s*:\s*)[^\s,;]+`) + traceSecretValue = regexp.MustCompile(`(?i)((?:api[_-]?key|token|secret|password|authorization|cookie|csrf[_-]?token)\s*[=:]\s*)[^\s,;}]+`) +) + +type streamState struct { + sessionID string + + endSeen bool + runCompleted bool + runID string + runURL string + lastEventID string + + profileID string + profileVersionID string + modelName string + finalContent string + deltaContent strings.Builder + fallbackContent string + usage TokenUsage + fallbackUsage TokenUsage + fallbackModel string + + failureCode string + failureStatus string + processedIDs map[string]struct{} + eventTypeSet map[string]struct{} + eventTypes []string + traceEvents []TraceEvent + eventCount int + streamBytes int64 +} + +func newStreamState(sessionID string) *streamState { + return &streamState{ + sessionID: sessionID, + processedIDs: make(map[string]struct{}), + eventTypeSet: make(map[string]struct{}), + } +} + +type sseFrame struct { + eventType string + id string + data []string + dataBytes int +} + +func (s *streamState) consume(reader io.Reader, traceHandler TraceHandler) error { + if reader == nil { + return fmt.Errorf("%w: empty SSE body", ErrProtocol) + } + + scanner := bufio.NewScanner(reader) + scanner.Buffer(make([]byte, 64*1024), maxSSELineBytes) + frame := sseFrame{eventType: "message"} + flush := func() error { + if len(frame.data) == 0 && frame.eventType != "end" { + frame = sseFrame{eventType: "message"} + return nil + } + err := s.consumeFrame(frame, traceHandler) + frame = sseFrame{eventType: "message"} + return err + } + + for scanner.Scan() { + line := scanner.Text() + s.streamBytes += int64(len(line) + 1) + if s.streamBytes > maxSSEStreamBytes { + return fmt.Errorf("%w: SSE stream exceeds size limit", ErrProtocol) + } + if line == "" { + if err := flush(); err != nil { + return err + } + if s.endSeen { + return nil + } + continue + } + if strings.HasPrefix(line, ":") { + continue + } + + field, value, found := strings.Cut(line, ":") + if !found { + field, value = line, "" + } else if strings.HasPrefix(value, " ") { + value = value[1:] + } + switch field { + case "event": + if strings.TrimSpace(value) == "" { + frame.eventType = "message" + } else { + frame.eventType = strings.TrimSpace(value) + } + case "id": + if !strings.ContainsRune(value, '\x00') { + frame.id = strings.TrimSpace(value) + } + case "data": + frame.dataBytes += len(value) + 1 + if frame.dataBytes > maxSSEEventBytes { + return fmt.Errorf("%w: SSE event exceeds size limit", ErrProtocol) + } + frame.data = append(frame.data, value) + } + } + if err := scanner.Err(); err != nil { + return ErrStreamRead + } + if err := flush(); err != nil { + return err + } + if !s.endSeen { + return ErrStreamIncomplete + } + return nil +} + +func (s *streamState) consumeFrame(frame sseFrame, traceHandler TraceHandler) error { + s.eventCount++ + if s.eventCount > maxSSEEvents { + return fmt.Errorf("%w: too many SSE events", ErrProtocol) + } + if !validResourceID(frame.eventType) { + return fmt.Errorf("%w: invalid SSE event type", ErrProtocol) + } + if frame.id != "" && !validOptionalHeader(frame.id) { + return fmt.Errorf("%w: invalid SSE event ID", ErrProtocol) + } + if frame.id != "" { + if _, duplicate := s.processedIDs[frame.id]; duplicate { + return nil + } + s.processedIDs[frame.id] = struct{}{} + s.lastEventID = frame.id + } + s.addEventType(frame.eventType) + + data := strings.Join(frame.data, "\n") + switch frame.eventType { + case "end": + s.endSeen = true + return nil + case "error": + payload, _ := decodeJSONObject(data) + s.failureCode = safeValue(stringValue(payload, "code")) + if s.failureCode == "" { + s.failureCode = "stream_error" + } + return &RunError{Code: s.failureCode} + case "metadata": + payload, err := decodeJSONObject(data) + if err != nil { + return err + } + s.runID = valueOrExisting(safeValue(stringValue(payload, "run_id")), s.runID) + s.profileID = valueOrExisting(safeValue(stringValue(payload, "resolved_profile_id")), s.profileID) + s.profileVersionID = valueOrExisting(safeValue(stringValue(payload, "resolved_profile_version_id")), s.profileVersionID) + case "messages": + payload, err := decodeJSON(data) + if err != nil { + return err + } + return s.consumeMessages(payload) + case "values": + payload, err := decodeJSONObject(data) + if err != nil { + return err + } + return s.consumeMessages(payload["messages"]) + case "trace": + payload, err := decodeJSONObject(data) + if err != nil { + return err + } + return s.consumeTrace(payload, traceHandler) + } + return nil +} + +func (s *streamState) consumeTrace(payload map[string]any, traceHandler TraceHandler) error { + trace := payload + if _, hasEvent := payload["event"]; !hasEvent { + if nested, ok := payload["data"].(map[string]any); ok { + trace = nested + } + } + + s.runID = valueOrExisting(safeValue(stringValue(trace, "run_id")), s.runID) + eventName := stringValue(trace, "event") + var protocolErr error + switch eventName { + case "message.delta": + text := stringValue(trace, "text") + if s.deltaContent.Len()+len(text) > maxAnswerBytes { + return fmt.Errorf("%w: streamed answer exceeds size limit", ErrProtocol) + } + s.deltaContent.WriteString(text) + case "message.final": + if text := stringValue(trace, "text"); strings.TrimSpace(text) != "" { + if len(text) > maxAnswerBytes { + return fmt.Errorf("%w: final answer exceeds size limit", ErrProtocol) + } + s.finalContent = text + } + case "run.completed": + status := safeValue(stringValue(trace, "status")) + if status != "success" { + s.failureStatus = valueOrExisting(status, "non_success") + protocolErr = &RunError{Status: s.failureStatus} + } else { + s.runCompleted = true + } + case "run.failed": + if providerError, ok := trace["error"].(map[string]any); ok { + s.failureCode = safeValue(stringValue(providerError, "code")) + } + if s.failureCode == "" { + s.failureCode = "run_failed" + } + protocolErr = &RunError{Code: s.failureCode} + } + + toolCalls, _ := trace["tool_calls"].([]any) + if len(toolCalls) == 0 { + s.emitTrace(traceEvent(trace, trace), traceHandler) + return protocolErr + } + for _, rawCall := range toolCalls { + detail, ok := rawCall.(map[string]any) + if !ok { + continue + } + s.emitTrace(traceEvent(trace, detail), traceHandler) + } + return protocolErr +} + +func (s *streamState) consumeMessages(value any) error { + switch current := value.(type) { + case []any: + for _, item := range current { + if err := s.consumeMessages(item); err != nil { + return err + } + } + case map[string]any: + if nested, exists := current["messages"]; exists { + return s.consumeMessages(nested) + } + if stringValue(current, "type") != "ai" { + return nil + } + content := contentText(current["content"]) + if strings.TrimSpace(content) == "" { + return nil + } + if len(content) > maxAnswerBytes { + return fmt.Errorf("%w: final answer exceeds size limit", ErrProtocol) + } + responseMetadata, _ := current["response_metadata"].(map[string]any) + usageMetadata, _ := current["usage_metadata"].(map[string]any) + modelName := valueOrExisting(safeTrace(stringValue(responseMetadata, "model_name")), s.modelName) + usage := TokenUsage{ + Input: int64Value(usageMetadata, "input_tokens"), + Output: int64Value(usageMetadata, "output_tokens"), + Total: int64Value(usageMetadata, "total_tokens"), + } + if stringValue(responseMetadata, "finish_reason") == "stop" { + s.finalContent = content + s.modelName = modelName + s.usage = usage + return nil + } + s.fallbackContent = content + s.fallbackModel = modelName + s.fallbackUsage = usage + } + return nil +} + +func (s *streamState) result() (Result, error) { + if s.failureCode != "" || s.failureStatus != "" { + return Result{}, &RunError{Code: s.failureCode, Status: s.failureStatus} + } + if !s.endSeen { + return Result{}, fmt.Errorf("%w: missing end event", ErrProtocol) + } + if !s.runCompleted { + return Result{}, fmt.Errorf("%w: missing successful run.completed event", ErrProtocol) + } + answer := s.finalContent + if strings.TrimSpace(answer) == "" { + answer = s.deltaContent.String() + } + if strings.TrimSpace(answer) == "" { + answer = s.fallbackContent + s.modelName = s.fallbackModel + s.usage = s.fallbackUsage + } + if strings.TrimSpace(answer) == "" { + return Result{}, fmt.Errorf("%w: missing final AI answer", ErrProtocol) + } + + return Result{ + SessionID: s.sessionID, + RunID: s.runID, + ProfileID: s.profileID, + ProfileVersionID: s.profileVersionID, + ModelName: s.modelName, + Answer: answer, + Usage: s.usage, + EventTypes: append([]string(nil), s.eventTypes...), + TraceEvents: append([]TraceEvent(nil), s.traceEvents...), + LastEventID: s.lastEventID, + }, nil +} + +func (s *streamState) addEventType(eventType string) { + if _, exists := s.eventTypeSet[eventType]; exists { + return + } + s.eventTypeSet[eventType] = struct{}{} + s.eventTypes = append(s.eventTypes, eventType) +} + +func (s *streamState) emitTrace(event TraceEvent, traceHandler TraceHandler) { + if event.Event == "" { + return + } + if len(s.traceEvents) < maxTraceEvents { + s.traceEvents = append(s.traceEvents, event) + } + if traceHandler != nil { + traceHandler(event) + } +} + +func traceEvent(trace, detail map[string]any) TraceEvent { + toolCallID := valueOrExisting(stringValue(detail, "tool_call_id"), stringValue(trace, "tool_call_id")) + toolName := valueOrExisting(stringValue(detail, "name"), stringValue(trace, "name")) + return TraceEvent{ + Event: safeTrace(stringValue(trace, "event")), + RunID: safeTrace(stringValue(trace, "run_id")), + MessageID: safeTrace(stringValue(trace, "message_id")), + ToolCallID: safeTrace(toolCallID), + ToolName: safeTrace(toolName), + Text: safeTrace(stringValue(trace, "text")), + Status: safeTrace(stringValue(trace, "status")), + Timestamp: safeTrace(scalarString(trace["ts"])), + } +} + +func decodeJSON(data string) (any, error) { + if strings.TrimSpace(data) == "" { + return nil, fmt.Errorf("%w: empty JSON event", ErrProtocol) + } + decoder := json.NewDecoder(strings.NewReader(data)) + decoder.UseNumber() + var value any + if err := decoder.Decode(&value); err != nil { + return nil, fmt.Errorf("%w: invalid JSON event", ErrProtocol) + } + if decoder.Decode(&struct{}{}) != io.EOF { + return nil, fmt.Errorf("%w: multiple JSON values in one event", ErrProtocol) + } + return value, nil +} + +func decodeJSONObject(data string) (map[string]any, error) { + value, err := decodeJSON(data) + if err != nil { + return nil, err + } + object, ok := value.(map[string]any) + if !ok { + return nil, fmt.Errorf("%w: JSON event must be an object", ErrProtocol) + } + return object, nil +} + +func contentText(value any) string { + switch current := value.(type) { + case string: + return current + case []any: + parts := make([]string, 0, len(current)) + for _, item := range current { + if text := contentText(item); strings.TrimSpace(text) != "" { + parts = append(parts, text) + } + } + return strings.Join(parts, "\n") + case map[string]any: + if text := contentText(current["text"]); strings.TrimSpace(text) != "" { + return text + } + return contentText(current["content"]) + default: + return "" + } +} + +func stringValue(object map[string]any, key string) string { + if object == nil { + return "" + } + value, ok := object[key].(string) + if !ok { + return "" + } + return value +} + +func int64Value(object map[string]any, key string) int64 { + if object == nil { + return 0 + } + switch value := object[key].(type) { + case json.Number: + parsed, _ := value.Int64() + return nonNegative(parsed) + case float64: + return nonNegative(int64(value)) + case int64: + return nonNegative(value) + case string: + parsed, _ := strconv.ParseInt(value, 10, 64) + return nonNegative(parsed) + default: + return 0 + } +} + +func nonNegative(value int64) int64 { + if value < 0 { + return 0 + } + return value +} + +func scalarString(value any) string { + switch current := value.(type) { + case string: + return current + case json.Number: + return current.String() + case float64: + return strconv.FormatFloat(current, 'f', -1, 64) + default: + return "" + } +} + +func valueOrExisting(value, existing string) string { + if value == "" { + return existing + } + return value +} + +func safeValue(value string) string { + value = strings.TrimSpace(value) + if !safeProviderValue.MatchString(value) { + return "" + } + return value +} + +func safeTrace(value string) string { + value = strings.TrimSpace(strings.ToValidUTF8(value, "�")) + if value == "" { + return "" + } + value = traceJSONSecret.ReplaceAllString(value, "${1}***${2}") + value = traceHeaderSecret.ReplaceAllString(value, "${1}***") + value = traceSecretValue.ReplaceAllString(value, "${1}***") + runes := []rune(value) + if len(runes) > maxTraceText { + return string(runes[:maxTraceText]) + } + return value +} diff --git a/internal/integration/superagent/sse_test.go b/internal/integration/superagent/sse_test.go new file mode 100644 index 0000000..42fb250 --- /dev/null +++ b/internal/integration/superagent/sse_test.go @@ -0,0 +1,169 @@ +package superagent + +import ( + "errors" + "strings" + "testing" +) + +func TestStreamStateConsumesCurrentTraceProtocol(t *testing.T) { + state := newStreamState("session-1") + var traces []TraceEvent + stream := strings.Join([]string{ + "event: metadata\nid: 1\ndata: {\"run_id\":\"run-1\",\"resolved_profile_id\":\"profile-1\",\"resolved_profile_version_id\":\"version-1\"}\n", + "event: trace\nid: 2\ndata: {\"event\":\"message.delta\",\"run_id\":\"run-1\",\"text\":\"Hel\"}\n", + "event: trace\nid: 3\ndata: {\"event\":\"message.delta\",\"run_id\":\"run-1\",\"text\":\"lo\"}\n", + "event: trace\nid: 4\ndata: {\"event\":\"message.final\",\"run_id\":\"run-1\",\"message_id\":\"message-1\",\"text\":\"Hello\"}\n", + "event: trace\nid: 5\ndata: {\"event\":\"run.completed\",\"run_id\":\"run-1\",\"status\":\"success\"}\n", + "event: end\nid: 6\n", + }, "\n") + + if err := state.consume(strings.NewReader(stream), func(event TraceEvent) { + traces = append(traces, event) + }); err != nil { + t.Fatalf("consume() error = %v", err) + } + result, err := state.result() + if err != nil { + t.Fatalf("result() error = %v", err) + } + if result.Answer != "Hello" || result.RunID != "run-1" || result.ProfileID != "profile-1" || result.ProfileVersionID != "version-1" { + t.Fatalf("unexpected result: %#v", result) + } + if result.LastEventID != "6" { + t.Fatalf("LastEventID = %q, want 6", result.LastEventID) + } + if len(traces) != 4 || traces[2].MessageID != "message-1" { + t.Fatalf("unexpected traces: %#v", traces) + } +} + +func TestStreamStateConsumesLegacyValuesAndStructuredContent(t *testing.T) { + state := newStreamState("session-legacy") + stream := "event: values\n" + + "data: {\"messages\":[{\"type\":\"human\",\"content\":\"ignored\"},{\"type\":\"ai\",\"content\":[{\"text\":\"first\"},{\"content\":\"second\"}],\"response_metadata\":{\"finish_reason\":\"stop\",\"model_name\":\"model-1\"},\"usage_metadata\":{\"input_tokens\":2,\"output_tokens\":3,\"total_tokens\":5}}]}\n\n" + + "event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" + + "event: end\n\n" + + if err := state.consume(strings.NewReader(stream), nil); err != nil { + t.Fatalf("consume() error = %v", err) + } + result, err := state.result() + if err != nil { + t.Fatalf("result() error = %v", err) + } + if result.Answer != "first\nsecond" || result.ModelName != "model-1" { + t.Fatalf("unexpected legacy result: %#v", result) + } + if result.Usage != (TokenUsage{Input: 2, Output: 3, Total: 5}) { + t.Fatalf("Usage = %#v", result.Usage) + } +} + +func TestStreamStateSupportsMultilineDataAndDeduplicatesIDs(t *testing.T) { + state := newStreamState("session-1") + stream := ": heartbeat\n\n" + + "event: trace\nid: delta-1\ndata: {\"event\":\"message.delta\",\ndata: \"text\":\"A\"}\n\n" + + "event: trace\nid: delta-1\ndata: {\"event\":\"message.delta\",\"text\":\"duplicate\"}\n\n" + + "event: trace\nid: done-1\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" + + "event: end\nid: end-1\n\n" + + if err := state.consume(strings.NewReader(stream), nil); err != nil { + t.Fatalf("consume() error = %v", err) + } + result, err := state.result() + if err != nil { + t.Fatalf("result() error = %v", err) + } + if result.Answer != "A" { + t.Fatalf("Answer = %q, want A", result.Answer) + } +} + +func TestStreamStateRejectsIncompleteOrFailedStreams(t *testing.T) { + tests := []struct { + name string + stream string + wantErr error + result bool + }{ + { + name: "missing end", + stream: "event: trace\ndata: {\"event\":\"message.final\",\"text\":\"partial\"}\n\n" + "event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n", + wantErr: ErrStreamIncomplete, + }, + { + name: "missing completed", + stream: "event: trace\ndata: {\"event\":\"message.final\",\"text\":\"partial\"}\n\n" + "event: end\n\n", + wantErr: ErrProtocol, + result: true, + }, + { + name: "missing answer", + stream: "event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" + "event: end\n\n", + wantErr: ErrProtocol, + result: true, + }, + { + name: "top-level error", + stream: "event: error\ndata: {\"code\":\"provider_error\",\"message\":\"do not expose me\"}\n\n", + wantErr: ErrRunFailed, + }, + { + name: "run failed", + stream: "event: trace\ndata: {\"event\":\"run.failed\",\"error\":{\"code\":\"tool_failed\",\"message\":\"private\"}}\n\n", + wantErr: ErrRunFailed, + }, + { + name: "completed non-success", + stream: "event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"error\"}\n\n", + wantErr: ErrRunFailed, + }, + { + name: "invalid JSON", + stream: "event: trace\ndata: not-json\n\n", + wantErr: ErrProtocol, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + state := newStreamState("session-1") + err := state.consume(strings.NewReader(tt.stream), nil) + if tt.result && err == nil { + _, err = state.result() + } + if !errors.Is(err, tt.wantErr) { + t.Fatalf("error = %v, want errors.Is(..., %v)", err, tt.wantErr) + } + if strings.Contains(strings.ToLower(err.Error()), "do not expose") || strings.Contains(strings.ToLower(err.Error()), "private") { + t.Fatalf("error leaked provider message: %v", err) + } + }) + } +} + +func TestTraceProjectionRedactsAndBoundsText(t *testing.T) { + state := newStreamState("session-1") + secretText := "token=highly-secret " + strings.Repeat("x", maxTraceText+100) + stream := "event: trace\ndata: {\"event\":\"message.final\",\"text\":" + quotedJSON(secretText) + "}\n\n" + + "event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" + + "event: end\n\n" + + if err := state.consume(strings.NewReader(stream), nil); err != nil { + t.Fatalf("consume() error = %v", err) + } + if len(state.traceEvents) == 0 { + t.Fatal("expected trace event") + } + text := state.traceEvents[0].Text + if strings.Contains(text, "highly-secret") || len(text) > maxTraceText { + t.Fatalf("trace text was not safely projected: length=%d text=%q", len(text), text[:min(len(text), 80)]) + } +} + +func quotedJSON(value string) string { + value = strings.ReplaceAll(value, "\\", "\\\\") + value = strings.ReplaceAll(value, "\"", "\\\"") + return "\"" + value + "\"" +} diff --git a/internal/integration/superagent/types.go b/internal/integration/superagent/types.go new file mode 100644 index 0000000..6174480 --- /dev/null +++ b/internal/integration/superagent/types.go @@ -0,0 +1,143 @@ +// Package superagent implements the outbound SuperAgent Open API boundary. +package superagent + +import ( + "context" + "errors" + "fmt" + "time" +) + +var ( + // ErrDisabled indicates that outbound SuperAgent access is intentionally disabled. + ErrDisabled = errors.New("superagent Open API is disabled") + // ErrInvalidConfig indicates invalid client configuration. + ErrInvalidConfig = errors.New("invalid superagent client configuration") + // ErrInvalidRequest indicates an invalid call made by an application service. + ErrInvalidRequest = errors.New("invalid superagent request") + // ErrHTTPStatus indicates a non-success response from the provider. + ErrHTTPStatus = errors.New("superagent returned a non-success HTTP status") + // ErrProtocol indicates a malformed or incomplete provider protocol response. + ErrProtocol = errors.New("invalid superagent protocol response") + // ErrRunFailed indicates that the provider explicitly reported a failed run. + ErrRunFailed = errors.New("superagent run failed") + // ErrStreamIncomplete indicates that an SSE stream ended before its end event. + ErrStreamIncomplete = errors.New("superagent stream ended before the end event") + // ErrStreamRead indicates a low-level SSE read failure. + ErrStreamRead = errors.New("superagent stream read failed") + // ErrRecoveryExhausted indicates that bounded run recovery did not complete. + ErrRecoveryExhausted = errors.New("superagent stream recovery exhausted") +) + +// Config controls the outbound HTTP client. The caller should keep APIKey in a +// secret store and must never log this struct. +type Config struct { + Enabled bool + BaseURL string + APIKey string + ConnectTimeout time.Duration + RecoveryMaxAttempts int + RecoveryInitialBackoff time.Duration + MaxMessageBytes int64 +} + +// Client is the provider-neutral capability exposed to application services. +// Sessions are created separately so a later chat service can reuse one session +// across multiple turns. +type Client interface { + CreateSession(context.Context, CreateSessionRequest) (Session, error) + StreamMessage(context.Context, StreamMessageRequest, TraceHandler) (Result, error) +} + +// CreateSessionRequest starts an externally correlated SuperAgent session. +type CreateSessionRequest struct { + ExternalSubjectID string + IdempotencyKey string + RequestID string + Metadata map[string]any +} + +// Session is the stable provider session handle returned to the application. +type Session struct { + ID string `json:"id"` + Status string `json:"status,omitempty"` +} + +// StreamMessageRequest sends one message to an existing session. +type StreamMessageRequest struct { + SessionID string + Message string + IdempotencyKey string + RequestID string + Metadata map[string]any +} + +// TraceHandler receives only the public, bounded trace projection. +type TraceHandler func(TraceEvent) + +// TraceEvent intentionally excludes raw tool inputs, raw tool outputs, headers, +// cookies, and provider payloads. +type TraceEvent struct { + Event string `json:"event"` + RunID string `json:"run_id,omitempty"` + MessageID string `json:"message_id,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` + ToolName string `json:"tool_name,omitempty"` + Text string `json:"text,omitempty"` + Status string `json:"status,omitempty"` + Timestamp string `json:"timestamp,omitempty"` +} + +// TokenUsage contains provider-reported model token counts. +type TokenUsage struct { + Input int64 `json:"input"` + Output int64 `json:"output"` + Total int64 `json:"total"` +} + +// Result is returned only after all strict stream success conditions hold. +type Result struct { + SessionID string `json:"session_id"` + RunID string `json:"run_id,omitempty"` + ProfileID string `json:"profile_id,omitempty"` + ProfileVersionID string `json:"profile_version_id,omitempty"` + ModelName string `json:"model_name,omitempty"` + Answer string `json:"answer"` + Usage TokenUsage `json:"usage"` + EventTypes []string `json:"event_types"` + TraceEvents []TraceEvent `json:"trace_events,omitempty"` + LastEventID string `json:"last_event_id,omitempty"` +} + +// HTTPStatusError exposes only bounded status metadata, never the response body. +type HTTPStatusError struct { + StatusCode int + ProviderCode string +} + +func (e *HTTPStatusError) Error() string { + if e.ProviderCode == "" { + return fmt.Sprintf("%s: status=%d", ErrHTTPStatus, e.StatusCode) + } + return fmt.Sprintf("%s: status=%d code=%s", ErrHTTPStatus, e.StatusCode, e.ProviderCode) +} + +func (e *HTTPStatusError) Unwrap() error { return ErrHTTPStatus } + +// RunError exposes a provider error code or terminal status, but not its raw message. +type RunError struct { + Code string + Status string +} + +func (e *RunError) Error() string { + if e.Code != "" { + return fmt.Sprintf("%s: code=%s", ErrRunFailed, e.Code) + } + if e.Status != "" { + return fmt.Sprintf("%s: status=%s", ErrRunFailed, e.Status) + } + return ErrRunFailed.Error() +} + +func (e *RunError) Unwrap() error { return ErrRunFailed } diff --git a/internal/migration/postgis_srid.go b/internal/migration/postgis_srid.go new file mode 100644 index 0000000..d496a62 --- /dev/null +++ b/internal/migration/postgis_srid.go @@ -0,0 +1,403 @@ +// Package migration contains explicit, operator-invoked database migrations. +package migration + +import ( + "context" + "fmt" + "slices" + "sort" + "strconv" + "strings" + + "github.com/jackc/pgx/v5" +) + +const ( + // SourceSRID is the unknown SRID marker currently stored by the imported data. + SourceSRID = 0 + // TargetSRID is the CRS confirmed by the data provider for the eight source tables. + TargetSRID = 4326 +) + +type tableSpec struct { + name string + geometryTypes []string +} + +var sourceTables = []tableSpec{ + {name: "st_2_mpslfh_t_slfh_syd", geometryTypes: []string{"POINT"}}, + {name: "st_2_xianyouxushuichiguan", geometryTypes: []string{"POINT"}}, + {name: "st_2_xianyoufanghuotongdao", geometryTypes: []string{"LINESTRING", "MULTILINESTRING"}}, + {name: "st_2_fanghuojianchazhan", geometryTypes: []string{"POINT"}}, + {name: "st_2_fanghuoliaowangshao", geometryTypes: []string{"POINT"}}, + {name: "st_2_fanghuowangge", geometryTypes: []string{"POLYGON", "MULTIPOLYGON"}}, + {name: "st_2_linqugongkuangqiye", geometryTypes: []string{"POLYGON", "MULTIPOLYGON"}}, + {name: "st_2_mudifenqu_mian", geometryTypes: []string{"POLYGON", "MULTIPOLYGON"}}, +} + +// TableState is a non-record-level summary used to prove the migration boundary. +type TableState struct { + Table string `json:"table"` + TableOwner string `json:"table_owner"` + ColumnType string `json:"column_type"` + TypmodConstrained bool `json:"typmod_constrained"` + TypmodSRID int `json:"typmod_srid"` + TypmodDimensions int `json:"typmod_dimensions"` + TypmodGeometryType string `json:"typmod_geometry_type"` + RowCount int64 `json:"row_count"` + NullGeometryCount int64 `json:"null_geometry_count"` + NonNullGeometryCount int64 `json:"non_null_geometry_count"` + SourceSRIDCount int64 `json:"source_srid_count"` + TargetSRIDCount int64 `json:"target_srid_count"` + UnexpectedSRIDCount int64 `json:"unexpected_srid_count"` + EmptyGeometryCount int64 `json:"empty_geometry_count"` + InvalidGeometryCount int64 `json:"invalid_geometry_count"` + OutOfBoundsGeometryCount int64 `json:"wgs84_out_of_bounds_count"` + GeometryTypes []string `json:"geometry_types"` + SRIDs []int `json:"srids"` + CoordinateDimensions []int `json:"coordinate_dimensions"` + ZGeometryCount int64 `json:"z_geometry_count"` + MGeometryCount int64 `json:"m_geometry_count"` + CanUpdate bool `json:"can_update"` + CanAlter bool `json:"can_alter"` + fingerprint string +} + +// Report summarizes whether the database is ready for the SRID metadata migration. +type Report struct { + TargetSRID int `json:"target_srid"` + CurrentRole string `json:"current_role"` + TransactionReadOnly bool `json:"transaction_read_only"` + CandidateGeometryCount int64 `json:"candidate_geometry_count"` + AlreadyTargetGeometryCount int64 `json:"already_target_geometry_count"` + NeedsTypmodRetype []string `json:"needs_typmod_retype"` + Tables []TableState `json:"tables"` +} + +// Result reports the committed changes without exposing business records or credentials. +type Result struct { + TargetSRID int `json:"target_srid"` + UpdatedRows map[string]int64 `json:"updated_rows"` + RetypedTables []string `json:"retyped_tables"` + GeometryPayloadUnchanged bool `json:"geometry_payload_unchanged"` + Before Report `json:"before"` + After Report `json:"after"` +} + +type rowQuerier interface { + QueryRow(context.Context, string, ...any) pgx.Row +} + +// Inspect reads only schema- and aggregate-level facts from the fixed source tables. +func Inspect(ctx context.Context, db rowQuerier) (Report, error) { + var report Report + report.TargetSRID = TargetSRID + if err := db.QueryRow(ctx, `SELECT current_user, current_setting('transaction_read_only')::boolean`).Scan(&report.CurrentRole, &report.TransactionReadOnly); err != nil { + return Report{}, fmt.Errorf("read transaction mode: %w", err) + } + + for _, spec := range sourceTables { + state, err := inspectTable(ctx, db, spec) + if err != nil { + return Report{}, fmt.Errorf("inspect table %s: %w", spec.name, err) + } + report.Tables = append(report.Tables, state) + report.CandidateGeometryCount += state.SourceSRIDCount + report.AlreadyTargetGeometryCount += state.TargetSRIDCount + if state.SourceSRIDCount > 0 && state.TypmodConstrained && state.TypmodSRID != TargetSRID { + report.NeedsTypmodRetype = append(report.NeedsTypmodRetype, state.Table) + } + } + if report.NeedsTypmodRetype == nil { + report.NeedsTypmodRetype = []string{} + } + return report, nil +} + +// ValidateForApply verifies that ST_SetSRID can be applied without interpreting unknown data. +func ValidateForApply(report Report) error { + if report.TransactionReadOnly { + return fmt.Errorf("migration connection is read-only") + } + if len(report.Tables) != len(sourceTables) { + return fmt.Errorf("inspection contains %d tables, want %d", len(report.Tables), len(sourceTables)) + } + for index, state := range report.Tables { + spec := sourceTables[index] + if state.Table != spec.name { + return fmt.Errorf("inspection table %d is %s, want %s", index, state.Table, spec.name) + } + if state.UnexpectedSRIDCount != 0 || state.NonNullGeometryCount != state.SourceSRIDCount+state.TargetSRIDCount { + return fmt.Errorf("table %s contains geometry outside SRID %d or %d", state.Table, SourceSRID, TargetSRID) + } + if state.OutOfBoundsGeometryCount != 0 { + return fmt.Errorf("table %s contains coordinates outside WGS84 bounds", state.Table) + } + for _, geometryType := range state.GeometryTypes { + if !slices.Contains(spec.geometryTypes, geometryType) { + return fmt.Errorf("table %s contains unexpected geometry type %s", state.Table, geometryType) + } + } + if state.SourceSRIDCount > 0 && !state.CanUpdate { + return fmt.Errorf("migration role cannot update table %s", state.Table) + } + if state.SourceSRIDCount > 0 && state.TypmodConstrained && state.TypmodSRID != TargetSRID { + if state.TypmodDimensions != 2 || !slices.Equal(state.CoordinateDimensions, []int{2}) { + return fmt.Errorf("table %s cannot be safely retyped as two-dimensional EPSG:4326 geometry", state.Table) + } + } + if state.SourceSRIDCount > 0 && state.TypmodConstrained && state.TypmodSRID != TargetSRID && !state.CanAlter { + return fmt.Errorf("migration role cannot retype the incompatible geometry column on table %s", state.Table) + } + } + return nil +} + +// ApplySRID4326 atomically marks confirmed source geometries as EPSG:4326. +// It never transforms coordinates, repairs geometry validity, or changes dimensions. +func ApplySRID4326(ctx context.Context, conn *pgx.Conn) (Result, error) { + tx, err := conn.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.Serializable, AccessMode: pgx.ReadWrite}) + if err != nil { + return Result{}, fmt.Errorf("begin SRID migration transaction: %w", err) + } + defer func() { _ = tx.Rollback(context.Background()) }() + + if _, err := tx.Exec(ctx, `SET LOCAL lock_timeout = '5s'`); err != nil { + return Result{}, fmt.Errorf("set migration lock timeout: %w", err) + } + if _, err := tx.Exec(ctx, `SET LOCAL statement_timeout = '120s'`); err != nil { + return Result{}, fmt.Errorf("set migration statement timeout: %w", err) + } + if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended('fire-safety-ymd:srid-4326', 0))`); err != nil { + return Result{}, fmt.Errorf("acquire SRID migration lock: %w", err) + } + + before, err := Inspect(ctx, tx) + if err != nil { + return Result{}, err + } + if err := ValidateForApply(before); err != nil { + return Result{}, err + } + + result := Result{ + TargetSRID: TargetSRID, + UpdatedRows: make(map[string]int64, len(sourceTables)), + Before: before, + } + for _, state := range before.Tables { + if state.SourceSRIDCount == 0 { + result.UpdatedRows[state.Table] = 0 + continue + } + + tableIdentifier := pgx.Identifier{"public", state.Table}.Sanitize() + if state.TypmodConstrained && state.TypmodSRID != TargetSRID { + statement := fmt.Sprintf(`ALTER TABLE %s ALTER COLUMN geom TYPE geometry(Geometry, 4326) USING ST_SetSRID(geom, 4326)`, tableIdentifier) + if _, err := tx.Exec(ctx, statement); err != nil { + return Result{}, fmt.Errorf("retype geometry column on table %s: %w", state.Table, err) + } + result.RetypedTables = append(result.RetypedTables, state.Table) + result.UpdatedRows[state.Table] = state.SourceSRIDCount + continue + } + + statement := fmt.Sprintf(`UPDATE %s SET geom = ST_SetSRID(geom, $1) WHERE geom IS NOT NULL AND ST_SRID(geom) = $2`, tableIdentifier) + commandTag, err := tx.Exec(ctx, statement, TargetSRID, SourceSRID) + if err != nil { + return Result{}, fmt.Errorf("set SRID metadata on table %s: %w", state.Table, err) + } + updated := commandTag.RowsAffected() + if updated != state.SourceSRIDCount { + return Result{}, fmt.Errorf("table %s updated %d rows, expected %d", state.Table, updated, state.SourceSRIDCount) + } + result.UpdatedRows[state.Table] = updated + } + if result.RetypedTables == nil { + result.RetypedTables = []string{} + } + + after, err := Inspect(ctx, tx) + if err != nil { + return Result{}, err + } + if err := validateAfter(before, after); err != nil { + return Result{}, err + } + result.After = after + result.GeometryPayloadUnchanged = true + + if err := tx.Commit(ctx); err != nil { + return Result{}, fmt.Errorf("commit SRID migration: %w", err) + } + return result, nil +} + +func inspectTable(ctx context.Context, db rowQuerier, spec tableSpec) (TableState, error) { + var state TableState + state.Table = spec.name + qualifiedName := "public." + spec.name + + metadataSQL := ` +SELECT + owner.rolname, + format_type(a.atttypid, a.atttypmod), + a.atttypmod >= 0, + CASE WHEN a.atttypmod >= 0 THEN postgis_typmod_srid(a.atttypmod) ELSE -1 END, + CASE WHEN a.atttypmod >= 0 THEN postgis_typmod_dims(a.atttypmod) ELSE -1 END, + CASE WHEN a.atttypmod >= 0 THEN postgis_typmod_type(a.atttypmod) ELSE '' END, + has_table_privilege(current_user, $1, 'UPDATE'), + ( + c.relowner = (SELECT oid FROM pg_roles WHERE rolname = current_user) + OR pg_has_role(current_user, c.relowner, 'USAGE') + OR (SELECT rolsuper FROM pg_roles WHERE rolname = current_user) + ) +FROM pg_attribute AS a +JOIN pg_class AS c ON c.oid = a.attrelid +JOIN pg_namespace AS n ON n.oid = c.relnamespace +JOIN pg_roles AS owner ON owner.oid = c.relowner +WHERE n.nspname = 'public' + AND c.relname = $2 + AND a.attname = 'geom' + AND NOT a.attisdropped` + if err := db.QueryRow(ctx, metadataSQL, qualifiedName, spec.name).Scan( + &state.TableOwner, + &state.ColumnType, + &state.TypmodConstrained, + &state.TypmodSRID, + &state.TypmodDimensions, + &state.TypmodGeometryType, + &state.CanUpdate, + &state.CanAlter, + ); err != nil { + return TableState{}, fmt.Errorf("inspect geom column metadata: %w", err) + } + + tableIdentifier := pgx.Identifier{"public", spec.name}.Sanitize() + dataSQL := fmt.Sprintf(` +SELECT + count(*)::bigint, + count(*) FILTER (WHERE geom IS NULL)::bigint, + count(*) FILTER (WHERE geom IS NOT NULL)::bigint, + count(*) FILTER (WHERE geom IS NOT NULL AND ST_SRID(geom) = $1)::bigint, + count(*) FILTER (WHERE geom IS NOT NULL AND ST_SRID(geom) = $2)::bigint, + count(*) FILTER (WHERE geom IS NOT NULL AND ST_SRID(geom) NOT IN ($1, $2))::bigint, + count(*) FILTER (WHERE geom IS NOT NULL AND ST_IsEmpty(geom))::bigint, + count(*) FILTER (WHERE geom IS NOT NULL AND NOT ST_IsValid(geom))::bigint, + count(*) FILTER ( + WHERE geom IS NOT NULL + AND NOT ST_IsEmpty(geom) + AND ( + ST_XMin(Box3D(geom)) < -180 OR ST_XMax(Box3D(geom)) > 180 + OR ST_YMin(Box3D(geom)) < -90 OR ST_YMax(Box3D(geom)) > 90 + ) + )::bigint, + COALESCE(string_agg(DISTINCT GeometryType(geom), ',' ORDER BY GeometryType(geom)) FILTER (WHERE geom IS NOT NULL), ''), + COALESCE(string_agg(DISTINCT ST_SRID(geom)::text, ',' ORDER BY ST_SRID(geom)::text) FILTER (WHERE geom IS NOT NULL), ''), + COALESCE(string_agg(DISTINCT ST_NDims(geom)::text, ',' ORDER BY ST_NDims(geom)::text) FILTER (WHERE geom IS NOT NULL), ''), + count(*) FILTER (WHERE geom IS NOT NULL AND ST_Zmflag(geom) IN (2, 3))::bigint, + count(*) FILTER (WHERE geom IS NOT NULL AND ST_Zmflag(geom) IN (1, 3))::bigint, + md5(COALESCE( + string_agg( + md5(encode(ST_AsBinary(geom), 'hex')), + '' ORDER BY md5(encode(ST_AsBinary(geom), 'hex')) + ) FILTER (WHERE geom IS NOT NULL), + '' + )) +FROM %s`, tableIdentifier) + + var geometryTypes string + var srids string + var dimensions string + if err := db.QueryRow(ctx, dataSQL, SourceSRID, TargetSRID).Scan( + &state.RowCount, + &state.NullGeometryCount, + &state.NonNullGeometryCount, + &state.SourceSRIDCount, + &state.TargetSRIDCount, + &state.UnexpectedSRIDCount, + &state.EmptyGeometryCount, + &state.InvalidGeometryCount, + &state.OutOfBoundsGeometryCount, + &geometryTypes, + &srids, + &dimensions, + &state.ZGeometryCount, + &state.MGeometryCount, + &state.fingerprint, + ); err != nil { + return TableState{}, fmt.Errorf("inspect aggregate geometry state: %w", err) + } + state.GeometryTypes = splitStrings(geometryTypes) + parsedSRIDs, err := splitInts(srids) + if err != nil { + return TableState{}, fmt.Errorf("parse SRIDs: %w", err) + } + state.SRIDs = parsedSRIDs + parsedDimensions, err := splitInts(dimensions) + if err != nil { + return TableState{}, fmt.Errorf("parse coordinate dimensions: %w", err) + } + state.CoordinateDimensions = parsedDimensions + return state, nil +} + +func validateAfter(before, after Report) error { + if len(before.Tables) != len(after.Tables) { + return fmt.Errorf("post-migration inspection contains %d tables, want %d", len(after.Tables), len(before.Tables)) + } + for index := range before.Tables { + oldState := before.Tables[index] + newState := after.Tables[index] + if oldState.Table != newState.Table { + return fmt.Errorf("post-migration table %d is %s, want %s", index, newState.Table, oldState.Table) + } + if newState.SourceSRIDCount != 0 || newState.UnexpectedSRIDCount != 0 || newState.TargetSRIDCount != newState.NonNullGeometryCount { + return fmt.Errorf("table %s did not converge to SRID %d", newState.Table, TargetSRID) + } + if err := compareGeometryPayload(oldState, newState); err != nil { + return err + } + } + return nil +} + +func compareGeometryPayload(before, after TableState) error { + if before.RowCount != after.RowCount || + before.NullGeometryCount != after.NullGeometryCount || + before.NonNullGeometryCount != after.NonNullGeometryCount || + before.EmptyGeometryCount != after.EmptyGeometryCount || + before.InvalidGeometryCount != after.InvalidGeometryCount || + before.OutOfBoundsGeometryCount != after.OutOfBoundsGeometryCount || + before.ZGeometryCount != after.ZGeometryCount || + before.MGeometryCount != after.MGeometryCount || + before.fingerprint != after.fingerprint || + !slices.Equal(before.GeometryTypes, after.GeometryTypes) || + !slices.Equal(before.CoordinateDimensions, after.CoordinateDimensions) { + return fmt.Errorf("table %s geometry payload changed outside the SRID metadata", before.Table) + } + return nil +} + +func splitStrings(value string) []string { + if strings.TrimSpace(value) == "" { + return []string{} + } + values := strings.Split(value, ",") + sort.Strings(values) + return values +} + +func splitInts(value string) ([]int, error) { + parts := splitStrings(value) + values := make([]int, 0, len(parts)) + for _, part := range parts { + parsed, err := strconv.Atoi(part) + if err != nil { + return nil, err + } + values = append(values, parsed) + } + sort.Ints(values) + return values, nil +} diff --git a/internal/migration/postgis_srid_test.go b/internal/migration/postgis_srid_test.go new file mode 100644 index 0000000..c7a434a --- /dev/null +++ b/internal/migration/postgis_srid_test.go @@ -0,0 +1,150 @@ +package migration + +import ( + "strings" + "testing" +) + +func TestValidateForApplyAcceptsConfirmedSourceState(t *testing.T) { + report := validMigrationReport() + if err := ValidateForApply(report); err != nil { + t.Fatalf("ValidateForApply() error = %v", err) + } +} + +func TestValidateForApplyRejectsUnsafeState(t *testing.T) { + tests := []struct { + name string + mutate func(*Report) + want string + }{ + { + name: "read only connection", + mutate: func(report *Report) { + report.TransactionReadOnly = true + }, + want: "read-only", + }, + { + name: "unexpected srid", + mutate: func(report *Report) { + report.Tables[0].UnexpectedSRIDCount = 1 + }, + want: "outside SRID", + }, + { + name: "out of bounds", + mutate: func(report *Report) { + report.Tables[0].OutOfBoundsGeometryCount = 1 + }, + want: "outside WGS84 bounds", + }, + { + name: "unexpected type", + mutate: func(report *Report) { + report.Tables[0].GeometryTypes = []string{"POLYGON"} + }, + want: "unexpected geometry type", + }, + { + name: "cannot update", + mutate: func(report *Report) { + report.Tables[0].CanUpdate = false + }, + want: "cannot update", + }, + { + name: "constrained source contains z", + mutate: func(report *Report) { + report.Tables[0].CoordinateDimensions = []int{2, 3} + }, + want: "cannot be safely retyped", + }, + { + name: "cannot alter constrained column", + mutate: func(report *Report) { + report.Tables[0].CanAlter = false + }, + want: "cannot retype", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + report := validMigrationReport() + test.mutate(&report) + err := ValidateForApply(report) + if err == nil || !strings.Contains(err.Error(), test.want) { + t.Fatalf("ValidateForApply() error = %v, want substring %q", err, test.want) + } + }) + } +} + +func TestValidateAfterAcceptsSRIDOnlyChange(t *testing.T) { + before := validMigrationReport() + after := migratedReport(before) + if err := validateAfter(before, after); err != nil { + t.Fatalf("validateAfter() error = %v", err) + } +} + +func TestValidateAfterRejectsGeometryPayloadChange(t *testing.T) { + before := validMigrationReport() + after := migratedReport(before) + after.Tables[2].ZGeometryCount-- + err := validateAfter(before, after) + if err == nil || !strings.Contains(err.Error(), "geometry payload changed") { + t.Fatalf("validateAfter() error = %v", err) + } +} + +func validMigrationReport() Report { + report := Report{TargetSRID: TargetSRID} + for index, spec := range sourceTables { + state := TableState{ + Table: spec.name, + ColumnType: "geometry", + TypmodConstrained: true, + TypmodSRID: SourceSRID, + TypmodDimensions: 2, + TypmodGeometryType: "Geometry", + RowCount: 3, + NonNullGeometryCount: 2, + NullGeometryCount: 1, + SourceSRIDCount: 2, + GeometryTypes: []string{spec.geometryTypes[0]}, + SRIDs: []int{SourceSRID}, + CoordinateDimensions: []int{2}, + CanUpdate: true, + CanAlter: true, + fingerprint: "same-payload", + } + if index == 2 { + state.TypmodConstrained = false + state.TypmodSRID = -1 + state.TypmodDimensions = -1 + state.TypmodGeometryType = "" + state.GeometryTypes = []string{"LINESTRING", "MULTILINESTRING"} + state.CoordinateDimensions = []int{2, 3} + state.ZGeometryCount = 1 + } + report.Tables = append(report.Tables, state) + report.CandidateGeometryCount += state.SourceSRIDCount + } + return report +} + +func migratedReport(before Report) Report { + after := before + after.Tables = append([]TableState(nil), before.Tables...) + after.CandidateGeometryCount = 0 + after.AlreadyTargetGeometryCount = 0 + for index := range after.Tables { + after.Tables[index].SourceSRIDCount = 0 + after.Tables[index].TargetSRIDCount = after.Tables[index].NonNullGeometryCount + after.Tables[index].SRIDs = []int{TargetSRID} + after.AlreadyTargetGeometryCount += after.Tables[index].TargetSRIDCount + } + return after +} diff --git a/internal/repository/doc.go b/internal/repository/doc.go new file mode 100644 index 0000000..c01d899 --- /dev/null +++ b/internal/repository/doc.go @@ -0,0 +1,2 @@ +// Package repository contains persistence adapters and repository contracts that are stable across use cases. +package repository diff --git a/internal/repository/postgis.go b/internal/repository/postgis.go new file mode 100644 index 0000000..86b7ff7 --- /dev/null +++ b/internal/repository/postgis.go @@ -0,0 +1,969 @@ +package repository + +import ( + "context" + "errors" + "fmt" + "slices" + "strconv" + "strings" + "time" + + "fire-safety-ymd/internal/domain" + + "github.com/jackc/pgx/v5/pgtype" + "github.com/jackc/pgx/v5/pgxpool" +) + +const mcpSpatialSRID = 4326 + +// PostGISOptions configures the read-only PostgreSQL/PostGIS connection pool. +type PostGISOptions struct { + DSN string + MaxConns int32 + ConnectTimeout time.Duration + QueryTimeout time.Duration +} + +// TableAudit is a safe schema and geometry summary that contains no business rows. +type TableAudit struct { + Table string `json:"table"` + RowCount int64 `json:"row_count"` + NullGeometryCount int64 `json:"null_geometry_count"` + EmptyGeometryCount int64 `json:"empty_geometry_count"` + InvalidCount int64 `json:"invalid_geometry_count"` + OutOfBoundsCount int64 `json:"wgs84_out_of_bounds_count"` + GeometryTypes []string `json:"geometry_types"` + SRIDs []int `json:"srids"` + HasGISTIndex bool `json:"has_gist_geometry_index"` +} + +// ReadinessReport contains non-sensitive PostGIS readiness information. +type ReadinessReport struct { + PostGISVersion string `json:"postgis_version"` + Tables []TableAudit `json:"tables"` + Warnings []string `json:"warnings"` +} + +// PostGIS implements the fixed read-only spatial repository over pgxpool. +type PostGIS struct { + pool *pgxpool.Pool +} + +type safeConnectionError struct { + message string + cause error +} + +func (e *safeConnectionError) Error() string { return e.message } +func (e *safeConnectionError) Unwrap() error { return e.cause } + +// OpenPostGIS opens and verifies a read-only PostgreSQL connection pool. +func OpenPostGIS(ctx context.Context, options PostGISOptions) (*PostGIS, error) { + if strings.TrimSpace(options.DSN) == "" || options.MaxConns <= 0 || options.ConnectTimeout <= 0 || options.QueryTimeout <= 0 { + return nil, errors.New("invalid PostGIS options") + } + poolConfig, err := pgxpool.ParseConfig(options.DSN) + if err != nil { + return nil, &safeConnectionError{message: "parse PostGIS configuration", cause: err} + } + poolConfig.MaxConns = options.MaxConns + poolConfig.ConnConfig.ConnectTimeout = options.ConnectTimeout + poolConfig.ConnConfig.RuntimeParams["application_name"] = "fire-safety-ymd-mcp" + poolConfig.ConnConfig.RuntimeParams["default_transaction_read_only"] = "on" + queryTimeoutMilliseconds := options.QueryTimeout.Milliseconds() + if queryTimeoutMilliseconds < 1 { + queryTimeoutMilliseconds = 1 + } + poolConfig.ConnConfig.RuntimeParams["statement_timeout"] = strconv.FormatInt(queryTimeoutMilliseconds, 10) + poolConfig.ConnConfig.RuntimeParams["lock_timeout"] = strconv.FormatInt(queryTimeoutMilliseconds, 10) + + pool, err := pgxpool.NewWithConfig(ctx, poolConfig) + if err != nil { + return nil, &safeConnectionError{message: "create PostGIS connection pool", cause: err} + } + connectCtx, cancel := context.WithTimeout(ctx, options.ConnectTimeout) + defer cancel() + if err := pool.Ping(connectCtx); err != nil { + pool.Close() + return nil, &safeConnectionError{message: "connect to PostGIS", cause: err} + } + return &PostGIS{pool: pool}, nil +} + +// Close releases all PostgreSQL connections. +func (p *PostGIS) Close() { + if p != nil && p.pool != nil { + p.pool.Close() + } +} + +// Audit inspects PostGIS and the allowlisted source tables without returning record values. +func (p *PostGIS) Audit(ctx context.Context) (ReadinessReport, error) { + var report ReadinessReport + if err := p.pool.QueryRow(ctx, `SELECT postgis_version()`).Scan(&report.PostGISVersion); err != nil { + return ReadinessReport{}, fmt.Errorf("query PostGIS version: %w", err) + } + + for _, spec := range spatialTableSpecs { + audit, err := p.auditTable(ctx, spec.name) + if err != nil { + return ReadinessReport{}, fmt.Errorf("audit allowlisted table %s: %w", spec.name, err) + } + report.Tables = append(report.Tables, audit) + report.Warnings = append(report.Warnings, tableAuditWarnings(audit)...) + } + if report.Warnings == nil { + report.Warnings = []string{} + } + return report, nil +} + +// ValidateForMCP rejects unsafe spatial metadata while allowing queries to exclude invalid records. +func (p *PostGIS) ValidateForMCP(ctx context.Context, expectedSRID int) (ReadinessReport, error) { + report, err := p.Audit(ctx) + if err != nil { + return ReadinessReport{}, err + } + if err := validateReadiness(report, expectedSRID); err != nil { + return report, err + } + return report, nil +} + +func validateReadiness(report ReadinessReport, expectedSRID int) error { + if expectedSRID != mcpSpatialSRID { + return fmt.Errorf("unsupported MCP spatial SRID: %d", expectedSRID) + } + if len(report.Tables) != len(spatialTableSpecs) { + return fmt.Errorf("readiness report contains %d tables, want %d", len(report.Tables), len(spatialTableSpecs)) + } + for index, audit := range report.Tables { + spec := spatialTableSpecs[index] + if audit.Table != spec.name { + return fmt.Errorf("readiness report table %d is %s, want %s", index, audit.Table, spec.name) + } + if audit.OutOfBoundsCount > 0 { + return fmt.Errorf("table %s contains coordinates outside WGS84 bounds", audit.Table) + } + for _, srid := range audit.SRIDs { + if srid != expectedSRID { + return fmt.Errorf("table %s contains unexpected SRID %d", audit.Table, srid) + } + } + for _, geometryType := range audit.GeometryTypes { + if !slices.Contains(spec.geometryTypes, geometryType) { + return fmt.Errorf("table %s contains unexpected geometry type %s", audit.Table, geometryType) + } + } + } + return nil +} + +func tableAuditWarnings(audit TableAudit) []string { + warnings := make([]string, 0, 5) + if audit.RowCount == 0 { + warnings = append(warnings, audit.Table+":empty_table") + } + if audit.NullGeometryCount > 0 { + warnings = append(warnings, audit.Table+":null_geometries_excluded") + } + if audit.EmptyGeometryCount > 0 { + warnings = append(warnings, audit.Table+":empty_geometries_excluded") + } + if audit.InvalidCount > 0 { + warnings = append(warnings, audit.Table+":invalid_geometries_excluded") + } + if !audit.HasGISTIndex { + warnings = append(warnings, audit.Table+":missing_gist_geometry_index") + } + return warnings +} + +func (p *PostGIS) auditTable(ctx context.Context, table string) (TableAudit, error) { + query := fmt.Sprintf(` +SELECT + count(*)::bigint, + count(*) FILTER (WHERE geom IS NULL)::bigint, + count(*) FILTER (WHERE geom IS NOT NULL AND ST_IsEmpty(geom))::bigint, + count(*) FILTER (WHERE geom IS NOT NULL AND NOT ST_IsValid(geom))::bigint, + count(*) FILTER ( + WHERE geom IS NOT NULL + AND NOT ST_IsEmpty(geom) + AND ( + ST_XMin(Box3D(geom)) < -180 OR ST_XMax(Box3D(geom)) > 180 + OR ST_YMin(Box3D(geom)) < -90 OR ST_YMax(Box3D(geom)) > 90 + ) + )::bigint, + COALESCE(string_agg(DISTINCT GeometryType(geom), ',' ORDER BY GeometryType(geom)) FILTER (WHERE geom IS NOT NULL), ''), + COALESCE(string_agg(DISTINCT ST_SRID(geom)::text, ',' ORDER BY ST_SRID(geom)::text) FILTER (WHERE geom IS NOT NULL), ''), + EXISTS ( + SELECT 1 + FROM pg_indexes + WHERE schemaname = 'public' + AND tablename = $1 + AND indexdef ILIKE '%%USING gist%%' + AND indexdef ILIKE '%%geom%%' + ) +FROM public.%s`, quoteIdentifier(table)) + + var audit TableAudit + var geometryTypes string + var srids string + audit.Table = table + if err := p.pool.QueryRow(ctx, query, table).Scan( + &audit.RowCount, + &audit.NullGeometryCount, + &audit.EmptyGeometryCount, + &audit.InvalidCount, + &audit.OutOfBoundsCount, + &geometryTypes, + &srids, + &audit.HasGISTIndex, + ); err != nil { + return TableAudit{}, err + } + audit.GeometryTypes = splitNonEmpty(geometryTypes) + parsedSRIDs, err := parseSRIDs(srids) + if err != nil { + return TableAudit{}, err + } + audit.SRIDs = parsedSRIDs + return audit, nil +} + +// SearchPlaceCandidates returns bounded name matches from the fixed fire-safety source tables. +func (p *PostGIS) SearchPlaceCandidates(ctx context.Context, query domain.PlaceSearchQuery) ([]domain.PlaceCandidate, error) { + rows, err := p.pool.Query(ctx, placeCandidateSearchSQL, + query.PlaceName, + query.Scope.AllTowns, + query.Scope.AllowedTowns, + query.Limit, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]domain.PlaceCandidate, 0, query.Limit) + for rows.Next() { + var item domain.PlaceCandidate + if err := rows.Scan( + &item.PlaceType, + &item.SourceRecordID, + &item.Name, + &item.Town, + &item.Village, + &item.MatchedField, + &item.MatchedText, + &item.MatchKind, + &item.Location.Longitude, + &item.Location.Latitude, + &item.LocationKind, + ); err != nil { + return nil, err + } + items = append(items, item) + } + return items, rows.Err() +} + +// ResolveIncidentContext returns fire grids that cover the supplied point. +func (p *PostGIS) ResolveIncidentContext(ctx context.Context, point domain.Coordinate, scope domain.SpatialScope) ([]domain.IncidentContext, error) { + rows, err := p.pool.Query(ctx, resolveIncidentContextSQL, point.Longitude, point.Latitude, scope.AllTowns, scope.AllowedTowns) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]domain.IncidentContext, 0) + for rows.Next() { + var item domain.IncidentContext + if err := rows.Scan(&item.GridID, &item.Town, &item.AreaLabel); err != nil { + return nil, err + } + items = append(items, item) + } + return items, rows.Err() +} + +// FindNearbyWaterSources returns water-source and storage-pool candidates. +func (p *PostGIS) FindNearbyWaterSources(ctx context.Context, query domain.NearbyQuery) ([]domain.WaterSource, error) { + rows, err := p.pool.Query(ctx, nearbyWaterSourcesSQL, + query.Point.Longitude, + query.Point.Latitude, + query.Scope.AllTowns, + query.Scope.AllowedTowns, + query.RadiusMeters, + query.Limit, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]domain.WaterSource, 0, query.Limit) + for rows.Next() { + var item domain.WaterSource + var name, village, resourceType, reportedStatus pgtype.Text + var capacity pgtype.Float8 + var sourceTimestamp pgtype.Int8 + if err := rows.Scan( + &item.Category, + &item.SourceRecordID, + &name, + &item.Town, + &village, + &item.Location.Longitude, + &item.Location.Latitude, + &item.DistanceMeters, + &capacity, + &resourceType, + &reportedStatus, + &sourceTimestamp, + ); err != nil { + return nil, err + } + item.Name = optionalText(name) + item.Village = optionalText(village) + item.CapacityCubicMeters = optionalFloat64(capacity) + item.ResourceType = optionalText(resourceType) + item.ReportedStatus = optionalText(reportedStatus) + item.SourceTimestampRaw = optionalInt64(sourceTimestamp) + items = append(items, item) + } + return items, rows.Err() +} + +// FindCommandPostCandidates returns nearby check stations and lookout posts. +func (p *PostGIS) FindCommandPostCandidates(ctx context.Context, query domain.NearbyQuery) ([]domain.CommandPostCandidate, error) { + rows, err := p.pool.Query(ctx, commandPostCandidatesSQL, + query.Point.Longitude, + query.Point.Latitude, + query.Scope.AllTowns, + query.Scope.AllowedTowns, + query.RadiusMeters, + query.Limit, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]domain.CommandPostCandidate, 0, query.Limit) + for rows.Next() { + var item domain.CommandPostCandidate + var name, village, reportedStatus, managementUnit pgtype.Text + if err := rows.Scan( + &item.FacilityType, + &item.SourceRecordID, + &name, + &item.Town, + &village, + &item.Location.Longitude, + &item.Location.Latitude, + &item.DistanceMeters, + &reportedStatus, + &managementUnit, + ); err != nil { + return nil, err + } + item.Name = optionalText(name) + item.Village = optionalText(village) + item.ReportedStatus = optionalText(reportedStatus) + item.ManagementUnit = optionalText(managementUnit) + items = append(items, item) + } + return items, rows.Err() +} + +// ListNearbyAccessLines returns nearby fire access lines and closest access points. +func (p *PostGIS) ListNearbyAccessLines(ctx context.Context, query domain.NearbyQuery) ([]domain.AccessLine, error) { + rows, err := p.pool.Query(ctx, nearbyAccessLinesSQL, + query.Point.Longitude, + query.Point.Latitude, + query.Scope.AllTowns, + query.Scope.AllowedTowns, + query.RadiusMeters, + query.Limit, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]domain.AccessLine, 0, query.Limit) + for rows.Next() { + var item domain.AccessLine + var name, sourceUpdated pgtype.Text + var length pgtype.Float8 + if err := rows.Scan( + &item.SourceRecordID, + &item.Town, + &name, + &item.DistanceMeters, + &item.NearestPoint.Longitude, + &item.NearestPoint.Latitude, + &length, + &sourceUpdated, + ); err != nil { + return nil, err + } + item.Name = optionalText(name) + item.LengthMeters = optionalFloat64(length) + item.SourceUpdatedRaw = optionalText(sourceUpdated) + items = append(items, item) + } + return items, rows.Err() +} + +// GetResponsibleUnits returns non-personal responsibility details from covering grids. +func (p *PostGIS) GetResponsibleUnits(ctx context.Context, point domain.Coordinate, scope domain.SpatialScope) ([]domain.ResponsibleUnit, error) { + rows, err := p.pool.Query(ctx, responsibleUnitsSQL, point.Longitude, point.Latitude, scope.AllTowns, scope.AllowedTowns) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]domain.ResponsibleUnit, 0) + for rows.Next() { + var item domain.ResponsibleUnit + if err := rows.Scan(&item.GridID, &item.Town, &item.AreaLabel, &item.FireTeam); err != nil { + return nil, err + } + items = append(items, item) + } + return items, rows.Err() +} + +// FindNearbyRiskAreas returns nearby cemetery and forest-enterprise polygons. +func (p *PostGIS) FindNearbyRiskAreas(ctx context.Context, query domain.NearbyQuery) ([]domain.RiskArea, error) { + rows, err := p.pool.Query(ctx, nearbyRiskAreasSQL, + query.Point.Longitude, + query.Point.Latitude, + query.Scope.AllTowns, + query.Scope.AllowedTowns, + query.RadiusMeters, + query.Limit, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + items := make([]domain.RiskArea, 0, query.Limit) + for rows.Next() { + var item domain.RiskArea + var name, village, direction pgtype.Text + if err := rows.Scan( + &item.RiskType, + &item.SourceRecordID, + &name, + &item.Town, + &village, + &direction, + &item.CoversPoint, + &item.DistanceMeters, + ); err != nil { + return nil, err + } + item.Name = optionalText(name) + item.Village = optionalText(village) + item.Direction = optionalText(direction) + items = append(items, item) + } + return items, rows.Err() +} + +type spatialTableSpec struct { + name string + geometryTypes []string +} + +var spatialTableSpecs = []spatialTableSpec{ + {name: "st_2_mpslfh_t_slfh_syd", geometryTypes: []string{"POINT"}}, + {name: "st_2_xianyouxushuichiguan", geometryTypes: []string{"POINT"}}, + {name: "st_2_xianyoufanghuotongdao", geometryTypes: []string{"LINESTRING", "MULTILINESTRING"}}, + {name: "st_2_fanghuojianchazhan", geometryTypes: []string{"POINT"}}, + {name: "st_2_fanghuoliaowangshao", geometryTypes: []string{"POINT"}}, + {name: "st_2_fanghuowangge", geometryTypes: []string{"POLYGON", "MULTIPOLYGON"}}, + {name: "st_2_linqugongkuangqiye", geometryTypes: []string{"POLYGON", "MULTIPOLYGON"}}, + {name: "st_2_mudifenqu_mian", geometryTypes: []string{"POLYGON", "MULTIPOLYGON"}}, +} + +func quoteIdentifier(identifier string) string { + return `"` + strings.ReplaceAll(identifier, `"`, `""`) + `"` +} + +func splitNonEmpty(value string) []string { + if value == "" { + return []string{} + } + return strings.Split(value, ",") +} + +func parseSRIDs(value string) ([]int, error) { + parts := splitNonEmpty(value) + result := make([]int, 0, len(parts)) + for _, part := range parts { + srid, err := strconv.Atoi(part) + if err != nil { + return nil, fmt.Errorf("parse geometry SRID: %w", err) + } + result = append(result, srid) + } + return result, nil +} + +func optionalText(value pgtype.Text) string { + if !value.Valid { + return "" + } + return value.String +} + +func optionalFloat64(value pgtype.Float8) *float64 { + if !value.Valid { + return nil + } + result := value.Float64 + return &result +} + +func optionalInt64(value pgtype.Int8) *int64 { + if !value.Valid { + return nil + } + result := value.Int64 + return &result +} + +const placeCandidateSearchSQL = ` +WITH candidates AS ( + SELECT + 'water_source'::text AS place_type, + w.objectid::text AS source_record_id, + NULLIF(w.name, '') AS name, + COALESCE(w.zj, '') AS town, + NULL::text AS village, + ST_Force2D(w.geom) AS location, + 'recorded_point'::text AS location_kind + FROM public.st_2_mpslfh_t_slfh_syd AS w + WHERE ($2::boolean OR w.zj = ANY($3::text[])) + AND w.geom IS NOT NULL + AND ST_SRID(w.geom) = 4326 + AND GeometryType(w.geom) = 'POINT' + AND ST_IsValid(w.geom) + AND NOT ST_IsEmpty(w.geom) + + UNION ALL + + SELECT + 'storage_pool'::text, + r.gid::text, + NULLIF(r.community, ''), + COALESCE(r.auth, ''), + NULLIF(r.community, ''), + ST_Force2D(r.geom), + 'recorded_point'::text + FROM public.st_2_xianyouxushuichiguan AS r + WHERE ($2::boolean OR r.auth = ANY($3::text[])) + AND r.geom IS NOT NULL + AND ST_SRID(r.geom) = 4326 + AND GeometryType(r.geom) = 'POINT' + AND ST_IsValid(r.geom) + AND NOT ST_IsEmpty(r.geom) + + UNION ALL + + SELECT + 'fire_access_line'::text, + a.gid::text, + NULLIF(a.name, ''), + COALESCE(a.auth, ''), + NULL::text, + ST_StartPoint( + CASE + WHEN GeometryType(a.geom) = 'LINESTRING' THEN ST_Force2D(a.geom) + ELSE ST_GeometryN(ST_Force2D(a.geom), 1) + END + ), + 'representative_point'::text + FROM public.st_2_xianyoufanghuotongdao AS a + WHERE ($2::boolean OR a.auth = ANY($3::text[])) + AND a.geom IS NOT NULL + AND ST_SRID(a.geom) = 4326 + AND GeometryType(a.geom) IN ('LINESTRING', 'MULTILINESTRING') + AND ST_IsValid(a.geom) + AND NOT ST_IsEmpty(a.geom) + + UNION ALL + + SELECT + 'fire_check_station'::text, + s.gid::text, + NULLIF(s.jczmc, ''), + COALESCE(s.zjmc, ''), + NULLIF(s.cmc, ''), + ST_Force2D(s.geom), + 'recorded_point'::text + FROM public.st_2_fanghuojianchazhan AS s + WHERE ($2::boolean OR s.zjmc = ANY($3::text[])) + AND s.geom IS NOT NULL + AND ST_SRID(s.geom) = 4326 + AND GeometryType(s.geom) = 'POINT' + AND ST_IsValid(s.geom) + AND NOT ST_IsEmpty(s.geom) + + UNION ALL + + SELECT + 'fire_lookout'::text, + l.gid::text, + NULLIF(l.lwsmc, ''), + COALESCE(l.zjmc, ''), + NULLIF(l.cmc, ''), + ST_Force2D(l.geom), + 'recorded_point'::text + FROM public.st_2_fanghuoliaowangshao AS l + WHERE ($2::boolean OR l.zjmc = ANY($3::text[])) + AND l.geom IS NOT NULL + AND ST_SRID(l.geom) = 4326 + AND GeometryType(l.geom) = 'POINT' + AND ST_IsValid(l.geom) + AND NOT ST_IsEmpty(l.geom) + + UNION ALL + + SELECT + 'fire_grid'::text, + g.gid::text, + NULLIF(g.name, ''), + COALESCE(g.auth, ''), + NULL::text, + ST_PointOnSurface(ST_Force2D(g.geom)), + 'representative_point'::text + FROM public.st_2_fanghuowangge AS g + WHERE ($2::boolean OR g.auth = ANY($3::text[])) + AND g.geom IS NOT NULL + AND ST_SRID(g.geom) = 4326 + AND GeometryType(g.geom) IN ('POLYGON', 'MULTIPOLYGON') + AND ST_IsValid(g.geom) + AND NOT ST_IsEmpty(g.geom) + + UNION ALL + + SELECT + 'forest_enterprise'::text, + e.gid::text, + NULLIF(e.lx, ''), + COALESCE(e.auth, ''), + NULLIF(e.c, ''), + ST_PointOnSurface(ST_Force2D(e.geom)), + 'representative_point'::text + FROM public.st_2_linqugongkuangqiye AS e + WHERE ($2::boolean OR e.auth = ANY($3::text[])) + AND e.geom IS NOT NULL + AND ST_SRID(e.geom) = 4326 + AND GeometryType(e.geom) IN ('POLYGON', 'MULTIPOLYGON') + AND ST_IsValid(e.geom) + AND NOT ST_IsEmpty(e.geom) + + UNION ALL + + SELECT + 'cemetery_area'::text, + c.gid::text, + NULLIF(c.name, ''), + COALESCE(c.zj, ''), + NULLIF(c.cmc, ''), + ST_PointOnSurface(ST_Force2D(c.geom)), + 'representative_point'::text + FROM public.st_2_mudifenqu_mian AS c + WHERE ($2::boolean OR c.zj = ANY($3::text[])) + AND c.geom IS NOT NULL + AND ST_SRID(c.geom) = 4326 + AND GeometryType(c.geom) IN ('POLYGON', 'MULTIPOLYGON') + AND ST_IsValid(c.geom) + AND NOT ST_IsEmpty(c.geom) +), matched AS ( + SELECT + candidates.*, + CASE + WHEN lower(COALESCE(name, '')) = lower($1::text) THEN 'name' + WHEN lower(COALESCE(village, '')) = lower($1::text) THEN 'village' + WHEN lower(COALESCE(town, '')) = lower($1::text) THEN 'town' + WHEN strpos(lower(COALESCE(name, '')), lower($1::text)) = 1 THEN 'name' + WHEN strpos(lower(COALESCE(village, '')), lower($1::text)) = 1 THEN 'village' + WHEN strpos(lower(COALESCE(town, '')), lower($1::text)) = 1 THEN 'town' + WHEN strpos(lower(COALESCE(name, '')), lower($1::text)) > 0 THEN 'name' + WHEN strpos(lower(COALESCE(village, '')), lower($1::text)) > 0 THEN 'village' + ELSE 'town' + END AS matched_field, + CASE + WHEN lower(COALESCE(name, '')) = lower($1::text) THEN 0 + WHEN lower(COALESCE(village, '')) = lower($1::text) THEN 1 + WHEN lower(COALESCE(town, '')) = lower($1::text) THEN 2 + WHEN strpos(lower(COALESCE(name, '')), lower($1::text)) = 1 THEN 3 + WHEN strpos(lower(COALESCE(village, '')), lower($1::text)) = 1 THEN 4 + WHEN strpos(lower(COALESCE(town, '')), lower($1::text)) = 1 THEN 5 + WHEN strpos(lower(COALESCE(name, '')), lower($1::text)) > 0 THEN 6 + WHEN strpos(lower(COALESCE(village, '')), lower($1::text)) > 0 THEN 7 + ELSE 8 + END AS match_rank + FROM candidates + WHERE strpos(lower(COALESCE(name, '')), lower($1::text)) > 0 + OR strpos(lower(COALESCE(village, '')), lower($1::text)) > 0 + OR strpos(lower(COALESCE(town, '')), lower($1::text)) > 0 +) +SELECT + place_type, + source_record_id, + LEFT(COALESCE(name, ''), 255), + LEFT(COALESCE(town, ''), 100), + LEFT(COALESCE(village, ''), 255), + matched_field, + LEFT(CASE matched_field WHEN 'name' THEN name WHEN 'village' THEN village ELSE town END, 255), + CASE WHEN match_rank <= 2 THEN 'exact' ELSE 'partial' END, + round(ST_X(location)::numeric, 6)::float8, + round(ST_Y(location)::numeric, 6)::float8, + location_kind +FROM matched +ORDER BY match_rank, town, village, name, place_type, source_record_id +LIMIT $4` + +const resolveIncidentContextSQL = ` +WITH incident AS ( + SELECT ST_SetSRID(ST_MakePoint($1, $2), 4326) AS geom +) +SELECT g.gid::text, LEFT(COALESCE(g.auth, ''), 100), LEFT(COALESCE(g.name, ''), 255) +FROM public.st_2_fanghuowangge AS g +CROSS JOIN incident AS i +WHERE ($3::boolean OR g.auth = ANY($4::text[])) + AND g.geom IS NOT NULL + AND ST_SRID(g.geom) = 4326 + AND GeometryType(g.geom) IN ('POLYGON', 'MULTIPOLYGON') + AND ST_IsValid(g.geom) + AND NOT ST_IsEmpty(g.geom) + AND ST_Covers(g.geom, i.geom) +ORDER BY ST_Area(g.geom::geography), g.gid +LIMIT 20` + +const nearbyWaterSourcesSQL = ` +WITH incident AS ( + SELECT ST_SetSRID(ST_MakePoint($1, $2), 4326) AS geom +), candidates AS ( + SELECT + 'water_source'::text AS category, + w.objectid::text AS source_record_id, + LEFT(NULLIF(w.name, ''), 255) AS name, + LEFT(COALESCE(w.zj, ''), 100) AS town, + NULL::text AS village, + ST_X(w.geom) AS longitude, + ST_Y(w.geom) AS latitude, + ST_Distance(w.geom::geography, i.geom::geography) AS distance_meters, + w.xsrl::float8 AS capacity_cubic_meters, + LEFT(NULLIF(w.lx, ''), 100) AS resource_type, + LEFT(NULLIF(w.syzt, ''), 100) AS reported_status, + w.hc_datetime AS source_timestamp_raw + FROM public.st_2_mpslfh_t_slfh_syd AS w + CROSS JOIN incident AS i + WHERE ($3::boolean OR w.zj = ANY($4::text[])) + AND w.geom IS NOT NULL + AND ST_SRID(w.geom) = 4326 + AND GeometryType(w.geom) = 'POINT' + AND ST_IsValid(w.geom) + AND NOT ST_IsEmpty(w.geom) + AND w.geom && ST_Expand(i.geom, $5 / GREATEST(1000.0, 111320.0 * abs(cos(radians($2))))) + AND ST_DWithin(w.geom::geography, i.geom::geography, $5) + + UNION ALL + + SELECT + 'storage_pool'::text, + r.gid::text, + LEFT(NULLIF(r.community, ''), 255), + LEFT(COALESCE(r.auth, ''), 100), + LEFT(NULLIF(r.community, ''), 255), + ST_X(r.geom), + ST_Y(r.geom), + ST_Distance(r.geom::geography, i.geom::geography), + r.lfs::float8, + LEFT(NULLIF(r.lx, ''), 100), + NULL::text, + NULL::bigint + FROM public.st_2_xianyouxushuichiguan AS r + CROSS JOIN incident AS i + WHERE ($3::boolean OR r.auth = ANY($4::text[])) + AND r.geom IS NOT NULL + AND ST_SRID(r.geom) = 4326 + AND GeometryType(r.geom) = 'POINT' + AND ST_IsValid(r.geom) + AND NOT ST_IsEmpty(r.geom) + AND r.geom && ST_Expand(i.geom, $5 / GREATEST(1000.0, 111320.0 * abs(cos(radians($2))))) + AND ST_DWithin(r.geom::geography, i.geom::geography, $5) +) +SELECT category, source_record_id, name, town, village, + round(longitude::numeric, 6)::float8, round(latitude::numeric, 6)::float8, + round(distance_meters::numeric, 1)::float8, capacity_cubic_meters, resource_type, reported_status, + source_timestamp_raw +FROM candidates +ORDER BY distance_meters, category, source_record_id +LIMIT $6` + +const commandPostCandidatesSQL = ` +WITH incident AS ( + SELECT ST_SetSRID(ST_MakePoint($1, $2), 4326) AS geom +), candidates AS ( + SELECT + 'fire_check_station'::text AS facility_type, + s.gid::text AS source_record_id, + LEFT(NULLIF(s.jczmc, ''), 255) AS name, + LEFT(COALESCE(s.zjmc, ''), 100) AS town, + LEFT(NULLIF(s.cmc, ''), 255) AS village, + ST_X(s.geom) AS longitude, + ST_Y(s.geom) AS latitude, + ST_Distance(s.geom::geography, i.geom::geography) AS distance_meters, + LEFT(NULLIF(s.syzt, ''), 100) AS reported_status, + LEFT(NULLIF(s.gldw, ''), 255) AS management_unit + FROM public.st_2_fanghuojianchazhan AS s + CROSS JOIN incident AS i + WHERE ($3::boolean OR s.zjmc = ANY($4::text[])) + AND s.geom IS NOT NULL + AND ST_SRID(s.geom) = 4326 + AND GeometryType(s.geom) = 'POINT' + AND ST_IsValid(s.geom) + AND NOT ST_IsEmpty(s.geom) + AND s.geom && ST_Expand(i.geom, $5 / GREATEST(1000.0, 111320.0 * abs(cos(radians($2))))) + AND ST_DWithin(s.geom::geography, i.geom::geography, $5) + + UNION ALL + + SELECT + 'fire_lookout'::text, + l.gid::text, + LEFT(NULLIF(l.lwsmc, ''), 255), + LEFT(COALESCE(l.zjmc, ''), 100), + LEFT(NULLIF(l.cmc, ''), 255), + ST_X(l.geom), + ST_Y(l.geom), + ST_Distance(l.geom::geography, i.geom::geography), + LEFT(NULLIF(l.syzt, ''), 100), + NULL::text + FROM public.st_2_fanghuoliaowangshao AS l + CROSS JOIN incident AS i + WHERE ($3::boolean OR l.zjmc = ANY($4::text[])) + AND l.geom IS NOT NULL + AND ST_SRID(l.geom) = 4326 + AND GeometryType(l.geom) = 'POINT' + AND ST_IsValid(l.geom) + AND NOT ST_IsEmpty(l.geom) + AND l.geom && ST_Expand(i.geom, $5 / GREATEST(1000.0, 111320.0 * abs(cos(radians($2))))) + AND ST_DWithin(l.geom::geography, i.geom::geography, $5) +) +SELECT facility_type, source_record_id, name, town, village, + round(longitude::numeric, 6)::float8, round(latitude::numeric, 6)::float8, + round(distance_meters::numeric, 1)::float8, reported_status, management_unit +FROM candidates +ORDER BY distance_meters, facility_type, source_record_id +LIMIT $6` + +const nearbyAccessLinesSQL = ` +WITH incident AS ( + SELECT ST_SetSRID(ST_MakePoint($1, $2), 4326) AS geom +), candidates AS ( + SELECT + a.gid::text AS source_record_id, + LEFT(COALESCE(a.auth, ''), 100) AS town, + LEFT(NULLIF(a.name, ''), 255) AS name, + ST_Distance(a.geom::geography, i.geom::geography) AS distance_meters, + ST_ClosestPoint(a.geom, i.geom) AS nearest_point, + ST_Length(a.geom::geography) AS length_meters, + a.sttime::text AS source_updated_raw + FROM public.st_2_xianyoufanghuotongdao AS a + CROSS JOIN incident AS i + WHERE ($3::boolean OR a.auth = ANY($4::text[])) + AND a.geom IS NOT NULL + AND ST_SRID(a.geom) = 4326 + AND GeometryType(a.geom) IN ('LINESTRING', 'MULTILINESTRING') + AND ST_IsValid(a.geom) + AND NOT ST_IsEmpty(a.geom) + AND a.geom && ST_Expand(i.geom, $5 / GREATEST(1000.0, 111320.0 * abs(cos(radians($2))))) + AND ST_DWithin(a.geom::geography, i.geom::geography, $5) +) +SELECT source_record_id, town, name, round(distance_meters::numeric, 1)::float8, + round(ST_X(nearest_point)::numeric, 6)::float8, + round(ST_Y(nearest_point)::numeric, 6)::float8, + length_meters, source_updated_raw +FROM candidates +ORDER BY distance_meters, source_record_id +LIMIT $6` + +const responsibleUnitsSQL = ` +WITH incident AS ( + SELECT ST_SetSRID(ST_MakePoint($1, $2), 4326) AS geom +) +SELECT g.gid::text, LEFT(COALESCE(g.auth, ''), 100), LEFT(COALESCE(g.name, ''), 255), LEFT(COALESCE(g.fhzd, ''), 255) +FROM public.st_2_fanghuowangge AS g +CROSS JOIN incident AS i +WHERE ($3::boolean OR g.auth = ANY($4::text[])) + AND g.geom IS NOT NULL + AND ST_SRID(g.geom) = 4326 + AND GeometryType(g.geom) IN ('POLYGON', 'MULTIPOLYGON') + AND ST_IsValid(g.geom) + AND NOT ST_IsEmpty(g.geom) + AND ST_Covers(g.geom, i.geom) +ORDER BY ST_Area(g.geom::geography), g.gid +LIMIT 20` + +const nearbyRiskAreasSQL = ` +WITH incident AS ( + SELECT ST_SetSRID(ST_MakePoint($1, $2), 4326) AS geom +), candidates AS ( + SELECT + 'cemetery_area'::text AS risk_type, + c.gid::text AS source_record_id, + LEFT(NULLIF(c.name, ''), 255) AS name, + LEFT(COALESCE(c.zj, ''), 100) AS town, + LEFT(NULLIF(c.cmc, ''), 255) AS village, + LEFT(NULLIF(c.fw, ''), 255) AS direction, + ST_Covers(c.geom, i.geom) AS covers_point, + ST_Distance(c.geom::geography, i.geom::geography) AS distance_meters + FROM public.st_2_mudifenqu_mian AS c + CROSS JOIN incident AS i + WHERE ($3::boolean OR c.zj = ANY($4::text[])) + AND c.geom IS NOT NULL + AND ST_SRID(c.geom) = 4326 + AND GeometryType(c.geom) IN ('POLYGON', 'MULTIPOLYGON') + AND ST_IsValid(c.geom) + AND NOT ST_IsEmpty(c.geom) + AND c.geom && ST_Expand(i.geom, $5 / GREATEST(1000.0, 111320.0 * abs(cos(radians($2))))) + AND ST_DWithin(c.geom::geography, i.geom::geography, $5) + + UNION ALL + + SELECT + 'forest_enterprise'::text, + e.gid::text, + LEFT(NULLIF(e.lx, ''), 255), + LEFT(COALESCE(e.auth, ''), 100), + LEFT(NULLIF(e.c, ''), 255), + LEFT(NULLIF(e.fw, ''), 255), + ST_Covers(e.geom, i.geom), + ST_Distance(e.geom::geography, i.geom::geography) + FROM public.st_2_linqugongkuangqiye AS e + CROSS JOIN incident AS i + WHERE ($3::boolean OR e.auth = ANY($4::text[])) + AND e.geom IS NOT NULL + AND ST_SRID(e.geom) = 4326 + AND GeometryType(e.geom) IN ('POLYGON', 'MULTIPOLYGON') + AND ST_IsValid(e.geom) + AND NOT ST_IsEmpty(e.geom) + AND e.geom && ST_Expand(i.geom, $5 / GREATEST(1000.0, 111320.0 * abs(cos(radians($2))))) + AND ST_DWithin(e.geom::geography, i.geom::geography, $5) +) +SELECT risk_type, source_record_id, name, town, village, direction, covers_point, + round(distance_meters::numeric, 1)::float8 +FROM candidates +ORDER BY covers_point DESC, distance_meters, risk_type, source_record_id +LIMIT $6` diff --git a/internal/repository/postgis_test.go b/internal/repository/postgis_test.go new file mode 100644 index 0000000..2ae162b --- /dev/null +++ b/internal/repository/postgis_test.go @@ -0,0 +1,231 @@ +package repository + +import ( + "context" + "slices" + "strings" + "testing" + "time" +) + +func TestOpenPostGISDoesNotExposeInvalidDSN(t *testing.T) { + secret := "do-not-leak" + _, err := OpenPostGIS(context.Background(), PostGISOptions{ + DSN: "postgresql://user:" + secret + "%zz@localhost/database", + MaxConns: 1, + ConnectTimeout: time.Second, + QueryTimeout: time.Second, + }) + if err == nil { + t.Fatal("OpenPostGIS() error = nil") + } + if strings.Contains(err.Error(), secret) { + t.Fatalf("OpenPostGIS() leaked DSN content: %v", err) + } +} + +func TestValidateReadinessAcceptsExpectedGeometryContracts(t *testing.T) { + report := validReadinessReport() + if err := validateReadiness(report, 4326); err != nil { + t.Fatalf("validateReadiness() error = %v", err) + } +} + +func TestValidateReadinessRejectsUnsafeSpatialMetadata(t *testing.T) { + tests := []struct { + name string + mutate func(*ReadinessReport) + want string + }{ + { + name: "SRID zero", + mutate: func(report *ReadinessReport) { + report.Tables[0].SRIDs = []int{0} + }, + want: "unexpected SRID 0", + }, + { + name: "mixed SRID", + mutate: func(report *ReadinessReport) { + report.Tables[1].SRIDs = []int{4326, 3857} + }, + want: "unexpected SRID 3857", + }, + { + name: "unexpected geometry type", + mutate: func(report *ReadinessReport) { + report.Tables[2].GeometryTypes = []string{"POLYGON"} + }, + want: "unexpected geometry type POLYGON", + }, + { + name: "coordinate outside WGS84", + mutate: func(report *ReadinessReport) { + report.Tables[0].OutOfBoundsCount = 1 + }, + want: "outside WGS84 bounds", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + report := validReadinessReport() + tt.mutate(&report) + err := validateReadiness(report, 4326) + if err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("validateReadiness() error = %v, want %q", err, tt.want) + } + }) + } +} + +func TestValidateReadinessAllowsInvalidGeometriesThatQueriesExclude(t *testing.T) { + report := validReadinessReport() + report.Tables[5].InvalidCount = 4 + report.Tables[6].InvalidCount = 4 + report.Tables[7].InvalidCount = 27 + + if err := validateReadiness(report, 4326); err != nil { + t.Fatalf("validateReadiness() error = %v", err) + } +} + +func TestTableAuditWarningsExposeExcludedInvalidGeometries(t *testing.T) { + audit := TableAudit{ + Table: "st_2_fanghuowangge", + RowCount: 2, + NullGeometryCount: 1, + EmptyGeometryCount: 1, + InvalidCount: 1, + HasGISTIndex: false, + } + + warnings := tableAuditWarnings(audit) + for _, want := range []string{ + "st_2_fanghuowangge:null_geometries_excluded", + "st_2_fanghuowangge:empty_geometries_excluded", + "st_2_fanghuowangge:invalid_geometries_excluded", + "st_2_fanghuowangge:missing_gist_geometry_index", + } { + if !slices.Contains(warnings, want) { + t.Fatalf("warnings = %#v, want %q", warnings, want) + } + } +} + +func TestSpatialQueriesKeepScopeAndBoundsParameterized(t *testing.T) { + nearbyQueries := []struct { + name string + query string + wantScopePredicates int + }{ + {name: "water sources", query: nearbyWaterSourcesSQL, wantScopePredicates: 2}, + {name: "command posts", query: commandPostCandidatesSQL, wantScopePredicates: 2}, + {name: "access lines", query: nearbyAccessLinesSQL, wantScopePredicates: 1}, + {name: "risk areas", query: nearbyRiskAreasSQL, wantScopePredicates: 2}, + } + for _, tt := range nearbyQueries { + query := tt.query + for _, required := range []string{"$3::boolean", "ANY($4::text[])", "$5", "LIMIT $6", "ST_SRID", "ST_IsValid", "ST_IsEmpty", "ST_Expand"} { + if !strings.Contains(query, required) { + t.Fatalf("%s query lacks %q", tt.name, required) + } + } + if got := strings.Count(query, "($3::boolean OR"); got != tt.wantScopePredicates { + t.Fatalf("%s query has %d scope predicates, want %d", tt.name, got, tt.wantScopePredicates) + } + if got := strings.Count(query, "ANY($4::text[])"); got != tt.wantScopePredicates { + t.Fatalf("%s query has %d town predicates, want %d", tt.name, got, tt.wantScopePredicates) + } + } + for index, query := range []string{resolveIncidentContextSQL, responsibleUnitsSQL} { + for _, required := range []string{"$3::boolean", "ANY($4::text[])", "ST_SRID", "ST_IsValid", "ST_Covers"} { + if !strings.Contains(query, required) { + t.Fatalf("cover query %d lacks %q", index, required) + } + } + if got := strings.Count(query, "($3::boolean OR"); got != 1 { + t.Fatalf("cover query %d has %d scope predicates, want 1", index, got) + } + } +} + +func TestPlaceCandidateSearchKeepsTermScopeAndLimitParameterized(t *testing.T) { + for _, required := range []string{ + "lower($1::text)", + "($2::boolean OR", + "ANY($3::text[])", + "LIMIT $4", + "strpos", + "recorded_point", + "representative_point", + "ST_Force2D", + "ST_IsValid", + } { + if !strings.Contains(placeCandidateSearchSQL, required) { + t.Fatalf("place candidate query lacks %q", required) + } + } + if got := strings.Count(placeCandidateSearchSQL, "($2::boolean OR"); got != len(spatialTableSpecs) { + t.Fatalf("place candidate query has %d scope predicates, want %d", got, len(spatialTableSpecs)) + } + if got := strings.Count(placeCandidateSearchSQL, "ST_IsValid("); got != len(spatialTableSpecs) { + t.Fatalf("place candidate query has %d validity filters, want %d", got, len(spatialTableSpecs)) + } + for _, spec := range spatialTableSpecs { + if !strings.Contains(placeCandidateSearchSQL, "public."+spec.name) { + t.Fatalf("place candidate query does not include %s", spec.name) + } + } +} + +func TestSpatialQueriesDoNotSelectSensitiveContactColumns(t *testing.T) { + queries := strings.Join([]string{ + placeCandidateSearchSQL, + resolveIncidentContextSQL, + nearbyWaterSourcesSQL, + commandPostCandidatesSQL, + nearbyAccessLinesSQL, + responsibleUnitsSQL, + nearbyRiskAreasSQL, + }, "\n") + for _, forbidden := range []string{ + "bpld", + "bplddh", + "csj", + "csjdh", + "lxdh", + "zjwgfzr", + "zjfzrlxfs", + "fhzddc", + "fhzddclxfs", + "chzbry", + "shzbry", + "zbry", + ".bz", + } { + if strings.Contains(strings.ToLower(queries), strings.ToLower(forbidden)) { + t.Fatalf("spatial SQL references sensitive contact column %s", forbidden) + } + } +} + +func TestQuoteIdentifierEscapesEmbeddedQuotes(t *testing.T) { + if got, want := quoteIdentifier(`safe"name`), `"safe""name"`; got != want { + t.Fatalf("quoteIdentifier() = %q, want %q", got, want) + } +} + +func validReadinessReport() ReadinessReport { + report := ReadinessReport{Tables: make([]TableAudit, 0, len(spatialTableSpecs))} + for _, spec := range spatialTableSpecs { + report.Tables = append(report.Tables, TableAudit{ + Table: spec.name, + RowCount: 2, + GeometryTypes: []string{spec.geometryTypes[0]}, + SRIDs: []int{4326}, + HasGISTIndex: true, + }) + } + return report +} diff --git a/internal/service/chat.go b/internal/service/chat.go new file mode 100644 index 0000000..0365f0a --- /dev/null +++ b/internal/service/chat.go @@ -0,0 +1,455 @@ +package service + +import ( + "context" + "crypto/rand" + "encoding/base64" + "errors" + "fmt" + "regexp" + "strings" + "sync" + "time" + "unicode" +) + +const ( + maximumChatMessageBytes int64 = 16 * 1024 * 1024 + maximumChatSessionTTL = 24 * time.Hour + maximumChatSessions = 10_000 +) + +var ( + // ErrChatInvalidArgument indicates an invalid user-facing chat request. + ErrChatInvalidArgument = errors.New("invalid chat argument") + // ErrChatConversationNotFound indicates an unknown, expired, or lost in-memory conversation. + ErrChatConversationNotFound = errors.New("chat conversation not found") + // ErrChatConversationBusy indicates that a conversation already has an active run. + ErrChatConversationBusy = errors.New("chat conversation is busy") + // ErrChatCapacityReached indicates that the bounded in-memory conversation store is full. + ErrChatCapacityReached = errors.New("chat conversation capacity reached") + // ErrChatUpstreamUnavailable indicates a provider transport or availability failure. + ErrChatUpstreamUnavailable = errors.New("chat upstream unavailable") + // ErrChatUpstreamProtocol indicates that the provider returned an unsafe or incomplete protocol result. + ErrChatUpstreamProtocol = errors.New("chat upstream protocol error") + // ErrChatRunFailed indicates that the provider explicitly reported a failed run. + ErrChatRunFailed = errors.New("chat upstream run failed") + // ErrChatInternal indicates a local failure that cannot be attributed to user input. + ErrChatInternal = errors.New("chat internal failure") + + conversationIDPattern = regexp.MustCompile(`^conv_[A-Za-z0-9_-]{16,128}$`) +) + +// ChatAgent is the provider-neutral outbound boundary used by ChatService. +type ChatAgent interface { + CreateSession(context.Context, AgentCreateSessionRequest) (AgentSession, error) + StreamMessage(context.Context, AgentMessageRequest, func(AgentTraceEvent)) (AgentMessageResult, error) +} + +// AgentCreateSessionRequest contains only server-controlled identity and metadata. +type AgentCreateSessionRequest struct { + ExternalSubjectID string + IdempotencyKey string + RequestID string + Metadata map[string]any +} + +// AgentSession is the provider's opaque session handle. +type AgentSession struct { + ID string +} + +// AgentMessageRequest sends one turn to an already-created provider session. +type AgentMessageRequest struct { + SessionID string + Message string + IdempotencyKey string + RequestID string + Metadata map[string]any +} + +// AgentTraceEvent is the deliberately small public progress projection. It +// excludes text, identifiers, tool inputs, tool outputs, and provider payloads. +type AgentTraceEvent struct { + Event string `json:"event"` + ToolName string `json:"tool_name,omitempty"` + Status string `json:"status,omitempty"` +} + +// ChatTokenUsage contains non-negative provider-reported token counts. +type ChatTokenUsage struct { + Input int64 `json:"input"` + Output int64 `json:"output"` + Total int64 `json:"total"` +} + +// AgentMessageResult is returned by a ChatAgent only after strict provider success. +type AgentMessageResult struct { + RunID string + ModelID string + Answer string + Usage ChatTokenUsage +} + +// ChatRequest is the validated application-level request for one turn. +type ChatRequest struct { + Message string + ConversationID string + RequestID string +} + +// ChatResult is safe for the inbound handler to expose after strict success. +type ChatResult struct { + ConversationID string `json:"conversation_id"` + RunID string `json:"run_id,omitempty"` + ModelID string `json:"model_id,omitempty"` + Answer string `json:"answer"` + Usage ChatTokenUsage `json:"usage"` +} + +// ChatTurn represents an exclusively acquired conversation turn. Call Close if +// Stream is not called; Stream releases or invalidates the conversation itself. +type ChatTurn interface { + ConversationID() string + Reused() bool + Stream(context.Context, func(AgentTraceEvent)) (ChatResult, error) + Close() +} + +// ChatOptions controls the bounded, process-local conversation store. +type ChatOptions struct { + ExternalSubjectID string + MaxMessageBytes int64 + SessionTTL time.Duration + MaxSessions int + + now func() time.Time + newID func(string) (string, error) +} + +// ChatService owns local-to-provider session mapping and per-conversation run exclusion. +type ChatService struct { + agent ChatAgent + externalSubjectID string + maxMessageBytes int64 + sessionTTL time.Duration + maxSessions int + now func() time.Time + newID func(string) (string, error) + + mu sync.Mutex + conversations map[string]*chatConversation + reserved map[string]struct{} +} + +type chatConversation struct { + providerSessionID string + busy bool + lastUsed time.Time +} + +// NewChatService constructs a bounded in-memory chat service. +func NewChatService(agent ChatAgent, options ChatOptions) (*ChatService, error) { + if agent == nil { + return nil, errors.New("chat agent is required") + } + options.ExternalSubjectID = strings.TrimSpace(options.ExternalSubjectID) + if options.ExternalSubjectID == "" || len(options.ExternalSubjectID) > 512 || containsControl(options.ExternalSubjectID) { + return nil, errors.New("chat external subject ID is invalid") + } + if options.MaxMessageBytes <= 0 || options.MaxMessageBytes > maximumChatMessageBytes { + return nil, errors.New("chat message byte limit is outside the supported range") + } + if options.SessionTTL <= 0 || options.SessionTTL > maximumChatSessionTTL { + return nil, errors.New("chat session TTL is outside the supported range") + } + if options.MaxSessions <= 0 || options.MaxSessions > maximumChatSessions { + return nil, errors.New("chat session capacity is outside the supported range") + } + if options.now == nil { + options.now = time.Now + } + if options.newID == nil { + options.newID = randomChatID + } + return &ChatService{ + agent: agent, + externalSubjectID: options.ExternalSubjectID, + maxMessageBytes: options.MaxMessageBytes, + sessionTTL: options.SessionTTL, + maxSessions: options.MaxSessions, + now: options.now, + newID: options.newID, + conversations: make(map[string]*chatConversation), + reserved: make(map[string]struct{}), + }, nil +} + +// Prepare validates a request and exclusively acquires either a new or existing conversation. +func (s *ChatService) Prepare(ctx context.Context, request ChatRequest) (ChatTurn, error) { + message := strings.TrimSpace(request.Message) + if message == "" || int64(len([]byte(message))) > s.maxMessageBytes { + return nil, fmt.Errorf("%w: message is required and bounded", ErrChatInvalidArgument) + } + requestID := strings.TrimSpace(request.RequestID) + if requestID == "" { + generated, err := s.newID("req_") + if err != nil { + return nil, fmt.Errorf("%w: generate request ID", ErrChatInternal) + } + requestID = generated + } + if !validCorrelationID(requestID) { + return nil, fmt.Errorf("%w: request ID is invalid", ErrChatInvalidArgument) + } + + conversationID := strings.TrimSpace(request.ConversationID) + if conversationID != "" { + if !conversationIDPattern.MatchString(conversationID) { + return nil, fmt.Errorf("%w: conversation ID is invalid", ErrChatInvalidArgument) + } + return s.prepareExisting(message, conversationID, requestID) + } + return s.prepareNew(ctx, message, requestID) +} + +func (s *ChatService) prepareExisting(message, conversationID, requestID string) (ChatTurn, error) { + now := s.now() + s.mu.Lock() + s.evictExpiredLocked(now) + record, exists := s.conversations[conversationID] + if !exists { + s.mu.Unlock() + return nil, ErrChatConversationNotFound + } + if record.busy { + s.mu.Unlock() + return nil, ErrChatConversationBusy + } + record.busy = true + record.lastUsed = now + s.mu.Unlock() + + idempotencyKey, err := s.newID("turn_") + if err != nil || !validCorrelationID(idempotencyKey) { + s.release(conversationID, record, true) + return nil, fmt.Errorf("%w: generate turn ID", ErrChatInternal) + } + return &chatTurn{ + service: s, + record: record, + conversationID: conversationID, + message: message, + requestID: requestID, + idempotencyKey: idempotencyKey, + reused: true, + }, nil +} + +func (s *ChatService) prepareNew(ctx context.Context, message, requestID string) (ChatTurn, error) { + conversationID, err := s.reserveConversationID() + if err != nil { + return nil, err + } + releaseReservation := true + defer func() { + if releaseReservation { + s.mu.Lock() + delete(s.reserved, conversationID) + s.mu.Unlock() + } + }() + + createKey, err := s.newID("session_") + if err != nil || !validCorrelationID(createKey) { + return nil, fmt.Errorf("%w: generate session ID", ErrChatInternal) + } + session, err := s.agent.CreateSession(ctx, AgentCreateSessionRequest{ + ExternalSubjectID: s.externalSubjectID, + IdempotencyKey: createKey, + RequestID: requestID, + Metadata: map[string]any{"source": "fire-safety-ymd-chat-api", "api_version": "v1"}, + }) + if err != nil { + return nil, err + } + if !validOpaqueProviderID(session.ID) { + return nil, fmt.Errorf("%w: provider returned an invalid session", ErrChatUpstreamProtocol) + } + idempotencyKey, err := s.newID("turn_") + if err != nil || !validCorrelationID(idempotencyKey) { + return nil, fmt.Errorf("%w: generate turn ID", ErrChatInternal) + } + + record := &chatConversation{providerSessionID: session.ID, busy: true, lastUsed: s.now()} + s.mu.Lock() + delete(s.reserved, conversationID) + s.conversations[conversationID] = record + s.mu.Unlock() + releaseReservation = false + + return &chatTurn{ + service: s, + record: record, + conversationID: conversationID, + message: message, + requestID: requestID, + idempotencyKey: idempotencyKey, + }, nil +} + +func (s *ChatService) reserveConversationID() (string, error) { + for attempt := 0; attempt < 4; attempt++ { + conversationID, err := s.newID("conv_") + if err != nil || !conversationIDPattern.MatchString(conversationID) { + return "", fmt.Errorf("%w: generate conversation ID", ErrChatInternal) + } + now := s.now() + s.mu.Lock() + s.evictExpiredLocked(now) + if len(s.conversations)+len(s.reserved) >= s.maxSessions { + s.mu.Unlock() + return "", ErrChatCapacityReached + } + _, conversationExists := s.conversations[conversationID] + _, reservationExists := s.reserved[conversationID] + if !conversationExists && !reservationExists { + s.reserved[conversationID] = struct{}{} + s.mu.Unlock() + return conversationID, nil + } + s.mu.Unlock() + } + return "", fmt.Errorf("%w: conversation ID collision", ErrChatInternal) +} + +func (s *ChatService) evictExpiredLocked(now time.Time) { + for conversationID, record := range s.conversations { + if !record.busy && now.Sub(record.lastUsed) >= s.sessionTTL { + delete(s.conversations, conversationID) + } + } +} + +func (s *ChatService) release(conversationID string, record *chatConversation, keep bool) { + s.mu.Lock() + defer s.mu.Unlock() + current, exists := s.conversations[conversationID] + if !exists || current != record { + return + } + if !keep { + delete(s.conversations, conversationID) + return + } + record.busy = false + record.lastUsed = s.now() +} + +type chatTurn struct { + service *ChatService + record *chatConversation + conversationID string + message string + requestID string + idempotencyKey string + reused bool + + streamMu sync.Mutex + streamed bool + closed bool + finish sync.Once +} + +func (t *chatTurn) ConversationID() string { return t.conversationID } + +func (t *chatTurn) Reused() bool { return t.reused } + +func (t *chatTurn) Stream(ctx context.Context, traceHandler func(AgentTraceEvent)) (ChatResult, error) { + t.streamMu.Lock() + if t.streamed || t.closed { + t.streamMu.Unlock() + return ChatResult{}, fmt.Errorf("%w: turn is no longer available", ErrChatInternal) + } + t.streamed = true + t.streamMu.Unlock() + + result, err := t.service.agent.StreamMessage(ctx, AgentMessageRequest{ + SessionID: t.record.providerSessionID, + Message: t.message, + IdempotencyKey: t.idempotencyKey, + RequestID: t.requestID, + Metadata: map[string]any{"source": "fire-safety-ymd-chat-api", "api_version": "v1"}, + }, traceHandler) + if err != nil { + t.complete(false) + return ChatResult{}, err + } + if strings.TrimSpace(result.Answer) == "" || result.Usage.Input < 0 || result.Usage.Output < 0 || result.Usage.Total < 0 { + t.complete(false) + return ChatResult{}, fmt.Errorf("%w: provider returned an invalid final result", ErrChatUpstreamProtocol) + } + t.complete(true) + return ChatResult{ + ConversationID: t.conversationID, + RunID: result.RunID, + ModelID: result.ModelID, + Answer: result.Answer, + Usage: result.Usage, + }, nil +} + +func (t *chatTurn) Close() { + t.streamMu.Lock() + if t.streamed { + t.closed = true + t.streamMu.Unlock() + return + } + t.closed = true + keep := t.reused + t.streamMu.Unlock() + t.complete(keep) +} + +func (t *chatTurn) complete(keep bool) { + t.finish.Do(func() { + t.service.release(t.conversationID, t.record, keep) + }) +} + +func randomChatID(prefix string) (string, error) { + value := make([]byte, 24) + if _, err := rand.Read(value); err != nil { + return "", err + } + return prefix + base64.RawURLEncoding.EncodeToString(value), nil +} + +func validCorrelationID(value string) bool { + if value == "" || len(value) > 512 { + return false + } + for _, character := range value { + if !(character >= 'a' && character <= 'z') && + !(character >= 'A' && character <= 'Z') && + !(character >= '0' && character <= '9') && + !strings.ContainsRune("._:-", character) { + return false + } + } + return true +} + +func validOpaqueProviderID(value string) bool { + return strings.TrimSpace(value) != "" && len(value) <= 512 && !containsControl(value) +} + +func containsControl(value string) bool { + for _, character := range value { + if unicode.IsControl(character) { + return true + } + } + return false +} diff --git a/internal/service/chat_test.go b/internal/service/chat_test.go new file mode 100644 index 0000000..83aa4c6 --- /dev/null +++ b/internal/service/chat_test.go @@ -0,0 +1,278 @@ +package service + +import ( + "context" + "errors" + "fmt" + "strings" + "sync" + "testing" + "time" +) + +func TestChatServiceCreatesAndReusesProviderSession(t *testing.T) { + agent := &fakeChatAgent{} + chat := newTestChatService(t, agent, nil, 10) + + first, err := chat.Prepare(context.Background(), ChatRequest{Message: " 第一问 ", RequestID: "request-1"}) + if err != nil { + t.Fatalf("Prepare(first) error = %v", err) + } + if first.Reused() || !conversationIDPattern.MatchString(first.ConversationID()) { + t.Fatalf("first turn = id %q reused %t", first.ConversationID(), first.Reused()) + } + firstResult, err := first.Stream(context.Background(), func(event AgentTraceEvent) { + if event.Event != "tool.started" || event.ToolName != "fire_safety_test" { + t.Errorf("trace event = %#v", event) + } + }) + if err != nil { + t.Fatalf("Stream(first) error = %v", err) + } + if firstResult.Answer != "answer-1" || firstResult.ConversationID != first.ConversationID() { + t.Fatalf("first result = %#v", firstResult) + } + if firstResult.ModelID != "test-model" { + t.Fatalf("first model ID = %q", firstResult.ModelID) + } + + second, err := chat.Prepare(context.Background(), ChatRequest{ + Message: "第二问", + ConversationID: first.ConversationID(), + RequestID: "request-2", + }) + if err != nil { + t.Fatalf("Prepare(second) error = %v", err) + } + if !second.Reused() { + t.Fatal("second.Reused() = false") + } + if _, err := second.Stream(context.Background(), nil); err != nil { + t.Fatalf("Stream(second) error = %v", err) + } + + agent.mu.Lock() + defer agent.mu.Unlock() + if len(agent.creates) != 1 || len(agent.messages) != 2 { + t.Fatalf("provider calls: creates=%d messages=%d", len(agent.creates), len(agent.messages)) + } + if agent.creates[0].ExternalSubjectID != "chat-test-subject" { + t.Fatalf("ExternalSubjectID = %q", agent.creates[0].ExternalSubjectID) + } + if agent.messages[0].SessionID != agent.messages[1].SessionID { + t.Fatalf("provider sessions differ: %q and %q", agent.messages[0].SessionID, agent.messages[1].SessionID) + } + if agent.messages[0].Message != "第一问" || agent.messages[1].Message != "第二问" { + t.Fatalf("provider messages = %q, %q", agent.messages[0].Message, agent.messages[1].Message) + } + if agent.messages[0].IdempotencyKey == agent.messages[1].IdempotencyKey { + t.Fatal("turn idempotency keys were reused") + } +} + +func TestChatServiceRejectsConcurrentRunOnSameConversation(t *testing.T) { + block := make(chan struct{}) + started := make(chan struct{}) + agent := &fakeChatAgent{block: block, started: started} + chat := newTestChatService(t, agent, nil, 10) + turn, err := chat.Prepare(context.Background(), ChatRequest{Message: "first", RequestID: "request-1"}) + if err != nil { + t.Fatalf("Prepare() error = %v", err) + } + + done := make(chan error, 1) + go func() { + _, streamErr := turn.Stream(context.Background(), nil) + done <- streamErr + }() + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("provider stream did not start") + } + _, err = chat.Prepare(context.Background(), ChatRequest{ + Message: "overlap", + ConversationID: turn.ConversationID(), + RequestID: "request-2", + }) + if !errors.Is(err, ErrChatConversationBusy) { + t.Fatalf("Prepare(overlap) error = %v, want busy", err) + } + close(block) + if err := <-done; err != nil { + t.Fatalf("Stream() error = %v", err) + } + + next, err := chat.Prepare(context.Background(), ChatRequest{ + Message: "after completion", + ConversationID: turn.ConversationID(), + RequestID: "request-3", + }) + if err != nil { + t.Fatalf("Prepare(after completion) error = %v", err) + } + next.Close() +} + +func TestChatServiceInvalidatesConversationAfterUncertainStreamFailure(t *testing.T) { + agent := &fakeChatAgent{streamErr: ErrChatUpstreamProtocol} + chat := newTestChatService(t, agent, nil, 10) + turn, err := chat.Prepare(context.Background(), ChatRequest{Message: "first", RequestID: "request-1"}) + if err != nil { + t.Fatalf("Prepare() error = %v", err) + } + if _, err := turn.Stream(context.Background(), nil); !errors.Is(err, ErrChatUpstreamProtocol) { + t.Fatalf("Stream() error = %v", err) + } + _, err = chat.Prepare(context.Background(), ChatRequest{ + Message: "retry", + ConversationID: turn.ConversationID(), + RequestID: "request-2", + }) + if !errors.Is(err, ErrChatConversationNotFound) { + t.Fatalf("Prepare(retry) error = %v, want not found", err) + } +} + +func TestChatServiceExpiresIdleSessionsAndEnforcesCapacity(t *testing.T) { + now := time.Date(2026, 9, 5, 12, 0, 0, 0, time.UTC) + clock := func() time.Time { return now } + agent := &fakeChatAgent{} + chat := newTestChatService(t, agent, clock, 1) + first, err := chat.Prepare(context.Background(), ChatRequest{Message: "first", RequestID: "request-1"}) + if err != nil { + t.Fatalf("Prepare(first) error = %v", err) + } + if _, err := first.Stream(context.Background(), nil); err != nil { + t.Fatalf("Stream(first) error = %v", err) + } + + if _, err := chat.Prepare(context.Background(), ChatRequest{Message: "new", RequestID: "request-2"}); !errors.Is(err, ErrChatCapacityReached) { + t.Fatalf("Prepare(at capacity) error = %v", err) + } + now = now.Add(31 * time.Minute) + second, err := chat.Prepare(context.Background(), ChatRequest{Message: "new", RequestID: "request-3"}) + if err != nil { + t.Fatalf("Prepare(after TTL) error = %v", err) + } + second.Close() + + _, err = chat.Prepare(context.Background(), ChatRequest{ + Message: "old", + ConversationID: first.ConversationID(), + RequestID: "request-4", + }) + if !errors.Is(err, ErrChatConversationNotFound) { + t.Fatalf("Prepare(expired) error = %v", err) + } +} + +func TestChatServiceReleasesReservationAfterCreateFailure(t *testing.T) { + agent := &fakeChatAgent{createErr: ErrChatUpstreamUnavailable} + chat := newTestChatService(t, agent, nil, 1) + if _, err := chat.Prepare(context.Background(), ChatRequest{Message: "first", RequestID: "request-1"}); !errors.Is(err, ErrChatUpstreamUnavailable) { + t.Fatalf("Prepare(first) error = %v", err) + } + agent.mu.Lock() + agent.createErr = nil + agent.mu.Unlock() + turn, err := chat.Prepare(context.Background(), ChatRequest{Message: "second", RequestID: "request-2"}) + if err != nil { + t.Fatalf("Prepare(second) error = %v", err) + } + turn.Close() +} + +func TestChatServiceRejectsInvalidRequests(t *testing.T) { + chat := newTestChatService(t, &fakeChatAgent{}, nil, 10) + tests := []ChatRequest{ + {Message: "", RequestID: "request-1"}, + {Message: strings.Repeat("a", 1025), RequestID: "request-1"}, + {Message: "ok", RequestID: "invalid request id"}, + {Message: "ok", RequestID: "request-1", ConversationID: "provider-session-1"}, + } + for _, request := range tests { + if _, err := chat.Prepare(context.Background(), request); !errors.Is(err, ErrChatInvalidArgument) { + t.Fatalf("Prepare(%#v) error = %v, want invalid argument", request, err) + } + } +} + +func newTestChatService(t *testing.T, agent ChatAgent, now func() time.Time, maxSessions int) *ChatService { + t.Helper() + var idMu sync.Mutex + idSequence := 0 + newID := func(prefix string) (string, error) { + idMu.Lock() + defer idMu.Unlock() + idSequence++ + return fmt.Sprintf("%s%024d", prefix, idSequence), nil + } + chat, err := NewChatService(agent, ChatOptions{ + ExternalSubjectID: "chat-test-subject", + MaxMessageBytes: 1024, + SessionTTL: 30 * time.Minute, + MaxSessions: maxSessions, + now: now, + newID: newID, + }) + if err != nil { + t.Fatalf("NewChatService() error = %v", err) + } + return chat +} + +type fakeChatAgent struct { + mu sync.Mutex + creates []AgentCreateSessionRequest + messages []AgentMessageRequest + createErr error + streamErr error + block <-chan struct{} + started chan<- struct{} + startOnce sync.Once + sessionSeq int +} + +func (a *fakeChatAgent) CreateSession(_ context.Context, request AgentCreateSessionRequest) (AgentSession, error) { + a.mu.Lock() + defer a.mu.Unlock() + a.creates = append(a.creates, request) + if a.createErr != nil { + return AgentSession{}, a.createErr + } + a.sessionSeq++ + return AgentSession{ID: fmt.Sprintf("provider-session-%d", a.sessionSeq)}, nil +} + +func (a *fakeChatAgent) StreamMessage(ctx context.Context, request AgentMessageRequest, trace func(AgentTraceEvent)) (AgentMessageResult, error) { + a.mu.Lock() + a.messages = append(a.messages, request) + messageNumber := len(a.messages) + streamErr := a.streamErr + block := a.block + started := a.started + a.mu.Unlock() + if started != nil { + a.startOnce.Do(func() { close(started) }) + } + if block != nil { + select { + case <-ctx.Done(): + return AgentMessageResult{}, ctx.Err() + case <-block: + } + } + if trace != nil { + trace(AgentTraceEvent{Event: "tool.started", ToolName: "fire_safety_test", Status: "running"}) + } + if streamErr != nil { + return AgentMessageResult{}, streamErr + } + return AgentMessageResult{ + RunID: fmt.Sprintf("run-%d", messageNumber), + ModelID: "test-model", + Answer: fmt.Sprintf("answer-%d", messageNumber), + Usage: ChatTokenUsage{Input: 1, Output: 2, Total: 3}, + }, nil +} diff --git a/internal/service/doc.go b/internal/service/doc.go new file mode 100644 index 0000000..58cbd9d --- /dev/null +++ b/internal/service/doc.go @@ -0,0 +1,2 @@ +// Package service orchestrates application use cases over domain rules and external ports. +package service diff --git a/internal/service/spatial.go b/internal/service/spatial.go new file mode 100644 index 0000000..998d530 --- /dev/null +++ b/internal/service/spatial.go @@ -0,0 +1,334 @@ +package service + +import ( + "context" + "errors" + "fmt" + "math" + "slices" + "strings" + "time" + "unicode" + + "fire-safety-ymd/internal/domain" +) + +const ( + statusOK = "ok" + statusNoResults = "no_results" + allTownsScopeWarning = "results_include_all_towns_in_configured_database" + townAllowlistScopeWarning = "results_limited_to_server_authorized_towns" + invalidGeometryWarning = "source_records_with_invalid_geometries_are_excluded" +) + +var ( + // ErrInvalidArgument indicates a caller-controlled spatial query is invalid. + ErrInvalidArgument = errors.New("invalid spatial query argument") + // ErrDataSourceUnavailable indicates a repository query could not complete safely. + ErrDataSourceUnavailable = errors.New("spatial data source unavailable") + // ErrQueryTimeout indicates the bounded repository query timed out. + ErrQueryTimeout = errors.New("spatial query timeout") +) + +// SpatialRepository is the fixed read-only persistence port used by spatial use cases. +type SpatialRepository interface { + SearchPlaceCandidates(context.Context, domain.PlaceSearchQuery) ([]domain.PlaceCandidate, error) + ResolveIncidentContext(context.Context, domain.Coordinate, domain.SpatialScope) ([]domain.IncidentContext, error) + FindNearbyWaterSources(context.Context, domain.NearbyQuery) ([]domain.WaterSource, error) + FindCommandPostCandidates(context.Context, domain.NearbyQuery) ([]domain.CommandPostCandidate, error) + ListNearbyAccessLines(context.Context, domain.NearbyQuery) ([]domain.AccessLine, error) + GetResponsibleUnits(context.Context, domain.Coordinate, domain.SpatialScope) ([]domain.ResponsibleUnit, error) + FindNearbyRiskAreas(context.Context, domain.NearbyQuery) ([]domain.RiskArea, error) +} + +// ResultMetadata records the provenance and bounds of a tool result. +type ResultMetadata struct { + GeneratedAt string `json:"generated_at"` + DataSources []string `json:"data_sources"` + SpatialReference string `json:"spatial_reference"` + ResultCount int `json:"result_count"` + SearchRadiusMeters *float64 `json:"search_radius_meters,omitempty"` +} + +// QueryResult is the common structured MCP result envelope. +type QueryResult[T any] struct { + Status string `json:"status"` + Data T `json:"data"` + Metadata ResultMetadata `json:"metadata"` + Warnings []string `json:"warnings"` +} + +// SpatialService executes bounded read-only forest-fire spatial queries. +type SpatialService struct { + repository SpatialRepository + scope domain.SpatialScope + queryTimeout time.Duration + now func() time.Time +} + +// NewSpatialService creates a spatial service with a trusted server-side data scope. +func NewSpatialService(repository SpatialRepository, scope domain.SpatialScope, queryTimeout time.Duration) (*SpatialService, error) { + if repository == nil { + return nil, errors.New("spatial repository is required") + } + if scope.AllTowns && len(scope.AllowedTowns) > 0 { + return nil, errors.New("all-towns scope must not include an allowlist") + } + if !scope.AllTowns && len(scope.AllowedTowns) == 0 { + return nil, errors.New("at least one trusted town is required") + } + if queryTimeout <= 0 { + return nil, errors.New("query timeout must be greater than zero") + } + return &SpatialService{ + repository: repository, + scope: domain.SpatialScope{ + AllTowns: scope.AllTowns, + AllowedTowns: slices.Clone(scope.AllowedTowns), + }, + queryTimeout: queryTimeout, + now: time.Now, + }, nil +} + +// SearchPlaceCandidates searches existing fire-safety records without treating a match as a confirmed incident point. +func (s *SpatialService) SearchPlaceCandidates(ctx context.Context, placeName string, limit int) (QueryResult[[]domain.PlaceCandidate], error) { + placeName = strings.TrimSpace(placeName) + if count := len([]rune(placeName)); count < 2 || count > 100 || containsControlCharacter(placeName) { + return QueryResult[[]domain.PlaceCandidate]{}, fmt.Errorf("%w: place_name must contain 2 to 100 characters without controls", ErrInvalidArgument) + } + if limit == 0 { + limit = 10 + } + if limit < 1 || limit > 20 { + return QueryResult[[]domain.PlaceCandidate]{}, fmt.Errorf("%w: limit must be between 1 and 20", ErrInvalidArgument) + } + + queryCtx, cancel := context.WithTimeout(ctx, s.queryTimeout) + defer cancel() + items, err := s.repository.SearchPlaceCandidates(queryCtx, domain.PlaceSearchQuery{ + PlaceName: placeName, + Limit: limit, + Scope: s.scope, + }) + if err != nil { + return QueryResult[[]domain.PlaceCandidate]{}, classifyRepositoryError("search place candidates", err) + } + return result(items, []string{ + "water_source", + "storage_pool", + "fire_access_line", + "fire_check_station", + "fire_lookout", + "fire_grid", + "forest_enterprise", + "cemetery_area", + }, nil, s.now, s.scopeWarning(), []string{ + "place_search_is_limited_to_existing_fire_safety_records", + "place_candidates_require_user_confirmation", + "representative_points_are_not_exact_incident_locations", + }), nil +} + +// ResolveIncidentContext returns fire grids covering a WGS84 point. +func (s *SpatialService) ResolveIncidentContext(ctx context.Context, point domain.Coordinate) (QueryResult[[]domain.IncidentContext], error) { + if err := validateCoordinate(point); err != nil { + return QueryResult[[]domain.IncidentContext]{}, err + } + queryCtx, cancel := context.WithTimeout(ctx, s.queryTimeout) + defer cancel() + + items, err := s.repository.ResolveIncidentContext(queryCtx, point, s.scope) + if err != nil { + return QueryResult[[]domain.IncidentContext]{}, classifyRepositoryError("resolve incident context", err) + } + return result(items, []string{"fire_grid"}, nil, s.now, s.scopeWarning(), []string{ + "grid_records_may_be_incomplete_or_stale", + }), nil +} + +// FindNearbyWaterSources returns bounded water-source candidates by distance. +func (s *SpatialService) FindNearbyWaterSources(ctx context.Context, point domain.Coordinate, radiusMeters float64, limit int) (QueryResult[[]domain.WaterSource], error) { + query, err := s.nearbyQuery(point, radiusMeters, 10_000, 30_000, limit) + if err != nil { + return QueryResult[[]domain.WaterSource]{}, err + } + queryCtx, cancel := context.WithTimeout(ctx, s.queryTimeout) + defer cancel() + + items, err := s.repository.FindNearbyWaterSources(queryCtx, query) + if err != nil { + return QueryResult[[]domain.WaterSource]{}, classifyRepositoryError("find nearby water sources", err) + } + return result(items, []string{ + "water_source", + "storage_pool", + }, &query.RadiusMeters, s.now, s.scopeWarning(), []string{ + "resource_availability_not_verified", + "distance_is_spatial_proximity_not_operational_accessibility", + "source_timestamp_raw_unit_and_timezone_are_unconfirmed", + }), nil +} + +// FindCommandPostCandidates returns existing nearby facilities that require field assessment. +func (s *SpatialService) FindCommandPostCandidates(ctx context.Context, point domain.Coordinate, radiusMeters float64, limit int) (QueryResult[[]domain.CommandPostCandidate], error) { + query, err := s.nearbyQuery(point, radiusMeters, 10_000, 20_000, limit) + if err != nil { + return QueryResult[[]domain.CommandPostCandidate]{}, err + } + queryCtx, cancel := context.WithTimeout(ctx, s.queryTimeout) + defer cancel() + + items, err := s.repository.FindCommandPostCandidates(queryCtx, query) + if err != nil { + return QueryResult[[]domain.CommandPostCandidate]{}, classifyRepositoryError("find command post candidates", err) + } + return result(items, []string{ + "fire_check_station", + "fire_lookout", + }, &query.RadiusMeters, s.now, s.scopeWarning(), []string{ + "candidate_only_requires_field_safety_communications_capacity_and_access_review", + "facility_availability_not_verified", + }), nil +} + +// ListNearbyAccessLines returns nearby mapped fire-access lines, not calculated routes. +func (s *SpatialService) ListNearbyAccessLines(ctx context.Context, point domain.Coordinate, radiusMeters float64, limit int) (QueryResult[[]domain.AccessLine], error) { + query, err := s.nearbyQuery(point, radiusMeters, 5_000, 10_000, limit) + if err != nil { + return QueryResult[[]domain.AccessLine]{}, err + } + queryCtx, cancel := context.WithTimeout(ctx, s.queryTimeout) + defer cancel() + + items, err := s.repository.ListNearbyAccessLines(queryCtx, query) + if err != nil { + return QueryResult[[]domain.AccessLine]{}, classifyRepositoryError("list nearby access lines", err) + } + return result(items, []string{"fire_access_line"}, &query.RadiusMeters, s.now, s.scopeWarning(), []string{ + "candidate_lines_only_not_a_route_plan", + "passability_surface_width_slope_vehicle_limits_and_closures_are_unknown", + "source_updated_raw_has_no_confirmed_timezone", + }), nil +} + +// GetResponsibleUnits returns non-personal team responsibility recorded on covering grids. +func (s *SpatialService) GetResponsibleUnits(ctx context.Context, point domain.Coordinate) (QueryResult[[]domain.ResponsibleUnit], error) { + if err := validateCoordinate(point); err != nil { + return QueryResult[[]domain.ResponsibleUnit]{}, err + } + queryCtx, cancel := context.WithTimeout(ctx, s.queryTimeout) + defer cancel() + + items, err := s.repository.GetResponsibleUnits(queryCtx, point, s.scope) + if err != nil { + return QueryResult[[]domain.ResponsibleUnit]{}, classifyRepositoryError("get responsible units", err) + } + return result(items, []string{"fire_grid"}, nil, s.now, s.scopeWarning(), []string{ + "source_does_not_store_live_team_positions", + "source_does_not_store_confirmed_assembly_sites_or_readiness", + }), nil +} + +// FindNearbyRiskAreas returns nearby cemetery and forest-enterprise polygons. +func (s *SpatialService) FindNearbyRiskAreas(ctx context.Context, point domain.Coordinate, radiusMeters float64, limit int) (QueryResult[[]domain.RiskArea], error) { + query, err := s.nearbyQuery(point, radiusMeters, 3_000, 10_000, limit) + if err != nil { + return QueryResult[[]domain.RiskArea]{}, err + } + queryCtx, cancel := context.WithTimeout(ctx, s.queryTimeout) + defer cancel() + + items, err := s.repository.FindNearbyRiskAreas(queryCtx, query) + if err != nil { + return QueryResult[[]domain.RiskArea]{}, classifyRepositoryError("find nearby risk areas", err) + } + return result(items, []string{ + "cemetery_area", + "forest_enterprise", + }, &query.RadiusMeters, s.now, s.scopeWarning(), []string{ + "risk_records_may_be_incomplete_or_stale", + "proximity_does_not_establish_current_hazard_severity", + }), nil +} + +func (s *SpatialService) nearbyQuery(point domain.Coordinate, radiusMeters, defaultRadius, maximumRadius float64, limit int) (domain.NearbyQuery, error) { + if err := validateCoordinate(point); err != nil { + return domain.NearbyQuery{}, err + } + if radiusMeters == 0 { + radiusMeters = defaultRadius + } + if radiusMeters < 100 || radiusMeters > maximumRadius || math.IsNaN(radiusMeters) || math.IsInf(radiusMeters, 0) { + return domain.NearbyQuery{}, fmt.Errorf("%w: radius_meters must be between 100 and %.0f", ErrInvalidArgument, maximumRadius) + } + if limit == 0 { + limit = 10 + } + if limit < 1 || limit > 20 { + return domain.NearbyQuery{}, fmt.Errorf("%w: limit must be between 1 and 20", ErrInvalidArgument) + } + return domain.NearbyQuery{ + Point: point, + RadiusMeters: radiusMeters, + Limit: limit, + Scope: s.scope, + }, nil +} + +func validateCoordinate(point domain.Coordinate) error { + if math.IsNaN(point.Longitude) || math.IsInf(point.Longitude, 0) || point.Longitude < -180 || point.Longitude > 180 { + return fmt.Errorf("%w: longitude must be between -180 and 180", ErrInvalidArgument) + } + if math.IsNaN(point.Latitude) || math.IsInf(point.Latitude, 0) || point.Latitude < -90 || point.Latitude > 90 { + return fmt.Errorf("%w: latitude must be between -90 and 90", ErrInvalidArgument) + } + return nil +} + +func containsControlCharacter(value string) bool { + for _, character := range value { + if unicode.IsControl(character) { + return true + } + } + return false +} + +func classifyRepositoryError(action string, err error) error { + if errors.Is(err, context.DeadlineExceeded) || errors.Is(err, context.Canceled) { + return fmt.Errorf("%w: %s", ErrQueryTimeout, action) + } + return fmt.Errorf("%w: %s: %w", ErrDataSourceUnavailable, action, err) +} + +func (s *SpatialService) scopeWarning() string { + if s.scope.AllTowns { + return allTownsScopeWarning + } + return townAllowlistScopeWarning +} + +func result[T any](items []T, dataSources []string, radius *float64, now func() time.Time, scopeWarning string, warnings []string) QueryResult[[]T] { + if items == nil { + items = []T{} + } + status := statusOK + if len(items) == 0 { + status = statusNoResults + } + resultWarnings := []string{scopeWarning, invalidGeometryWarning} + resultWarnings = append(resultWarnings, warnings...) + return QueryResult[[]T]{ + Status: status, + Data: items, + Metadata: ResultMetadata{ + GeneratedAt: now().UTC().Format(time.RFC3339Nano), + DataSources: slices.Clone(dataSources), + SpatialReference: "EPSG:4326", + ResultCount: len(items), + SearchRadiusMeters: radius, + }, + Warnings: resultWarnings, + } +} diff --git a/internal/service/spatial_test.go b/internal/service/spatial_test.go new file mode 100644 index 0000000..612d6d5 --- /dev/null +++ b/internal/service/spatial_test.go @@ -0,0 +1,334 @@ +package service + +import ( + "context" + "errors" + "strings" + "testing" + "time" + + "fire-safety-ymd/internal/domain" +) + +func TestSpatialServiceSearchesUserConfirmablePlaceCandidates(t *testing.T) { + var captured domain.PlaceSearchQuery + repository := &fakeSpatialRepository{ + places: func(_ context.Context, query domain.PlaceSearchQuery) ([]domain.PlaceCandidate, error) { + captured = query + return []domain.PlaceCandidate{{ + SourceRecordID: "station-1", + PlaceType: "fire_check_station", + Name: "观水检查站", + MatchedField: "name", + MatchedText: "观水检查站", + MatchKind: "partial", + LocationKind: "recorded_point", + }}, nil + }, + } + spatial, err := NewSpatialService(repository, domain.SpatialScope{AllTowns: true}, time.Second) + if err != nil { + t.Fatalf("NewSpatialService() error = %v", err) + } + + result, err := spatial.SearchPlaceCandidates(context.Background(), " 观水 ", 0) + if err != nil { + t.Fatalf("SearchPlaceCandidates() error = %v", err) + } + if captured.PlaceName != "观水" || captured.Limit != 10 || !captured.Scope.AllTowns { + t.Fatalf("captured query = %#v", captured) + } + if result.Status != "ok" || result.Metadata.ResultCount != 1 || len(result.Data) != 1 { + t.Fatalf("result = %#v", result) + } + for _, warning := range []string{ + invalidGeometryWarning, + "place_search_is_limited_to_existing_fire_safety_records", + "place_candidates_require_user_confirmation", + "representative_points_are_not_exact_incident_locations", + } { + if !containsString(result.Warnings, warning) { + t.Fatalf("warnings = %#v, want %q", result.Warnings, warning) + } + } +} + +func TestSpatialServiceRejectsInvalidPlaceSearch(t *testing.T) { + spatial, err := NewSpatialService(&fakeSpatialRepository{}, domain.SpatialScope{AllTowns: true}, time.Second) + if err != nil { + t.Fatalf("NewSpatialService() error = %v", err) + } + tests := []struct { + name string + placeName string + limit int + }{ + {name: "empty", placeName: " "}, + {name: "too short", placeName: "山"}, + {name: "too long", placeName: strings.Repeat("山", 101)}, + {name: "control character", placeName: "观水\n镇"}, + {name: "limit", placeName: "观水镇", limit: 21}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := spatial.SearchPlaceCandidates(context.Background(), tt.placeName, tt.limit) + if !errors.Is(err, ErrInvalidArgument) { + t.Fatalf("error = %v, want ErrInvalidArgument", err) + } + }) + } +} + +func TestSpatialServiceUsesTrustedScopeAndDefaults(t *testing.T) { + var captured domain.NearbyQuery + repository := &fakeSpatialRepository{ + water: func(_ context.Context, query domain.NearbyQuery) ([]domain.WaterSource, error) { + captured = query + return []domain.WaterSource{{SourceRecordID: "152", Category: "water_source"}}, nil + }, + } + towns := []string{"莒格庄镇"} + spatial, err := NewSpatialService(repository, domain.SpatialScope{AllowedTowns: towns}, time.Second) + if err != nil { + t.Fatalf("NewSpatialService() error = %v", err) + } + towns[0] = "被调用方篡改" + fixedNow := time.Date(2026, 9, 4, 8, 30, 0, 0, time.FixedZone("CST", 8*60*60)) + spatial.now = func() time.Time { return fixedNow } + + result, err := spatial.FindNearbyWaterSources(context.Background(), domain.Coordinate{Longitude: 121.75, Latitude: 37.19}, 0, 0) + if err != nil { + t.Fatalf("FindNearbyWaterSources() error = %v", err) + } + if captured.RadiusMeters != 10_000 || captured.Limit != 10 { + t.Fatalf("captured bounds = radius %.0f limit %d", captured.RadiusMeters, captured.Limit) + } + if captured.Scope.AllTowns || len(captured.Scope.AllowedTowns) != 1 || captured.Scope.AllowedTowns[0] != "莒格庄镇" { + t.Fatalf("trusted scope = %#v", captured.Scope) + } + if result.Status != "ok" || result.Metadata.ResultCount != 1 || result.Metadata.GeneratedAt != "2026-09-04T00:30:00Z" { + t.Fatalf("unexpected result metadata: %#v", result) + } + if result.Metadata.SearchRadiusMeters == nil || *result.Metadata.SearchRadiusMeters != 10_000 { + t.Fatalf("search radius metadata = %#v", result.Metadata.SearchRadiusMeters) + } + if got, want := result.Metadata.DataSources, []string{"water_source", "storage_pool"}; len(got) != len(want) || got[0] != want[0] || got[1] != want[1] { + t.Fatalf("data sources = %#v, want %#v", got, want) + } + if !containsString(result.Warnings, townAllowlistScopeWarning) { + t.Fatalf("warnings = %#v, want town allowlist scope warning", result.Warnings) + } + if !containsString(result.Warnings, invalidGeometryWarning) { + t.Fatalf("warnings = %#v, want invalid-geometry exclusion warning", result.Warnings) + } +} + +func TestSpatialServiceUsesAllTownsScope(t *testing.T) { + var captured domain.NearbyQuery + repository := &fakeSpatialRepository{ + water: func(_ context.Context, query domain.NearbyQuery) ([]domain.WaterSource, error) { + captured = query + return []domain.WaterSource{}, nil + }, + } + spatial, err := NewSpatialService(repository, domain.SpatialScope{AllTowns: true}, time.Second) + if err != nil { + t.Fatalf("NewSpatialService() error = %v", err) + } + + result, err := spatial.FindNearbyWaterSources(context.Background(), domain.Coordinate{Longitude: 121.75, Latitude: 37.19}, 100, 1) + if err != nil { + t.Fatalf("FindNearbyWaterSources() error = %v", err) + } + if !captured.Scope.AllTowns || len(captured.Scope.AllowedTowns) != 0 { + t.Fatalf("trusted scope = %#v, want all towns", captured.Scope) + } + if !containsString(result.Warnings, allTownsScopeWarning) || containsString(result.Warnings, townAllowlistScopeWarning) { + t.Fatalf("warnings = %#v, want only all-towns scope warning", result.Warnings) + } +} + +func TestNewSpatialServiceRejectsInvalidScope(t *testing.T) { + tests := []struct { + name string + scope domain.SpatialScope + }{ + {name: "empty allowlist", scope: domain.SpatialScope{}}, + {name: "all with allowlist", scope: domain.SpatialScope{AllTowns: true, AllowedTowns: []string{"莒格庄镇"}}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if _, err := NewSpatialService(&fakeSpatialRepository{}, tt.scope, time.Second); err == nil { + t.Fatal("NewSpatialService() error = nil") + } + }) + } +} + +func TestSpatialServiceReturnsExplicitNoResults(t *testing.T) { + spatial, err := NewSpatialService(&fakeSpatialRepository{}, domain.SpatialScope{AllowedTowns: []string{"高陵镇"}}, time.Second) + if err != nil { + t.Fatalf("NewSpatialService() error = %v", err) + } + + result, err := spatial.ResolveIncidentContext(context.Background(), domain.Coordinate{Longitude: 121, Latitude: 37}) + if err != nil { + t.Fatalf("ResolveIncidentContext() error = %v", err) + } + if result.Status != "no_results" || result.Data == nil || len(result.Data) != 0 || result.Metadata.ResultCount != 0 { + t.Fatalf("unexpected no-results response: %#v", result) + } + if len(result.Warnings) == 0 { + t.Fatal("no-results response must retain data-quality warnings") + } +} + +func TestSpatialServiceRejectsUnboundedInputs(t *testing.T) { + repository := &fakeSpatialRepository{} + spatial, err := NewSpatialService(repository, domain.SpatialScope{AllowedTowns: []string{"高陵镇"}}, time.Second) + if err != nil { + t.Fatalf("NewSpatialService() error = %v", err) + } + + tests := []struct { + name string + call func() error + }{ + { + name: "longitude", + call: func() error { + _, err := spatial.ResolveIncidentContext(context.Background(), domain.Coordinate{Longitude: 181, Latitude: 37}) + return err + }, + }, + { + name: "water radius", + call: func() error { + _, err := spatial.FindNearbyWaterSources(context.Background(), domain.Coordinate{Longitude: 121, Latitude: 37}, 30_001, 1) + return err + }, + }, + { + name: "access radius", + call: func() error { + _, err := spatial.ListNearbyAccessLines(context.Background(), domain.Coordinate{Longitude: 121, Latitude: 37}, 10_001, 1) + return err + }, + }, + { + name: "limit", + call: func() error { + _, err := spatial.FindNearbyRiskAreas(context.Background(), domain.Coordinate{Longitude: 121, Latitude: 37}, 100, 21) + return err + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if err := tt.call(); !errors.Is(err, ErrInvalidArgument) { + t.Fatalf("error = %v, want ErrInvalidArgument", err) + } + }) + } +} + +func TestSpatialServiceClassifiesRepositoryFailures(t *testing.T) { + t.Run("data source", func(t *testing.T) { + spatial, err := NewSpatialService(&fakeSpatialRepository{ + water: func(context.Context, domain.NearbyQuery) ([]domain.WaterSource, error) { + return nil, errors.New("driver detail must stay internal") + }, + }, domain.SpatialScope{AllowedTowns: []string{"高陵镇"}}, time.Second) + if err != nil { + t.Fatalf("NewSpatialService() error = %v", err) + } + _, err = spatial.FindNearbyWaterSources(context.Background(), domain.Coordinate{Longitude: 121, Latitude: 37}, 100, 1) + if !errors.Is(err, ErrDataSourceUnavailable) { + t.Fatalf("error = %v, want ErrDataSourceUnavailable", err) + } + }) + + t.Run("timeout", func(t *testing.T) { + spatial, err := NewSpatialService(&fakeSpatialRepository{ + water: func(ctx context.Context, _ domain.NearbyQuery) ([]domain.WaterSource, error) { + <-ctx.Done() + return nil, ctx.Err() + }, + }, domain.SpatialScope{AllowedTowns: []string{"高陵镇"}}, 5*time.Millisecond) + if err != nil { + t.Fatalf("NewSpatialService() error = %v", err) + } + _, err = spatial.FindNearbyWaterSources(context.Background(), domain.Coordinate{Longitude: 121, Latitude: 37}, 100, 1) + if !errors.Is(err, ErrQueryTimeout) { + t.Fatalf("error = %v, want ErrQueryTimeout", err) + } + }) +} + +func containsString(values []string, target string) bool { + for _, value := range values { + if value == target { + return true + } + } + return false +} + +type fakeSpatialRepository struct { + places func(context.Context, domain.PlaceSearchQuery) ([]domain.PlaceCandidate, error) + resolve func(context.Context, domain.Coordinate, domain.SpatialScope) ([]domain.IncidentContext, error) + water func(context.Context, domain.NearbyQuery) ([]domain.WaterSource, error) + commandPost func(context.Context, domain.NearbyQuery) ([]domain.CommandPostCandidate, error) + access func(context.Context, domain.NearbyQuery) ([]domain.AccessLine, error) + units func(context.Context, domain.Coordinate, domain.SpatialScope) ([]domain.ResponsibleUnit, error) + risks func(context.Context, domain.NearbyQuery) ([]domain.RiskArea, error) +} + +func (f *fakeSpatialRepository) SearchPlaceCandidates(ctx context.Context, query domain.PlaceSearchQuery) ([]domain.PlaceCandidate, error) { + if f.places == nil { + return []domain.PlaceCandidate{}, nil + } + return f.places(ctx, query) +} + +func (f *fakeSpatialRepository) ResolveIncidentContext(ctx context.Context, point domain.Coordinate, scope domain.SpatialScope) ([]domain.IncidentContext, error) { + if f.resolve == nil { + return []domain.IncidentContext{}, nil + } + return f.resolve(ctx, point, scope) +} + +func (f *fakeSpatialRepository) FindNearbyWaterSources(ctx context.Context, query domain.NearbyQuery) ([]domain.WaterSource, error) { + if f.water == nil { + return []domain.WaterSource{}, nil + } + return f.water(ctx, query) +} + +func (f *fakeSpatialRepository) FindCommandPostCandidates(ctx context.Context, query domain.NearbyQuery) ([]domain.CommandPostCandidate, error) { + if f.commandPost == nil { + return []domain.CommandPostCandidate{}, nil + } + return f.commandPost(ctx, query) +} + +func (f *fakeSpatialRepository) ListNearbyAccessLines(ctx context.Context, query domain.NearbyQuery) ([]domain.AccessLine, error) { + if f.access == nil { + return []domain.AccessLine{}, nil + } + return f.access(ctx, query) +} + +func (f *fakeSpatialRepository) GetResponsibleUnits(ctx context.Context, point domain.Coordinate, scope domain.SpatialScope) ([]domain.ResponsibleUnit, error) { + if f.units == nil { + return []domain.ResponsibleUnit{}, nil + } + return f.units(ctx, point, scope) +} + +func (f *fakeSpatialRepository) FindNearbyRiskAreas(ctx context.Context, query domain.NearbyQuery) ([]domain.RiskArea, error) { + if f.risks == nil { + return []domain.RiskArea{}, nil + } + return f.risks(ctx, query) +} diff --git a/pkg/README.md b/pkg/README.md new file mode 100644 index 0000000..c96724e --- /dev/null +++ b/pkg/README.md @@ -0,0 +1,5 @@ +# pkg + +该目录仅用于未来确实需要被外部 Go module 复用的公共包。 + +当前没有这样的稳定 API。项目内部代码应优先放在 `internal/`,避免过早形成难以演进的公共契约。