Compare commits

..

18 Commits

Author SHA1 Message Date
rayd1o
1cb51b1172 release: bump version to 0.51.0 2026-05-11 09:49:08 +08:00
rayd1o
455b8360d0 release: bump version to 0.50.0 2026-05-10 22:06:01 +08:00
linkong
e1984c7a35 release: bump version to 0.49.0
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-05-08 17:42:27 +08:00
linkong
bb9183b8a4 release: bump version to 0.48.0 2026-05-07 18:06:06 +08:00
linkong
421234301a release: bump version to 0.47.0 2026-04-30 16:56:37 +08:00
linkong
f22079d33a release: bump version to 0.46.3 2026-04-30 14:46:19 +08:00
linkong
9f737fdb89 release: bump version to 0.46.2 2026-04-30 14:30:12 +08:00
linkong
7418ce2fc1 release: bump version to 0.46.1 2026-04-30 09:41:08 +08:00
rayd1o
b1a5934b80 release: bump version to 0.46.0 2026-04-30 04:42:29 +08:00
rayd1o
ba54545ac7 release: bump version to 0.45.0 2026-04-29 23:43:54 +08:00
linkong
9dafbf4f6e release: bump version to 0.44.2 2026-04-29 18:11:37 +08:00
linkong
a87537e903 release: bump version to 0.44.1 2026-04-29 18:07:35 +08:00
linkong
87594a95ff release: bump version to 0.44.0 2026-04-29 17:27:44 +08:00
linkong
2da25376bd release: bump version to 0.43.1 2026-04-28 16:21:33 +08:00
linkong
ac69d5d354 release: bump version to 0.43.0 2026-04-28 16:10:17 +08:00
rayd1o
1cd2dab0ee release: bump version to 0.42.2 2026-04-28 04:35:13 +08:00
rayd1o
42d019af36 release: bump version to 0.42.1 2026-04-28 04:29:44 +08:00
rayd1o
b4e8afb272 release: bump version to 0.42.0 2026-04-28 04:27:18 +08:00
263 changed files with 51502 additions and 4068 deletions

View File

@@ -12,6 +12,19 @@ allowed-tools: ["Read", "Edit", "Bash", "Grep", "Glob"]
`$ARGUMENTS` 非空,则只检查指定文件/目录;否则检查所有未提交修改(`git diff HEAD`)。
## 节省上下文规则
优先用确定性的 CLI 检查缩小范围,不要一上来把完整文件或大 diff 读入上下文:
```bash
git diff --name-only HEAD
git diff --unified=0 HEAD -- <path>
git diff --check
rg -n "TODO|FIXME|console\.log|debugger|print\(" <changed-paths>
```
只有 focused diff 不足以安全判断或修改时,才读取完整文件。
## 审查清单
按优先级检查以下问题(只报告在本次 diff 中**新增或修改**的代码里存在的问题):
@@ -59,8 +72,13 @@ git diff HEAD --name-only
### Step 2 — 逐文件阅读并分析
- 用 Read 工具读取完整文件(不只读 diff
- 对照审查清单,记录每个问题:文件名、行号、问题类型、建议修复方式
先从 focused diff 开始:
```bash
git diff --unified=0 HEAD -- <file>
```
`rg``git diff --check`、编译器或 linter 输出确认确定性问题。只有需要上下文时才用 Read 读取完整文件。对照审查清单,记录每个问题:文件名、行号、问题类型、建议修复方式。
### Step 3 — 报告问题清单
@@ -95,6 +113,7 @@ git diff HEAD --name-only
- 只改在审查清单中发现的问题,不做额外优化
- 每次 Edit 只修改确实有问题的行,保持 diff 最小
- 改完后用 `grep` 验证旧的坏代码已消失
- 优先做精确补丁;只有仓库已有对应格式化流程时,才运行格式化工具
### Step 5 — 输出总结

93
.claude/commands/docs.md Normal file
View File

@@ -0,0 +1,93 @@
---
description: Create or update repository documentation from current code changes
argument-hint: Optional: topic to document, or leave empty to infer from git diff
allowed-tools: ["Read", "Edit", "Write", "Bash", "Glob", "Grep"]
---
# /docs — Documentation Workflow
## Goal
Create or update documentation that explains why a change exists, how it behaves, and what maintainers need to know. Keep this command generic. Repository-specific coverage rules live in the repository and must be loaded separately.
## Repository Rules
Before deciding scope, check whether the repository has a documentation rules file:
```bash
test -f docs/documentation-coverage-rules.md && sed -n '1,240p' docs/documentation-coverage-rules.md
```
If it exists, apply it as the project-specific coverage checklist. If it does not exist, continue with the generic workflow below.
## Workflow
### Step 1 — Understand The Change
```bash
git diff HEAD --stat
git diff HEAD --name-only
git log --oneline -10
rg --files docs
```
If `$ARGUMENTS` specifies a topic, focus on that topic. Otherwise infer the documentation topic from the changed files. Do not read the full repository diff by default; inspect focused files only:
```bash
git diff HEAD -- <path>
rg -n "class |def |function |export |router|@router|interface |type " <path>
```
### Step 2 — Decide Scope
- Prefer updating an existing relevant document over creating a duplicate.
- Use one document for one coherent topic.
- Split documents only when the change crosses meaningful domains.
- Keep filenames lowercase and hyphenated.
- Apply the repository-specific rules file before writing.
For ambiguous or large documentation changes, briefly state the intended doc plan before editing. For clear small changes, proceed directly.
### Step 3 — Write
Explain:
- Background/problem: what was wrong or missing before.
- Core design decisions and rationale.
- Operational or user-facing impact.
- Relevant code paths, only when useful for future maintainers.
Style:
- Follow the repositorys existing language and heading conventions.
- Use fenced code blocks with language tags.
- Prefer tables for comparisons or parameter lists.
- Keep snippets concise and relevant.
### Step 4 — Verify
- Read the completed docs once for clarity and stale statements.
- Verify referenced paths exist with `test -e` or `rg --files`.
- Run applicable checks from `docs/documentation-coverage-rules.md`.
- Check Markdown links use readable user-facing titles unless repository rules allow otherwise.
### Step 5 — Report
Summarize changed docs and verification:
```md
Updated:
- path/to/doc.md — what changed
Verified:
- checks that passed
- checks that could not be run, if any
```
## Hard Constraints
- Do not leave placeholder docs.
- Do not duplicate bilingual files byte-for-byte.
- Do not reference PR numbers, issue numbers, or the current conversation unless explicitly requested.
- Do not write changelog-style lists without the reasoning and tradeoffs behind the change.
- Keep docs maintainable and concise.

View File

@@ -72,6 +72,8 @@ Verification
## 执行风格
- 重证据,轻口头判断
- 优先使用确定性工具证据:`rg``git diff --stat``git diff -- <path>`、测试、构建、lint、`curl`、数据库查询等能直接证明成功标准的方式
- 不把大段命令输出粘进回复;保留在工具调用里,回复只总结关键证据
- 重验收,轻自我感觉
- 优先用测试、日志、产物、对比结果来证明完成
- 对长期任务保持“未达标就继续”的节奏

View File

@@ -28,6 +28,19 @@ allowed-tools: ["Read", "Edit", "Bash", "Glob", "Grep"]
- `docs/CHANGELOG.md`
- `docs/version-history.md`
## 节省上下文规则
发版判断应以确定性 CLI 证据为主,优先使用紧凑命令和定点读取:
```bash
git status --short
git diff --stat HEAD
git diff --name-only HEAD
rg -n "version|^## |^Released:|当前开发版本|current" VERSION frontend/package.json pyproject.toml docs/CHANGELOG.md docs/version-history.md
```
除非需要判断某个代码变更是否属于本次发版,否则不要读取完整 diff。
## 执行步骤
### Step 1 — 环境检查
@@ -45,7 +58,7 @@ cat VERSION # 读取当前版本
### Step 2 — 确定发版类型与新版本号
-`$ARGUMENTS` 提供了明确类型(`feature` / `bugfix`),直接使用
- 否则根据当前 `git diff HEAD``git log` 推断
- 否则根据 `git diff --stat HEAD``git diff --name-only HEAD`、必要的 focused diff`git log` 推断
- 计算新版本号(例:`0.26.2` → bugfix → `0.26.3`
- **先输出发版计划供用户确认**
@@ -91,12 +104,13 @@ cat VERSION # 读取当前版本
针对本次变更范围做最小验证:
- Python 文件有修改:`python3 -m py_compile <changed_files>`
- Frontend 文件有修改:运行项目标准检查(若无则跳过并说明)
- Python 文件有修改:先用 `git diff --name-only HEAD -- '*.py'` 列出,再运行 `python3 -m py_compile <changed_files>`
- Frontend 文件有修改:先用 `git diff --name-only HEAD -- frontend` 判断范围,再运行项目标准检查(若无则跳过并说明)
- 版本号一致性检查:用 grep 确认 VERSION、package.json、pyproject.toml 中的版本号完全一致
```bash
grep -h "version" VERSION frontend/package.json pyproject.toml
cat VERSION
rg -n "\"version\":|^version =|version = " frontend/package.json pyproject.toml uv.lock
```
### Step 7 — 提交前预览

View File

@@ -21,6 +21,19 @@ If the user specifies a file or directory, check only that. Otherwise check all
Only report issues present in **newly added or modified** lines of this diff — do not audit unchanged code.
## Token-Saving Rule
Prefer deterministic CLI checks before reading files into model context:
```bash
git diff --name-only HEAD
git diff --unified=0 HEAD -- <path>
git diff --check
rg -n "TODO|FIXME|console\.log|debugger|print\(" <changed-paths>
```
Read full files only when the focused diff does not provide enough surrounding context to make a safe edit.
## Checklist
### 1. Duplicate Logic
@@ -64,7 +77,13 @@ Filter to the user-specified path if one was provided.
### Step 2 — Read and analyze each file
Read the full file (not just the diff) with the Read tool. For each file, record every issue found: filename, line number, category, and suggested fix.
Start with focused diffs:
```bash
git diff --unified=0 HEAD -- <file>
```
Use `rg`, `git diff --check`, and compiler/linter output for deterministic findings. Read the full file only for files that need surrounding context. For each issue found, record filename, line number, category, and suggested fix.
### Step 3 — Report findings before touching anything
@@ -99,6 +118,7 @@ Principles:
- Only fix issues identified in the checklist — no extra improvements
- Keep each Edit as small as possible
- After fixing, verify the old bad pattern is gone with grep
- Prefer `apply_patch` for targeted edits; use formatters only when the repository already uses them for the touched file type
### Step 5 — Summary

View File

@@ -0,0 +1,82 @@
---
name: docs
description: Create or update repository documentation from current code changes. Use when the user asks to write docs, update docs, summarize implementation changes into docs, or check documentation coverage. Load repository-specific coverage rules from docs/documentation-coverage-rules.md when present.
---
# Docs
Use this skill when the task is documentation work: creating, updating, checking, or summarizing docs for code or behavior changes.
## Goal
Write documentation that explains why a change exists, how it behaves, and what maintainers need to know. Keep the skill generic; repository-specific rules belong in the repository, not in this skill.
## Repository Rules
Before deciding scope, check whether the repository has a documentation rules file:
```bash
test -f docs/documentation-coverage-rules.md && sed -n '1,240p' docs/documentation-coverage-rules.md
```
If it exists, apply it as the project-specific coverage checklist. If it does not exist, continue with the generic workflow below.
## Workflow
1. Gather focused context:
```bash
git diff HEAD --stat
git diff HEAD --name-only
git log --oneline -10
rg --files docs
```
If the user gives a topic, focus on that topic. Otherwise infer the doc topic from changed files. Avoid reading large full diffs by default; inspect focused files and symbols:
```bash
git diff HEAD -- <path>
rg -n "class |def |function |export |router|@router|interface |type " <path>
```
2. Decide scope:
- Prefer updating an existing relevant doc over creating a duplicate.
- Use one document for one coherent topic.
- Split documents only when changes cross meaningful domains.
- Keep filenames lowercase and hyphenated.
3. Write the doc:
- Explain background/problem, design decisions, constraints, and operational impact.
- Keep code snippets short and directly relevant.
- List related files only when they help future maintainers navigate.
- Use the repositorys existing language, heading style, and naming conventions.
4. Verify:
- Read the completed doc once for clarity and stale statements.
- Verify important referenced paths exist with `test -e` or `rg --files`.
- Run repository-specific doc checks from `docs/documentation-coverage-rules.md` when present.
- For Markdown links, check that user-facing titles are readable and not raw filenames unless the repository rules allow it.
## Hard Constraints
- Do not leave placeholder docs or copied source text pretending to be documentation.
- Do not duplicate bilingual files byte-for-byte.
- Do not reference PR numbers, issue numbers, or the current conversation unless explicitly requested.
- Do not write changelog-style lists without the reasoning, constraints, and tradeoffs behind the change.
- Keep docs concise enough to maintain.
## Recommended Output
After editing, summarize:
```md
Updated:
- path/to/doc.md — what changed
Verified:
- checks that passed
- checks that could not be run, if any
```

View File

@@ -72,6 +72,8 @@ In Codex, only use actual subagents when the user explicitly asks for delegation
## Operating Rules
- Prefer objective checks over self-reported completion.
- Prefer deterministic tool evidence over long model summaries: use `rg`, `git diff --stat`, targeted `git diff -- <path>`, tests, builds, linters, `curl`, or database queries when they can prove a criterion.
- Do not paste large command output into the conversation; summarize the evidence and keep raw output in tool calls.
- Do not confuse progress with completion.
- If the worker says "done", verify it.
- If verification fails, continue from the gap instead of restarting blindly.

View File

@@ -18,7 +18,7 @@ Do not use this skill for ordinary commits that are not being released.
## Versioning Rules
- `feature` -> bump `+0.1.0`
- `feature` -> bump minor and reset patch to `0` (`x.y.z``x.(y+1).0`; for example `0.41.2``0.42.0`)
- `bugfix` -> bump `+0.0.1`
- `docs`, `maintenance`, and `refactor` do not bump by default unless the user explicitly wants a release
@@ -35,6 +35,19 @@ Use `git rev-parse --show-toplevel` to get the repo root. All paths are relative
- `docs/CHANGELOG.md`
- `docs/version-history.md`
## Token-Saving Rule
Release work should be driven by deterministic CLI evidence. Prefer compact commands and targeted file reads:
```bash
git status --short
git diff --stat HEAD
git diff --name-only HEAD
rg -n "version|^## |^Released:|current" VERSION frontend/package.json pyproject.toml docs/CHANGELOG.md docs/version-history.md
```
Do not inspect full diffs unless deciding whether changed code belongs in the release.
## Workflow
### Step 1 — Environment check
@@ -52,8 +65,10 @@ If unrelated uncommitted changes exist, list them and ask the user whether to in
### Step 2 — Determine release type and next version
- If the user provided an explicit type (`feature` / `bugfix`), use it
- Otherwise infer from `git diff HEAD` and recent `git log`
- Compute the next version (e.g. `0.26.2` → bugfix → `0.26.3`)
- Otherwise infer from `git diff --stat HEAD`, `git diff --name-only HEAD`, focused diffs for changed code, and recent `git log`
- Compute the next version:
- `feature`: increment minor and reset patch to `0` (e.g. `0.41.2``0.42.0`)
- `bugfix`: increment patch only (e.g. `0.26.2``0.26.3`)
- **Show the release plan before making any changes:**
```
@@ -104,12 +119,13 @@ Get today's date with `date +%Y-%m-%d`.
Run the smallest relevant validation for the changes in scope:
- Python files changed: `python3 -m py_compile <changed_files>`
- Frontend files changed: run the project-standard check if available; otherwise skip and say so
- Python files changed: list changed Python files with `git diff --name-only HEAD -- '*.py'`, then run `python3 -m py_compile <changed_files>`
- Frontend files changed: list changed frontend files with `git diff --name-only HEAD -- frontend`, then run the project-standard check if available; otherwise skip and say so
- Version consistency: confirm VERSION, package.json, pyproject.toml, and uv.lock all show the same version
```bash
grep -h "version" VERSION frontend/package.json pyproject.toml
cat VERSION
rg -n "\"version\":|^version =|version = " frontend/package.json pyproject.toml uv.lock
```
### Step 7 — Pre-commit preview

13
.dockerignore Normal file
View File

@@ -0,0 +1,13 @@
**
!pyproject.toml
!uv.lock
!aiprovider/
!aiprovider/**
aiprovider/.env
aiprovider/.env.*
!aiprovider/.env.example
**/__pycache__/
**/*.pyc
**/*.pyo

View File

@@ -236,6 +236,8 @@ bun run build
推荐按下面顺序排查和配置。
端口占用、`iphlpsvc` / portproxy、摄像头和依赖问题的集中排障入口见 [常见问题](/home/ray/dev/linkong/planet/docs/technical/zh/faq.md)。
### 1. 在 WSL 中启动服务
```bash

View File

@@ -23,8 +23,16 @@
- [ ] 保持 Earth 当前这批纯个人偏好设置继续走本地持久化:`旋转模式`、HUD 面板显示/隐藏、`地形透明度` 暂不升级到后端系统设置,避免把设备级偏好过早做成全局配置
- [ ] 如果后续明确需要“账号级同步 Earth 偏好”,再单独设计 `Earth user preferences`:优先按用户维度而不是全局系统设置保存,并规划 `localStorage -> backend` 的平滑迁移策略
- [ ] 为 Planet / Earth 补一个可用的日志查看系统:先明确前后端/AI Provider/采集任务的日志入口、最近日志聚合、筛选与 tail 能力,再决定是先做脚本级统一入口还是控制台内置日志面板
- [ ] 重写控制台 UI逐步抛弃 Ant Design建立自有组件体系并统一采用 `tabler.io` / Tabler Icons 作为控制台主图标库
- [ ] 把 Earth 态势新闻源从 [earth_news.py](/home/ray/dev/linkong/planet/backend/app/services/earth_news.py) 的硬编码列表抽成可配置目录,优先保持当前“实时聚合”链路不变,只先解决新闻源不可配置的问题
- [ ] 为 Earth 态势新闻设计后续采集器化方案:明确新闻数据模型、去重策略、区域映射、过期清理和 Earth/AI 复用方式,再决定何时把新闻从实时抓取升级成正式 collector
- [ ] AIS v3.1:修复船只聚合完整性,`/geo/vessels` 合并 raw observation 聚合结果与 legacy `vessel_position + vessel_static` 最新结果,确保 BarentsWatch-only 船只不会因为 AISStream 子集存在而消失,并增加 raw/legacy/final unique MMSI 诊断统计
- [ ] AIS v3.2:把 AISStream 从收满 `max_messages` 后结束的批采集改成长连接 streaming service持续写入 raw observations通过内部 `/ws``vessels` channel 推送新船、位置和航向增量Earth 前端按 MMSI upsert marker
- [ ] AIS v3.3:修正 AISStream 采集页面状态语义,使用 connecting/streaming/reconnecting/stopped 与 indeterminate 状态展示运行时长、消息数、unique MMSI、message rate、最近消息和错误不再用一次性 REST 进度条表示长连接
- [ ] AIS v3.4修复船只身份字段和名称聚合MMSI/IMO/callsign 按字符串显示且不带千分位符;查询并列出所有仍以 MMSI 号码或 `MMSI <number>` 作为船名的记录标注来源、最近观测、message types 和缺失原因,并把这批 fallback-name 船只纳入名称聚合修复集合
- [ ] Earth Live Sync建立统一态势实时同步链路新增 `earth_summary` WS channel任意采集器成功后广播轻量 summary invalidation前端收到后重新拉 `/api/v1/visualization/geo/summary` 并更新 HUD同时为 BGP 增加 `bgp` WS channel使 BGP incidents/anomalies/collectors 在不刷新页面时也能 upsert 图层;卫星采集完成后触发 summary 刷新,必要时按 TLE 版本重新 hydrate 卫星数据
- [ ] AIS v4开放船只多源聚合策略配置支持 source priority、字段级规则、freshness 窗口和高级保护开关;保存时校验未知字段、非法模式和危险动态字段锁定,并在聚合接口返回命中的配置版本
- [ ] AIS v5实现船舶资料 enrichment 与冲突治理,按 `mmsi + imo + name + callsign` 异步补充船型细分、AIS 大类、旗国、尺寸、建造年份、运营方和图片缓存;详情面板展示缓存资料和字段来源,不在实时 AIS 请求链路现场抓第三方页面
- [ ] 为 Earth 地球表面增加一层与基础纹理对齐的材质/纹理 overlay并在同层叠加国界轮廓参考线要求国界线与底图稳定对齐且 hover 到国家轮廓时能高亮当前国家,便于校准地表和增强交互
- [ ] 把 Earth 新闻接入通用巡航队列:按新闻发生地和时间排序生成巡航目标,巡航聚焦到新闻事件时显示对应新闻卡片,并保持实现边界为“通用巡航层 + 新闻业务适配层”,不要再把新闻逻辑直接耦合回 `main.js` 状态机
- [ ] 为未知位置的算力中心建立分层坐标补全链路:优先 `精确坐标 > 站点/园区命中 > 城市 > 州/省 > 国家内主要算力城市 > 国家质心`,并把每次回退的 `confidence / reason / precision` 明确写进统一 GeoJSON

View File

@@ -1 +1 @@
0.41.2
0.51.0

View File

@@ -32,6 +32,15 @@ AI_API_KEY=sk-cp-change-me
AI_MAX_TOKENS=1200
AI_ANTHROPIC_VERSION=2023-06-01
# Optional provider-specific keys used by Settings fallback before AI_API_KEY
# MINIMAX_API_KEY=sk-cp-change-me
# OPENAI_API_KEY=sk-change-me
# ANTHROPIC_API_KEY=sk-ant-change-me
# DEEPSEEK_API_KEY=sk-change-me
# DASHSCOPE_API_KEY=sk-change-me
# MOONSHOT_API_KEY=sk-change-me
# OPENROUTER_API_KEY=sk-or-change-me
# OpenAI-compatible example (vLLM / LM Studio / One API / local gateway)
# AI_PROVIDER=openai
# AI_PROVIDER_API=openai-completions

View File

@@ -1,3 +1,5 @@
# syntax=docker/dockerfile:1.7
ARG PYTHON_IMAGE=python:3.14-slim
ARG UV_IMAGE=ghcr.io/astral-sh/uv:latest
@@ -18,9 +20,10 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
&& rm -rf /var/lib/apt/lists/*
COPY pyproject.toml uv.lock /app/
RUN uv sync --frozen --no-dev
RUN --mount=type=cache,target=/root/.cache/uv \
uv sync --frozen --no-dev
COPY . /app
COPY aiprovider /app/aiprovider
EXPOSE 8010

View File

@@ -37,8 +37,26 @@ def verify_service_token(x_provider_token: str | None = Header(default=None)) ->
)
def get_provider_service() -> ProviderService:
return ProviderService()
def get_provider_service(
x_ai_provider: str | None = Header(default=None),
x_ai_provider_api: str | None = Header(default=None),
x_ai_base_url: str | None = Header(default=None),
x_ai_api_key: str | None = Header(default=None),
x_ai_model: str | None = Header(default=None),
x_ai_max_tokens: str | None = Header(default=None),
x_ai_anthropic_version: str | None = Header(default=None),
) -> ProviderService:
overrides = {
"provider": x_ai_provider,
"provider_api": x_ai_provider_api,
"base_url": x_ai_base_url,
"api_key": x_ai_api_key,
"model": x_ai_model,
"anthropic_version": x_ai_anthropic_version,
}
if x_ai_max_tokens:
overrides["max_tokens"] = x_ai_max_tokens
return ProviderService({key: value for key, value in overrides.items() if value not in (None, "")})
@app.get("/health")

View File

@@ -46,19 +46,22 @@ def _resolve_provider_api(provider: str, configured_api: str) -> str:
class ProviderService:
def __init__(self) -> None:
self.provider = _normalize_provider(settings.AI_PROVIDER)
def __init__(self, overrides: dict[str, Any] | None = None) -> None:
overrides = overrides or {}
self.provider = _normalize_provider(overrides.get("provider") or settings.AI_PROVIDER)
self.provider_api = _resolve_provider_api(
self.provider,
_normalize_provider_api(settings.AI_PROVIDER_API),
_normalize_provider_api(overrides.get("provider_api") or settings.AI_PROVIDER_API),
)
self.base_url = settings.AI_BASE_URL.rstrip("/")
self.api_key = settings.AI_API_KEY
self.default_model = settings.AI_MODEL
self.base_url = str(overrides.get("base_url") or settings.AI_BASE_URL).rstrip("/")
self.api_key = str(overrides.get("api_key") or settings.AI_API_KEY)
self.default_model = str(overrides.get("model") or settings.AI_MODEL)
self.timeout = settings.AI_TIMEOUT_SECONDS
self.http_retry_attempts = max(settings.AI_HTTP_RETRY_ATTEMPTS, 1)
self.max_tokens = settings.AI_MAX_TOKENS
self.anthropic_version = settings.AI_ANTHROPIC_VERSION
self.max_tokens = int(overrides.get("max_tokens") or settings.AI_MAX_TOKENS)
self.anthropic_version = str(
overrides.get("anthropic_version") or settings.AI_ANTHROPIC_VERSION
)
self.system_prompt = settings.AI_ANALYSIS_SYSTEM_PROMPT
def get_status(self) -> AIProviderStatusResponse:

View File

@@ -5,6 +5,7 @@ from app.api.v1 import (
users,
datasource_config,
datasources,
docs,
tasks,
dashboard,
websocket,
@@ -12,6 +13,7 @@ from app.api.v1 import (
settings,
collected_data,
visualization,
vessel_aggregation,
bgp,
news,
system_control,
@@ -28,12 +30,18 @@ api_router.include_router(
)
api_router.include_router(datasources.router, prefix="/datasources", tags=["datasources"])
api_router.include_router(collected_data.router, prefix="/collected", tags=["collected-data"])
api_router.include_router(docs.router, prefix="/docs", tags=["docs"])
api_router.include_router(tasks.router, prefix="/tasks", tags=["tasks"])
api_router.include_router(dashboard.router, prefix="/dashboard", tags=["dashboard"])
api_router.include_router(alerts.router, prefix="/alerts", tags=["alerts"])
api_router.include_router(settings.router, prefix="/settings", tags=["settings"])
api_router.include_router(system_control.router, prefix="/system", tags=["system"])
api_router.include_router(visualization.router, prefix="/visualization", tags=["visualization"])
api_router.include_router(
vessel_aggregation.router,
prefix="/vessel-aggregation",
tags=["vessel-aggregation"],
)
api_router.include_router(bgp.router, prefix="/bgp", tags=["bgp"])
api_router.include_router(tv.router, prefix="/tv", tags=["tv"])
api_router.include_router(news.router, prefix="/news", tags=["news"])

View File

@@ -28,7 +28,7 @@ async def login(
):
result = await db.execute(
text(
"SELECT id, username, email, password_hash, role, is_active FROM users WHERE username = :username"
"SELECT id, username, email, password_hash, role, is_active, gatekeeper_groups FROM users WHERE username = :username"
),
{"username": form_data.username},
)
@@ -46,6 +46,7 @@ async def login(
user.password_hash = row[3]
user.role = row[4]
user.is_active = row[5]
user.gatekeeper_groups = row[6] or []
if not verify_password(form_data.password, user.password_hash):
raise HTTPException(
@@ -73,6 +74,7 @@ async def login(
"id": user.id,
"username": user.username,
"role": user.role,
"gatekeeper_groups": user.gatekeeper_groups or [],
},
}
@@ -95,6 +97,7 @@ async def refresh_token(
"id": current_user.id,
"username": current_user.username,
"role": current_user.role,
"gatekeeper_groups": current_user.gatekeeper_groups or [],
},
}
@@ -111,6 +114,7 @@ async def get_me(current_user: User = Depends(get_current_user)):
"username": current_user.username,
"email": current_user.email,
"role": current_user.role,
"gatekeeper_groups": current_user.gatekeeper_groups or [],
"is_active": current_user.is_active,
"created_at": current_user.created_at,
}

View File

@@ -5,13 +5,26 @@ from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from pydantic import BaseModel
from app.core.security import get_current_user
from app.db.session import get_db
from app.models.bgp_anomaly import BGPAnomaly
from app.models.bgp_incident import BGPIncident
from app.models.bgp_observation import BGPObservation
from app.models.user import User
from app.services.bgp_collector_locations import (
build_bgp_collector_location_query,
collect_bgp_collector_location_candidates,
get_bgp_collector_location_dict,
)
from app.services.bgp_collectors import build_bgp_collector_coverage
from app.services.ai_client import get_ai_provider_client
from app.api.v1.settings import get_web_search_client
from app.services.location.llm_fallback import (
collect_llm_location_fallback_candidate,
collect_location_search_evidence,
)
router = APIRouter()
@@ -264,6 +277,119 @@ async def get_bgp_collector_summary(
}
class CollectBGPCollectorLocationRequest(BaseModel):
city: Optional[str] = None
country: Optional[str] = None
site: Optional[str] = None
operator: Optional[str] = None
@router.post("/collectors/{collector_id}/collect-location")
async def collect_bgp_collector_location(
collector_id: str,
payload: CollectBGPCollectorLocationRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""Run the shared location pipeline for a BGP route collector.
Mirrors ``POST /api/v1/visualization/compute-centers/{source_id}/collect-location``.
Returns ranked candidates from source coordinates and Nominatim queries
built around the collector's stored context (IXP / city / country). Stored
collector locations provide context only; they are not emitted as
candidates.
"""
if not collector_id or not collector_id.strip():
raise HTTPException(status_code=400, detail="collector_id is required")
legacy = get_bgp_collector_location_dict(collector_id) or {}
site = payload.site or legacy.get("matched_location_name")
city = payload.city or legacy.get("city")
country = payload.country or legacy.get("country")
operator = payload.operator or "RIPE NCC"
candidates, attempted_queries = collect_bgp_collector_location_candidates(
collector=collector_id,
site=site,
city=city,
country=country,
operator=operator,
)
llm_failure_reason = None
if not candidates:
query = build_bgp_collector_location_query(
collector=collector_id,
site=site,
city=city,
country=country,
operator=operator,
)
llm_result = None
try:
web_search_client = await get_web_search_client(db)
search_result = await collect_location_search_evidence(
web_search_client=web_search_client,
query=query,
entity_type="bgp_collector",
)
attempted_queries = [*attempted_queries, *search_result.attempted_queries]
if not search_result.evidence:
llm_failure_reason = search_result.failure_reason
raise RuntimeError(search_result.failure_reason or "no WebSearch evidence")
provider_client = await get_ai_provider_client(db)
llm_result = await collect_llm_location_fallback_candidate(
provider_client=provider_client,
query=query,
entity_type="bgp_collector",
attempted_queries=attempted_queries,
search_evidence=search_result.evidence,
)
except Exception as exc:
if llm_failure_reason is None:
llm_failure_reason = f"LLM location factcheck unavailable: {exc}"
attempted_queries = [
*attempted_queries,
f"llm_factcheck:bgp_collector:{collector_id or 'unknown'}",
]
if llm_result is not None:
attempted_queries = [*attempted_queries, *llm_result.attempted_queries]
candidates = llm_result.candidates
llm_failure_reason = llm_result.failure_reason
context = {
"collector": collector_id,
"site": site,
"city": city,
"country": country,
"operator": operator,
}
if not candidates:
return {
"collector_id": collector_id,
"name": collector_id,
"success": False,
"failure_reason": (
"No source coordinates or online geocoding result reached"
" city-level precision for this collector."
),
"candidates": [],
"attempted_queries": list(attempted_queries),
"llm_failure_reason": llm_failure_reason,
"context": context,
}
return {
"collector_id": collector_id,
"name": collector_id,
"success": True,
"candidates": [candidate.to_dict() for candidate in candidates],
"best_candidate": candidates[0].to_dict(),
"attempted_queries": list(attempted_queries),
"context": context,
}
@router.get("/overview/summary")
async def get_bgp_overview_summary(
current_user: User = Depends(get_current_user),

View File

@@ -1,20 +1,52 @@
"""DataSourceConfig API for user-defined data sources"""
from typing import Optional
from typing import Any, Optional
from datetime import datetime
import base64
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy import select, func
import json
import re
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import delete, select, func
from sqlalchemy.ext.asyncio import AsyncSession
from pydantic import BaseModel, Field
import httpx
from app.core.target_schema_registry import get_target_schema, list_target_schemas
from app.core.datasource_defaults import DEFAULT_DATASOURCES
from app.db.session import get_db
from app.models.user import User
from app.models.datasource_config import DataSourceConfig
from app.models.datasource_mapping import DataSourceMappingTemplate
from app.models.collected_data import CollectedData
from app.models.vessel import AISRawObservation, AISSourceHealth
from app.core.security import get_current_user
from app.core.cache import cache
from app.core.time import to_iso8601_utc
from app.schemas.ai import SituationalAnalysisRequest
from app.services.ai_client import AIProviderClient, get_ai_provider_client
from app.services.datasource_mapping import (
MappingError,
build_heuristic_mapping,
execute_mapping,
redact_for_llm,
stable_payload_hash,
)
from app.services.custom_datasource_runtime import (
CustomDatasourceRuntimeError,
fetch_rest_payload,
get_custom_stream_status,
run_mapped_rest_config,
run_mapped_websocket_config,
start_custom_stream,
stop_custom_stream,
test_websocket_config,
)
from app.services.datasource_connectivity import (
get_builtin_connection_status,
save_connectivity_success,
strip_connectivity_validation,
test_builtin_connectivity,
)
router = APIRouter()
@@ -22,7 +54,7 @@ router = APIRouter()
class DataSourceConfigCreate(BaseModel):
name: str = Field(..., min_length=1, max_length=100)
description: Optional[str] = None
source_type: str = Field(..., description="http, api, database")
source_type: str = Field(..., description="rest, websocket, http, api, database")
endpoint: str = Field(..., max_length=500)
auth_type: str = Field(default="none", description="none, bearer, api_key, basic")
auth_config: dict = Field(default={})
@@ -59,6 +91,70 @@ class DataSourceConfigResponse(BaseModel):
from_attributes = True
def _is_builtin_config_name(name: str | None) -> bool:
return bool(name and name in DEFAULT_DATASOURCES)
async def _ensure_builtin_connection_verified(
db: AsyncSession,
config_data: DataSourceConfigCreate,
) -> None:
if not _is_builtin_config_name(config_data.name):
return
status_result = await get_builtin_connection_status(
db,
config_data.name,
config_data.endpoint,
config_data.auth_type,
config_data.headers,
config_data.config,
)
if not status_result.get("connected"):
raise HTTPException(
status_code=400,
detail=status_result.get("message") or "请先完成连接验证,再保存内置采集器配置。",
)
class CustomSampleRequest(BaseModel):
datasource_config_id: Optional[int] = None
config: Optional[DataSourceConfigCreate] = None
limit_bytes: int = Field(default=200000, ge=1000, le=1000000)
class MappingProposeRequest(BaseModel):
sample_payload: Any
target_schema: str
use_ai: bool = True
class MappingPreviewRequest(BaseModel):
sample_payload: Any
target_schema: str
mapping_json: dict
limit: int = Field(default=20, ge=1, le=100)
class MappingTemplateCreate(BaseModel):
datasource_config_id: int
target_schema: str
mapping_json: dict
sample_payload: Any | None = None
sample_payload_hash: Optional[str] = None
validation_status: str = Field(default="draft", pattern="^(draft|valid|invalid)$")
is_active: bool = False
class MappingTemplateUpdate(BaseModel):
target_schema: Optional[str] = None
mapping_json: Optional[dict] = None
sample_payload: Any | None = None
sample_payload_hash: Optional[str] = None
validation_status: Optional[str] = Field(default=None, pattern="^(draft|valid|invalid)$")
is_active: Optional[bool] = None
async def test_endpoint(
endpoint: str,
auth_type: str,
@@ -96,6 +192,136 @@ async def test_endpoint(
}
def _build_request_headers(auth_type: str, auth_config: dict, headers: dict) -> dict[str, str]:
request_headers = {str(key): str(value) for key, value in (headers or {}).items()}
auth_type = str(auth_type or "none").lower()
auth_config = auth_config or {}
if auth_type == "bearer" and auth_config.get("token"):
request_headers["Authorization"] = f"Bearer {auth_config['token']}"
elif auth_type == "api_key" and auth_config.get("api_key"):
location = str(auth_config.get("in") or auth_config.get("location") or "header").lower()
if location != "query":
key_name = auth_config.get("key_name", "X-API-Key")
request_headers[str(key_name)] = str(auth_config["api_key"])
elif auth_type == "basic":
username = auth_config.get("username", "")
password = auth_config.get("password", "")
credentials = f"{username}:{password}"
encoded = base64.b64encode(credentials.encode()).decode()
request_headers["Authorization"] = f"Basic {encoded}"
return request_headers
def _build_query_params(auth_type: str, auth_config: dict, config: dict) -> dict[str, Any]:
params = {}
candidate = (config or {}).get("params") or (config or {}).get("query_params")
if isinstance(candidate, dict):
params.update(candidate)
auth_type = str(auth_type or "none").lower()
auth_config = auth_config or {}
if auth_type == "api_key" and auth_config.get("api_key"):
location = str(auth_config.get("in") or auth_config.get("location") or "header").lower()
if location == "query":
key_name = auth_config.get("key_name") or auth_config.get("param_name") or "api_key"
params[str(key_name)] = auth_config["api_key"]
return params
async def fetch_custom_sample_from_config(config: DataSourceConfig, limit_bytes: int) -> Any:
if str(config.source_type or "").lower() in {"websocket", "ws"}:
raise HTTPException(status_code=400, detail="WebSocket sources must use connection test or run-mapped stream.")
request_config = config.config or {}
method = str(request_config.get("method") or request_config.get("request_method") or "GET").upper()
if method not in {"GET", "POST"}:
raise HTTPException(status_code=400, detail="Only GET and POST sample requests are supported.")
headers = _build_request_headers(config.auth_type, config.auth_config or {}, config.headers or {})
params = _build_query_params(config.auth_type, config.auth_config or {}, request_config)
timeout = float(request_config.get("timeout", 30))
json_body = request_config.get("json_body")
if json_body is None and str(request_config.get("body_type") or "").lower() in {"json", ""}:
candidate = request_config.get("body")
if isinstance(candidate, (dict, list)):
json_body = candidate
async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client:
response = await client.request(
method,
config.endpoint,
headers=headers,
params=params or None,
json=json_body,
)
response.raise_for_status()
content = response.content[:limit_bytes]
if "application/json" in response.headers.get("content-type", ""):
return json.loads(content.decode(response.encoding or "utf-8"))
return {"text": content.decode(response.encoding or "utf-8", errors="replace")}
def _parse_mapping_from_ai_text(content: str) -> dict[str, Any] | None:
if not content:
return None
candidates = [content]
fenced = re.findall(r"```(?:json)?\s*(\{.*?\})\s*```", content, flags=re.DOTALL)
candidates = fenced + candidates
for candidate in candidates:
try:
parsed = json.loads(candidate)
except json.JSONDecodeError:
continue
if isinstance(parsed, dict) and isinstance(parsed.get("fields"), dict):
return parsed
return None
async def _get_config_for_sample(
payload: CustomSampleRequest,
db: AsyncSession,
) -> DataSourceConfig:
if payload.datasource_config_id is not None:
result = await db.execute(
select(DataSourceConfig).where(DataSourceConfig.id == payload.datasource_config_id)
)
config = result.scalar_one_or_none()
if not config:
raise HTTPException(status_code=404, detail="Configuration not found")
return config
if payload.config is None:
raise HTTPException(status_code=400, detail="datasource_config_id or config is required")
config_data = payload.config
return DataSourceConfig(
name=config_data.name,
description=config_data.description,
source_type=config_data.source_type,
endpoint=config_data.endpoint,
auth_type=config_data.auth_type,
auth_config=config_data.auth_config,
headers=config_data.headers,
config=config_data.config,
)
def serialize_mapping_template(template: DataSourceMappingTemplate) -> dict[str, Any]:
return {
"id": template.id,
"datasource_config_id": template.datasource_config_id,
"target_schema": template.target_schema,
"mapping_json": template.mapping_json,
"sample_payload_hash": template.sample_payload_hash,
"validation_status": template.validation_status,
"version": template.version,
"is_active": template.is_active,
"created_at": to_iso8601_utc(template.created_at),
"updated_at": to_iso8601_utc(template.updated_at),
}
@router.get("/configs")
async def list_configs(
active_only: bool = False,
@@ -105,7 +331,7 @@ async def list_configs(
"""List all user-defined data source configurations"""
query = select(DataSourceConfig)
if active_only:
query = query.where(DataSourceConfig.is_active == True)
query = query.where(DataSourceConfig.is_active)
query = query.order_by(DataSourceConfig.created_at.desc())
result = await db.execute(query)
@@ -132,6 +358,52 @@ async def list_configs(
}
@router.get("/configs/all")
async def list_all_datasources(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""List all data sources: YAML defaults + DB overrides"""
from app.core.data_sources import COLLECTOR_URL_KEYS, get_data_sources_config
config = get_data_sources_config()
db_query = await db.execute(select(DataSourceConfig))
db_configs = {c.name: c for c in db_query.scalars().all()}
result = []
for name, yaml_key in COLLECTOR_URL_KEYS.items():
yaml_url = config.get_yaml_url(name)
db_config = db_configs.get(name)
result.append(
{
"name": name,
"default_url": yaml_url,
"endpoint": db_config.endpoint if db_config else yaml_url,
"is_overridden": db_config is not None and db_config.endpoint != yaml_url
if yaml_url
else db_config is not None,
"is_active": db_config.is_active if db_config else True,
"source_type": db_config.source_type if db_config else "http",
"auth_type": db_config.auth_type if db_config else "none",
"auth_configured": {
"api_key": bool((db_config.auth_config or {}).get("api_key"))
if db_config
else False,
},
"headers": db_config.headers if db_config else {},
"config": strip_connectivity_validation(db_config.config if db_config else {}),
"config_id": db_config.id if db_config else None,
"description": db_config.description
if db_config
else f"Data source from YAML: {yaml_key}",
}
)
return {"total": len(result), "data": result}
@router.get("/configs/{config_id}")
async def get_config(
config_id: int,
@@ -176,7 +448,7 @@ async def create_config(
auth_type=config_data.auth_type,
auth_config=config_data.auth_config,
headers=config_data.headers,
config=config_data.config,
config=strip_connectivity_validation(config_data.config),
)
db.add(config)
@@ -208,6 +480,10 @@ async def update_config(
update_data = config_data.model_dump(exclude_unset=True)
for field, value in update_data.items():
if field == "config":
value = strip_connectivity_validation(value)
if field == "auth_config" and value == {} and (config.auth_config or {}):
continue
setattr(config, field, value)
await db.commit()
@@ -225,6 +501,8 @@ async def update_config(
@router.delete("/configs/{config_id}")
async def delete_config(
config_id: int,
delete_mappings: bool = Query(False),
delete_source_data: bool = Query(False),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
@@ -235,12 +513,59 @@ async def delete_config(
if not config:
raise HTTPException(status_code=404, detail="Configuration not found")
deleted_mappings = 0
deleted_records = {
"collected_data": 0,
"ais_raw_observations": 0,
"ais_source_health": 0,
}
if delete_source_data:
collected_result = await db.execute(
delete(CollectedData).where(CollectedData.source == config.name)
)
raw_result = await db.execute(
delete(AISRawObservation).where(AISRawObservation.source == config.name)
)
health_result = await db.execute(
delete(AISSourceHealth).where(AISSourceHealth.source == config.name)
)
deleted_records = {
"collected_data": collected_result.rowcount or 0,
"ais_raw_observations": raw_result.rowcount or 0,
"ais_source_health": health_result.rowcount or 0,
}
if delete_mappings or delete_source_data:
mapping_result = await db.execute(
delete(DataSourceMappingTemplate).where(
DataSourceMappingTemplate.datasource_config_id == config_id
)
)
deleted_mappings = mapping_result.rowcount or 0
await db.delete(config)
await db.commit()
cache.delete_pattern("datasource_configs:*")
return {"message": "Configuration deleted successfully"}
if delete_source_data and (config.config or {}).get("target_schema") == "vessel_ais":
from app.core.websocket.broadcaster import broadcaster
await broadcaster.broadcast_custom(
"vessels",
{
"action": "reload",
"source": config.name,
"reason": "custom_source_deleted",
},
)
return {
"message": "Configuration deleted successfully",
"deleted_mappings": deleted_mappings,
"deleted_records": deleted_records,
}
@router.post("/configs/{config_id}/test")
@@ -257,6 +582,8 @@ async def test_config(
raise HTTPException(status_code=404, detail="Configuration not found")
try:
if str(config.source_type or "").lower() in {"websocket", "ws"}:
return await test_websocket_config(config)
result = await test_endpoint(
endpoint=config.endpoint,
auth_type=config.auth_type,
@@ -287,6 +614,18 @@ async def test_new_config(
):
"""Test a new data source configuration without saving"""
try:
if str(config_data.source_type or "").lower() in {"websocket", "ws"}:
config = DataSourceConfig(
name=config_data.name,
description=config_data.description,
source_type=config_data.source_type,
endpoint=config_data.endpoint,
auth_type=config_data.auth_type,
auth_config=config_data.auth_config,
headers=config_data.headers,
config=config_data.config,
)
return await test_websocket_config(config)
result = await test_endpoint(
endpoint=config_data.endpoint,
auth_type=config_data.auth_type,
@@ -310,38 +649,363 @@ async def test_new_config(
}
@router.get("/configs/all")
async def list_all_datasources(
@router.post("/configs/builtin/connection-status")
async def get_builtin_config_connection_status(
config_data: DataSourceConfigCreate,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""List all data sources: YAML defaults + DB overrides"""
from app.core.data_sources import COLLECTOR_URL_KEYS, get_data_sources_config
if not _is_builtin_config_name(config_data.name):
raise HTTPException(status_code=400, detail="Only built-in datasource configs are supported.")
config = get_data_sources_config()
return await get_builtin_connection_status(
db,
config_data.name,
config_data.endpoint,
config_data.auth_type,
config_data.headers,
config_data.config,
)
db_query = await db.execute(select(DataSourceConfig))
db_configs = {c.name: c for c in db_query.scalars().all()}
result = []
for name, yaml_key in COLLECTOR_URL_KEYS.items():
yaml_url = config.get_yaml_url(name)
db_config = db_configs.get(name)
@router.post("/configs/builtin/connect")
async def connect_builtin_config(
config_data: DataSourceConfigCreate,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
if not _is_builtin_config_name(config_data.name):
raise HTTPException(status_code=400, detail="Only built-in datasource configs are supported.")
result.append(
result = await test_builtin_connectivity(
config_data.name,
config_data.endpoint,
config_data.auth_type,
config_data.headers,
config_data.config,
db,
config_data.auth_config,
)
if result.get("success") and result.get("checksum"):
validation = await save_connectivity_success(
db,
config_data.name,
result["checksum"],
result,
connected_by="connection_button",
)
await db.commit()
return {
**result,
"connected": True,
"validation": validation,
}
return {
**result,
"connected": False,
}
@router.post("/custom/sample")
async def fetch_custom_sample(
payload: CustomSampleRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""Fetch a sample payload for a saved or draft custom data source."""
config = await _get_config_for_sample(payload, db)
try:
sample = await fetch_custom_sample_from_config(config, payload.limit_bytes)
except httpx.HTTPStatusError as exc:
raise HTTPException(
status_code=exc.response.status_code,
detail=f"Sample request failed: HTTP {exc.response.status_code}",
) from exc
except httpx.HTTPError as exc:
raise HTTPException(status_code=502, detail=f"Sample request failed: {exc}") from exc
return {
"success": True,
"sample_payload": sample,
"sample_payload_hash": stable_payload_hash(sample),
"redacted_preview": redact_for_llm(sample),
}
@router.get("/target-schemas")
async def get_datasource_target_schemas(
current_user: User = Depends(get_current_user),
):
"""List target schemas available for custom datasource mapping."""
return {"data": list_target_schemas()}
@router.post("/mappings/propose")
async def propose_datasource_mapping(
payload: MappingProposeRequest,
current_user: User = Depends(get_current_user),
ai_client: AIProviderClient = Depends(get_ai_provider_client),
):
"""Generate a mapping draft for a sample payload and target schema."""
schema = get_target_schema(payload.target_schema)
redacted_sample = redact_for_llm(payload.sample_payload)
fallback_mapping = build_heuristic_mapping(redacted_sample, payload.target_schema)
ai_error: str | None = None
mapping = fallback_mapping
generated_by = "heuristic"
if payload.use_ai:
try:
response = await ai_client.analyze(
SituationalAnalysisRequest(
title=f"Generate datasource mapping for {schema.key}",
objective=(
"Return only JSON for a deterministic mapping DSL. "
"The JSON must contain source.items_path and fields. "
"Do not include prose or code."
),
context={
"target_schema": schema.to_dict(),
"sample_payload": redacted_sample,
"mapping_dsl_example": fallback_mapping,
},
observations=[
"Use JSONPath-like paths beginning with $.",
"Never generate executable code.",
"Use field types from the target schema.",
],
constraints=[
"Return a single JSON object.",
"Do not include credentials or secrets.",
"Mark uncertain optional fields with default null.",
],
)
)
parsed = _parse_mapping_from_ai_text(response.content)
if parsed:
mapping = parsed
generated_by = "ai_provider"
else:
ai_error = "AI provider did not return a valid mapping JSON object."
except HTTPException as exc:
ai_error = str(exc.detail)
mapping.setdefault("meta", {})
if isinstance(mapping["meta"], dict):
mapping["meta"].update(
{
"name": name,
"default_url": yaml_url,
"endpoint": db_config.endpoint if db_config else yaml_url,
"is_overridden": db_config is not None and db_config.endpoint != yaml_url
if yaml_url
else db_config is not None,
"is_active": db_config.is_active if db_config else True,
"source_type": db_config.source_type if db_config else "http",
"description": db_config.description
if db_config
else f"Data source from YAML: {yaml_key}",
"generated_by": generated_by,
"requires_review": True,
"ai_error": ai_error,
}
)
return {"total": len(result), "data": result}
return {
"target_schema": schema.to_dict(),
"mapping_json": mapping,
"sample_payload_hash": stable_payload_hash(payload.sample_payload),
"redacted_sample_payload": redacted_sample,
}
@router.post("/mappings/preview")
async def preview_datasource_mapping(
payload: MappingPreviewRequest,
current_user: User = Depends(get_current_user),
):
"""Preview deterministic mapping output for a sample payload."""
try:
preview = execute_mapping(
payload.sample_payload,
payload.mapping_json,
payload.target_schema,
limit=payload.limit,
)
except (MappingError, ValueError) as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return {
"success": preview["failed_count"] == 0,
"preview": preview,
"sample_payload_hash": stable_payload_hash(payload.sample_payload),
}
@router.get("/mappings")
async def list_datasource_mappings(
datasource_config_id: Optional[int] = None,
active_only: bool = False,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""List saved mapping templates."""
query = select(DataSourceMappingTemplate).order_by(
DataSourceMappingTemplate.datasource_config_id,
DataSourceMappingTemplate.version.desc(),
)
if datasource_config_id is not None:
query = query.where(DataSourceMappingTemplate.datasource_config_id == datasource_config_id)
if active_only:
query = query.where(DataSourceMappingTemplate.is_active.is_(True))
result = await db.execute(query)
mappings = result.scalars().all()
return {"total": len(mappings), "data": [serialize_mapping_template(item) for item in mappings]}
@router.post("/mappings")
async def create_datasource_mapping(
payload: MappingTemplateCreate,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""Save a mapping template for a datasource config."""
get_target_schema(payload.target_schema)
datasource = await db.get(DataSourceConfig, payload.datasource_config_id)
if not datasource:
raise HTTPException(status_code=404, detail="Configuration not found")
if payload.sample_payload is not None:
try:
execute_mapping(payload.sample_payload, payload.mapping_json, payload.target_schema, limit=100)
except (MappingError, ValueError) as exc:
raise HTTPException(status_code=400, detail=f"Mapping validation failed: {exc}") from exc
result = await db.execute(
select(func.max(DataSourceMappingTemplate.version)).where(
DataSourceMappingTemplate.datasource_config_id == payload.datasource_config_id,
DataSourceMappingTemplate.target_schema == payload.target_schema,
)
)
next_version = int(result.scalar() or 0) + 1
if payload.is_active:
await db.execute(
DataSourceMappingTemplate.__table__.update()
.where(DataSourceMappingTemplate.datasource_config_id == payload.datasource_config_id)
.values(is_active=False)
)
template = DataSourceMappingTemplate(
datasource_config_id=payload.datasource_config_id,
target_schema=payload.target_schema,
mapping_json=payload.mapping_json,
sample_payload_hash=payload.sample_payload_hash
or (stable_payload_hash(payload.sample_payload) if payload.sample_payload is not None else None),
validation_status=payload.validation_status,
version=next_version,
is_active=payload.is_active,
)
db.add(template)
await db.commit()
await db.refresh(template)
return {"message": "Mapping template saved successfully", "data": serialize_mapping_template(template)}
@router.put("/mappings/{mapping_id}")
async def update_datasource_mapping(
mapping_id: int,
payload: MappingTemplateUpdate,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""Update a mapping template in place."""
template = await db.get(DataSourceMappingTemplate, mapping_id)
if not template:
raise HTTPException(status_code=404, detail="Mapping template not found")
target_schema = payload.target_schema or template.target_schema
mapping_json = payload.mapping_json or template.mapping_json
get_target_schema(target_schema)
if payload.sample_payload is not None:
try:
execute_mapping(payload.sample_payload, mapping_json, target_schema, limit=100)
except (MappingError, ValueError) as exc:
raise HTTPException(status_code=400, detail=f"Mapping validation failed: {exc}") from exc
if payload.is_active is True:
await db.execute(
DataSourceMappingTemplate.__table__.update()
.where(DataSourceMappingTemplate.datasource_config_id == template.datasource_config_id)
.where(DataSourceMappingTemplate.id != template.id)
.values(is_active=False)
)
template.target_schema = target_schema
template.mapping_json = mapping_json
if payload.sample_payload_hash is not None:
template.sample_payload_hash = payload.sample_payload_hash
elif payload.sample_payload is not None:
template.sample_payload_hash = stable_payload_hash(payload.sample_payload)
if payload.validation_status is not None:
template.validation_status = payload.validation_status
if payload.is_active is not None:
template.is_active = payload.is_active
await db.commit()
await db.refresh(template)
return {"message": "Mapping template updated successfully", "data": serialize_mapping_template(template)}
@router.post("/{config_id}/run-mapped")
async def run_mapped_datasource(
config_id: int,
background: bool = Query(False, description="For WebSocket sources, start a background stream task."),
debug_max_messages: int | None = Query(None, ge=1),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""Run a saved custom datasource through its active deterministic mapping."""
datasource = await db.get(DataSourceConfig, config_id)
if not datasource:
raise HTTPException(status_code=404, detail="Configuration not found")
try:
if str(datasource.source_type or "").lower() in {"websocket", "ws"}:
if background and debug_max_messages is None:
started = start_custom_stream(config_id)
if not started:
raise HTTPException(status_code=409, detail="Custom WebSocket source is already running")
return {
"status": "started",
"datasource_config_id": config_id,
"stream": get_custom_stream_status(config_id),
}
return await run_mapped_websocket_config(
db,
datasource,
debug_max_messages=debug_max_messages,
)
return await run_mapped_rest_config(db, datasource)
except httpx.HTTPStatusError as exc:
raise HTTPException(
status_code=exc.response.status_code,
detail=f"Datasource request failed: HTTP {exc.response.status_code}",
) from exc
except httpx.HTTPError as exc:
raise HTTPException(status_code=502, detail=f"Datasource request failed: {exc}") from exc
except (CustomDatasourceRuntimeError, MappingError, ValueError) as exc:
raise HTTPException(status_code=400, detail=f"Mapping failed: {exc}") from exc
@router.post("/{config_id}/stop-mapped")
async def stop_mapped_datasource(
config_id: int,
current_user: User = Depends(get_current_user),
):
stopped = await stop_custom_stream(config_id)
return {
"status": "stopped" if stopped else "not_running",
"datasource_config_id": config_id,
"stream": get_custom_stream_status(config_id),
}
@router.get("/{config_id}/stream-status")
async def get_mapped_stream_status(
config_id: int,
current_user: User = Depends(get_current_user),
):
return get_custom_stream_status(config_id)

View File

@@ -9,6 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.core.time import to_iso8601_utc
from app.core.security import get_current_user
from app.core.data_sources import get_data_sources_config
from app.core.datasource_defaults import DEFAULT_DATASOURCES
from app.db.session import get_db
from app.models.collected_data import CollectedData
from app.models.data_snapshot import DataSnapshot
@@ -35,6 +36,17 @@ def format_frequency_label(minutes: int) -> str:
return f"{minutes}m"
def datasource_metadata(source: str) -> dict:
info = DEFAULT_DATASOURCES.get(source, {})
return {
"display_name": info.get("display_name") or info.get("name") or source,
"is_free": bool(info.get("is_free", True)),
"requires_credentials": bool(info.get("requires_credentials", False)),
"credential_provider": info.get("credential_provider"),
"credential_status": info.get("credential_status", "none"),
}
def is_due_for_collection(datasource: DataSource, now: datetime) -> bool:
if datasource.last_run_at is None:
return True
@@ -72,31 +84,6 @@ async def _load_latest_running_tasks(
return {task.datasource_id: task for task in result.scalars().all()}
async def _load_latest_completed_tasks(
db: AsyncSession,
datasource_ids: list[int],
) -> dict[int, CollectionTask]:
if not datasource_ids:
return {}
ranked_tasks = (
select(
CollectionTask.id.label("task_id"),
_task_rank_column(CollectionTask.completed_at),
)
.where(CollectionTask.datasource_id.in_(datasource_ids))
.where(CollectionTask.completed_at.isnot(None))
.where(CollectionTask.status.in_(("success", "failed", "cancelled")))
.subquery()
)
result = await db.execute(
select(CollectionTask)
.join(ranked_tasks, CollectionTask.id == ranked_tasks.c.task_id)
.where(ranked_tasks.c.row_num == 1)
)
return {task.datasource_id: task for task in result.scalars().all()}
async def _load_latest_task_ids(
db: AsyncSession,
datasource_ids: list[int],
@@ -123,21 +110,6 @@ async def _load_latest_task_ids(
return {datasource_id: task_id for datasource_id, task_id in result.all()}
async def _load_datasource_data_counts(
db: AsyncSession,
sources: list[str],
) -> dict[str, int]:
if not sources:
return {}
result = await db.execute(
select(CollectedData.source, func.count(CollectedData.id))
.where(CollectedData.source.in_(sources))
.group_by(CollectedData.source)
)
return {source: count for source, count in result.all()}
async def _load_datasource_endpoint_overrides(
db: AsyncSession,
sources: list[str],
@@ -161,7 +133,7 @@ async def _load_datasource_endpoint_overrides(
async def _load_datasource_list_context(
db: AsyncSession,
datasources: list[DataSource],
) -> tuple[dict[int, CollectionTask], dict[int, CollectionTask], dict[str, int], dict[str, str]]:
) -> tuple[dict[int, CollectionTask], dict[str, str]]:
datasource_ids = [datasource.id for datasource in datasources]
sources = [datasource.source for datasource in datasources]
@@ -185,10 +157,8 @@ async def _load_datasource_list_context(
if stale_datasource_ids:
running_tasks = await _load_latest_running_tasks(db, datasource_ids)
completed_tasks = await _load_latest_completed_tasks(db, datasource_ids)
data_counts = await _load_datasource_data_counts(db, sources)
endpoint_overrides = await _load_datasource_endpoint_overrides(db, sources)
return running_tasks, completed_tasks, data_counts, endpoint_overrides
return running_tasks, endpoint_overrides
async def get_datasource_record(db: AsyncSession, source_id: str) -> Optional[DataSource]:
@@ -401,27 +371,19 @@ async def list_datasources(
collector_list = []
config = get_data_sources_config()
running_tasks, completed_tasks, data_counts, endpoint_overrides = await _load_datasource_list_context(
db,
datasources,
)
running_tasks, endpoint_overrides = await _load_datasource_list_context(db, datasources)
for datasource in datasources:
running_task = running_tasks.get(datasource.id)
last_task = completed_tasks.get(datasource.id)
endpoint = endpoint_overrides.get(datasource.source) or config.get_yaml_url(
datasource.source,
)
data_count = data_counts.get(datasource.source, 0)
last_run_at = datasource.last_run_at or (last_task.completed_at if last_task else None)
last_run = to_iso8601_utc(last_run_at)
last_status = datasource.last_status or (last_task.status if last_task else None)
endpoint = endpoint_overrides.get(datasource.source) or config.get_yaml_url(datasource.source)
last_run_at = datasource.last_run_at
last_status = datasource.last_status
collector_list.append(
{
"id": datasource.id,
"source": datasource.source,
"name": datasource.name,
**datasource_metadata(datasource.source),
"module": datasource.module,
"priority": datasource.priority,
"frequency": format_frequency_label(datasource.frequency_minutes),
@@ -429,15 +391,18 @@ async def list_datasources(
"is_active": datasource.is_active,
"collector_class": datasource.collector_class,
"endpoint": endpoint,
"last_run": last_run,
"last_run": to_iso8601_utc(last_run_at),
"last_run_at": to_iso8601_utc(last_run_at),
"last_status": last_status,
"last_records_processed": last_task.records_processed if last_task else None,
"data_count": data_count,
"is_running": running_task is not None,
"task_id": running_task.id if running_task else None,
"progress": running_task.progress if running_task else None,
"phase": running_task.phase if running_task else None,
"phase_progress": running_task.phase_progress if running_task else None,
"phase_message": running_task.phase_message if running_task else None,
"phase_current": running_task.phase_current if running_task else None,
"phase_total": running_task.phase_total if running_task else None,
"phase_unit": running_task.phase_unit if running_task else None,
"records_processed": running_task.records_processed if running_task else None,
"total_records": running_task.total_records if running_task else None,
}
@@ -576,6 +541,7 @@ async def get_datasource(
return {
"id": datasource.id,
"name": datasource.name,
**datasource_metadata(datasource.source),
"module": datasource.module,
"priority": datasource.priority,
"frequency": format_frequency_label(datasource.frequency_minutes),
@@ -665,6 +631,11 @@ async def trigger_datasource(
"message": "当前采集任务尚未完成,重新触发会丢失本次未完成进度。是否强制重新采集?",
"task_id": running_task.id,
"phase": running_task.phase,
"phase_progress": running_task.phase_progress,
"phase_message": running_task.phase_message,
"phase_current": running_task.phase_current,
"phase_total": running_task.phase_total,
"phase_unit": running_task.phase_unit,
"progress": running_task.progress,
"records_processed": running_task.records_processed,
"total_records": running_task.total_records,
@@ -748,13 +719,29 @@ async def get_task_status(
task = await get_running_task(db, datasource.id)
if not task:
return {"is_running": False, "task_id": None, "progress": None, "phase": None, "status": "idle"}
return {
"is_running": False,
"task_id": None,
"progress": None,
"phase": None,
"phase_progress": None,
"phase_message": None,
"phase_current": None,
"phase_total": None,
"phase_unit": None,
"status": "idle",
}
return {
"is_running": task.status == "running",
"task_id": task.id,
"progress": task.progress,
"phase": task.phase,
"phase_progress": task.phase_progress,
"phase_message": task.phase_message,
"phase_current": task.phase_current,
"phase_total": task.phase_total,
"phase_unit": task.phase_unit,
"records_processed": task.records_processed,
"total_records": task.total_records,
"status": task.status,

102
backend/app/api/v1/docs.py Normal file
View File

@@ -0,0 +1,102 @@
"""Authenticated documentation APIs."""
from __future__ import annotations
from fastapi import APIRouter, Depends, HTTPException, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from sqlalchemy import text
from app.core.security import decode_token
from app.db.session import async_session_factory
from app.models.user import User
from app.services.docs_gatekeeper import (
DOCS_BY_SLUG,
VALID_DOCS_LANGS,
can_read_doc,
catalog_for_user,
doc_path_for,
title_for,
)
router = APIRouter()
optional_bearer = HTTPBearer(auto_error=False)
async def get_optional_current_user(
credentials: HTTPAuthorizationCredentials | None = Depends(optional_bearer),
) -> User | None:
if credentials is None:
return None
payload = decode_token(credentials.credentials)
if payload is None or payload.get("type") != "access" or payload.get("sub") is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid token",
)
async with async_session_factory() as db:
result = await db.execute(
text(
"SELECT id, username, email, password_hash, role, is_active, gatekeeper_groups FROM users WHERE id = :id"
),
{"id": int(payload["sub"])},
)
row = result.fetchone()
if row is None or not row[5]:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="User not found or inactive",
)
user = User()
user.id = row[0]
user.username = row[1]
user.email = row[2]
user.password_hash = row[3]
user.role = row[4]
user.is_active = row[5]
user.gatekeeper_groups = row[6] or []
return user
@router.get("/catalog")
async def get_docs_catalog(current_user: User | None = Depends(get_optional_current_user)):
return {
"items": catalog_for_user(current_user),
"authenticated": current_user is not None,
}
@router.get("/{lang}/{slug}")
async def get_doc_content(
lang: str,
slug: str,
current_user: User | None = Depends(get_optional_current_user),
):
if lang not in VALID_DOCS_LANGS:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Document not found")
entry = DOCS_BY_SLUG.get(slug)
if entry is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Document not found")
path = doc_path_for(entry, lang)
if not path.exists():
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Document not found")
if not can_read_doc(entry, current_user):
if current_user is None:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Authentication required")
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Insufficient Docs permissions")
return {
"slug": entry.slug,
"filename": entry.filename,
"lang": lang,
"title": title_for(entry, lang),
"group": entry.group,
"order": entry.order,
"access": entry.access,
"markdown": path.read_text(encoding="utf-8"),
}

File diff suppressed because it is too large Load Diff

View File

@@ -27,7 +27,9 @@ async def list_tasks(
offset = (page - 1) * page_size
query = """
SELECT ct.id, ct.datasource_id, ds.name as datasource_name, ct.status,
ct.started_at, ct.completed_at, ct.records_processed, ct.error_message
ct.started_at, ct.completed_at, ct.records_processed, ct.error_message,
ct.phase, ct.phase_progress, ct.phase_message, ct.phase_current,
ct.phase_total, ct.phase_unit, ct.total_records, ct.progress
FROM collection_tasks ct
JOIN data_sources ds ON ct.datasource_id = ds.id
WHERE 1=1
@@ -66,6 +68,14 @@ async def list_tasks(
"completed_at": to_iso8601_utc(t[5]),
"records_processed": t[6],
"error_message": t[7],
"phase": t[8],
"phase_progress": t[9],
"phase_message": t[10],
"phase_current": t[11],
"phase_total": t[12],
"phase_unit": t[13],
"total_records": t[14],
"progress": t[15],
}
for t in tasks
],
@@ -81,7 +91,9 @@ async def get_task(
result = await db.execute(
text("""
SELECT ct.id, ct.datasource_id, ds.name as datasource_name, ct.status,
ct.started_at, ct.completed_at, ct.records_processed, ct.error_message
ct.started_at, ct.completed_at, ct.records_processed, ct.error_message,
ct.phase, ct.phase_progress, ct.phase_message, ct.phase_current,
ct.phase_total, ct.phase_unit, ct.total_records, ct.progress
FROM collection_tasks ct
JOIN data_sources ds ON ct.datasource_id = ds.id
WHERE ct.id = :id
@@ -105,6 +117,14 @@ async def get_task(
"completed_at": to_iso8601_utc(task[5]),
"records_processed": task[6],
"error_message": task[7],
"phase": task[8],
"phase_progress": task[9],
"phase_message": task[10],
"phase_current": task[11],
"phase_total": task[12],
"phase_unit": task[13],
"total_records": task[14],
"progress": task[15],
}

View File

@@ -1,3 +1,4 @@
import json
from typing import List
from fastapi import APIRouter, Depends, HTTPException, status
@@ -7,10 +8,12 @@ from sqlalchemy import text
from app.core.security import get_current_user, get_password_hash
from app.db.session import get_db
from app.models.user import User
from app.schemas.user import UserCreate, UserResponse, UserUpdate
from app.schemas.user import UserCreate, UserUpdate
router = APIRouter()
VALID_GATEKEEPER_GROUPS = {"docs_user", "docs_developer", "docs_admin"}
def check_permission(current_user: User, required_roles: List[str]) -> bool:
user_role_value = (
@@ -52,7 +55,7 @@ async def list_users(
offset = (page - 1) * page_size
query = text(
f"SELECT id, username, email, role, is_active, last_login_at, created_at FROM users WHERE {where_sql} ORDER BY created_at DESC LIMIT {page_size} OFFSET {offset}"
f"SELECT id, username, email, role, is_active, last_login_at, created_at, gatekeeper_groups FROM users WHERE {where_sql} ORDER BY created_at DESC LIMIT {page_size} OFFSET {offset}"
)
count_query = text(f"SELECT COUNT(*) FROM users WHERE {where_sql}")
@@ -75,6 +78,7 @@ async def list_users(
"is_active": u[4],
"last_login_at": u[5],
"created_at": u[6],
"gatekeeper_groups": u[7] or [],
}
for u in users
],
@@ -95,7 +99,7 @@ async def get_user(
result = await db.execute(
text(
"SELECT id, username, email, role, is_active, last_login_at, created_at FROM users WHERE id = :id"
"SELECT id, username, email, role, is_active, last_login_at, created_at, gatekeeper_groups FROM users WHERE id = :id"
),
{"id": user_id},
)
@@ -114,6 +118,7 @@ async def get_user(
"is_active": user[4],
"last_login_at": user[5],
"created_at": user[6],
"gatekeeper_groups": user[7] or [],
}
@@ -128,6 +133,12 @@ async def create_user(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only super_admin can create users",
)
invalid_groups = sorted(set(user_data.gatekeeper_groups) - VALID_GATEKEEPER_GROUPS)
if invalid_groups:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unsupported Gatekeeper groups: {', '.join(invalid_groups)}",
)
result = await db.execute(
text("SELECT id FROM users WHERE username = :username OR email = :email"),
@@ -142,13 +153,14 @@ async def create_user(
hashed_password = get_password_hash(user_data.password)
await db.execute(
text("""INSERT INTO users (username, email, password_hash, role, is_active, created_at, updated_at)
VALUES (:username, :email, :password_hash, :role, :is_active, NOW(), NOW())"""),
text("""INSERT INTO users (username, email, password_hash, role, gatekeeper_groups, is_active, created_at, updated_at)
VALUES (:username, :email, :password_hash, :role, CAST(:gatekeeper_groups AS jsonb), :is_active, NOW(), NOW())"""),
{
"username": user_data.username,
"email": user_data.email,
"password_hash": hashed_password,
"role": user_data.role,
"gatekeeper_groups": json.dumps(user_data.gatekeeper_groups),
"is_active": True,
},
)
@@ -172,6 +184,7 @@ async def create_user(
"username": user_data.username,
"email": user_data.email,
"role": user_data.role,
"gatekeeper_groups": user_data.gatekeeper_groups,
"is_active": True,
}
@@ -194,6 +207,18 @@ async def update_user(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only super_admin can change user role",
)
if not check_permission(current_user, ["super_admin"]) and user_data.gatekeeper_groups is not None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only super_admin can change Gatekeeper groups",
)
if user_data.gatekeeper_groups is not None:
invalid_groups = sorted(set(user_data.gatekeeper_groups) - VALID_GATEKEEPER_GROUPS)
if invalid_groups:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unsupported Gatekeeper groups: {', '.join(invalid_groups)}",
)
result = await db.execute(
text("SELECT id FROM users WHERE id = :id"),
@@ -213,6 +238,9 @@ async def update_user(
if user_data.role is not None:
update_fields.append("role = :role")
params["role"] = user_data.role
if user_data.gatekeeper_groups is not None:
update_fields.append("gatekeeper_groups = CAST(:gatekeeper_groups AS jsonb)")
params["gatekeeper_groups"] = json.dumps(user_data.gatekeeper_groups)
if user_data.is_active is not None:
update_fields.append("is_active = :is_active")
params["is_active"] = user_data.is_active

View File

@@ -0,0 +1,132 @@
"""v4 strategy + v5 conflict-promotion + enrichment APIs for vessel_ais."""
from __future__ import annotations
from typing import Any
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.security import get_current_user
from app.db.session import get_db
from app.models.user import User
from app.models.vessel import AISConflictRecord
from app.services.vessel_aggregation_strategy import (
StrategyValidationError,
load_strategy,
reset_strategy,
save_strategy,
)
from app.services.vessel_enrichment import (
get_vessel_enrichment_bundle,
upsert_vessel_media_enrichment,
upsert_vessel_profile_enrichment,
)
router = APIRouter()
@router.get("/strategy")
async def get_aggregation_strategy(db: AsyncSession = Depends(get_db)):
return await load_strategy(db)
@router.put("/strategy")
async def put_aggregation_strategy(
payload: dict[str, Any],
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
return await save_strategy(db, payload)
except StrategyValidationError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@router.delete("/strategy")
async def reset_aggregation_strategy(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await reset_strategy(db)
@router.post("/conflicts/{mmsi}/{field}/promote-to-rule")
async def promote_conflict_to_rule(
mmsi: int,
field: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""Lift the current conflict resolution into a persistent strategy rule."""
result = await db.execute(
select(AISConflictRecord)
.where(AISConflictRecord.target_schema == "vessel_ais")
.where(AISConflictRecord.entity_key == str(mmsi))
.where(AISConflictRecord.field == field)
.order_by(AISConflictRecord.updated_at.desc(), AISConflictRecord.id.desc())
.limit(1)
)
record = result.scalar_one_or_none()
if record is None or not record.selected_source:
raise HTTPException(status_code=404, detail="Conflict record with selected_source not found")
strategy = await load_strategy(db)
vessel_ais = dict(strategy.get("vessel_ais") or {})
field_rules = dict(vessel_ais.get("field_rules") or {})
field_rules[field] = {"mode": "source_priority", "source_priority": [record.selected_source]}
vessel_ais["field_rules"] = field_rules
incoming = {"version": int(strategy.get("version") or 0), "vessel_ais": vessel_ais}
try:
return await save_strategy(db, incoming)
except StrategyValidationError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@router.delete("/conflicts/{mmsi}/{field}/promote-to-rule")
async def revert_conflict_rule(
mmsi: int,
field: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
strategy = await load_strategy(db)
vessel_ais = dict(strategy.get("vessel_ais") or {})
field_rules = dict(vessel_ais.get("field_rules") or {})
if field in field_rules:
del field_rules[field]
vessel_ais["field_rules"] = field_rules
incoming = {"version": int(strategy.get("version") or 0), "vessel_ais": vessel_ais}
try:
return await save_strategy(db, incoming)
except StrategyValidationError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@router.get("/enrichment/{mmsi}")
async def get_vessel_enrichment(mmsi: int, db: AsyncSession = Depends(get_db)):
return await get_vessel_enrichment_bundle(db, mmsi)
@router.put("/enrichment/{mmsi}/profile")
async def put_vessel_profile_enrichment(
mmsi: int,
payload: dict[str, Any],
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await upsert_vessel_profile_enrichment(db, mmsi=mmsi, payload=payload)
@router.put("/enrichment/{mmsi}/media")
async def put_vessel_media_enrichment(
mmsi: int,
payload: dict[str, Any],
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await upsert_vessel_media_enrichment(db, mmsi=mmsi, payload=payload)

File diff suppressed because it is too large Load Diff

View File

@@ -40,16 +40,16 @@ async def authenticate_token(token: str) -> Optional[dict]:
@router.websocket("/ws")
async def websocket_endpoint(
websocket: WebSocket,
token: str = Query(...),
token: str | None = Query(None),
):
"""WebSocket endpoint for real-time data"""
logger.info_event(
"WebSocket connection attempt",
event="auth.websocket.connection_attempt",
context={"token_preview": f"{token[:8]}..."},
context={"token_preview": f"{token[:8]}..." if token else "anonymous"},
)
payload = await authenticate_token(token)
if payload is None:
payload = await authenticate_token(token) if token else None
if token and payload is None:
logger.warning_event(
"WebSocket authentication failed, closing connection",
event="auth.websocket.connection_rejected",
@@ -57,7 +57,17 @@ async def websocket_endpoint(
await websocket.close(code=4001)
return
user_id = str(payload.get("sub"))
is_anonymous = payload is None
user_id = str(payload.get("sub")) if payload else f"anonymous:{id(websocket)}"
supported_channels = ["vessels"] if is_anonymous else [
"gpu_clusters",
"submarine_cables",
"ixp_nodes",
"alerts",
"dashboard",
"datasource_tasks",
"vessels",
]
await manager.connect(websocket, user_id)
try:
@@ -68,14 +78,7 @@ async def websocket_endpoint(
"connection_id": f"conn_{user_id}",
"server_version": settings.VERSION,
"heartbeat_interval": 30,
"supported_channels": [
"gpu_clusters",
"submarine_cables",
"ixp_nodes",
"alerts",
"dashboard",
"datasource_tasks",
],
"supported_channels": supported_channels,
},
}
)
@@ -93,12 +96,24 @@ async def websocket_endpoint(
)
elif data.get("type") == "subscribe":
channels = data.get("data", {}).get("channels", [])
if is_anonymous:
channels = [channel for channel in channels if channel in supported_channels]
manager.subscribe(websocket, channels)
await websocket.send_json(
{
"type": "subscription_confirmed",
"data": {"action": "subscribe", "channels": channels},
}
)
elif data.get("type") == "unsubscribe":
channels = data.get("data", {}).get("channels", [])
manager.unsubscribe(websocket, channels)
await websocket.send_json(
{
"type": "subscription_confirmed",
"data": {"action": "unsubscribe", "channels": channels},
}
)
elif data.get("type") == "control_frame":
await websocket.send_json(
{"type": "control_acknowledged", "data": {"received": True}}

View File

@@ -4,8 +4,8 @@ from typing import Any, Dict, Optional
FIELD_ALIASES = {
"country": ("country",),
"city": ("city",),
"latitude": ("latitude",),
"longitude": ("longitude",),
"latitude": ("latitude", "lat"),
"longitude": ("longitude", "lon", "lng"),
"value": ("value",),
"unit": ("unit",),
"cores": ("cores",),
@@ -14,6 +14,28 @@ FIELD_ALIASES = {
"power": ("power",),
}
NESTED_FIELD_ALIASES = {
"latitude": (
("location", "latitude"),
("location", "lat"),
("geo", "latitude"),
("geo", "lat"),
("coordinates", "latitude"),
("coordinates", "lat"),
),
"longitude": (
("location", "longitude"),
("location", "lon"),
("location", "lng"),
("geo", "longitude"),
("geo", "lon"),
("geo", "lng"),
("coordinates", "longitude"),
("coordinates", "lon"),
("coordinates", "lng"),
),
}
def get_metadata_field(metadata: Optional[Dict[str, Any]], field: str, fallback: Any = None) -> Any:
if isinstance(metadata, dict):
@@ -21,9 +43,34 @@ def get_metadata_field(metadata: Optional[Dict[str, Any]], field: str, fallback:
value = metadata.get(key)
if value not in (None, ""):
return value
for path in NESTED_FIELD_ALIASES.get(field, ()):
current: Any = metadata
for key in path:
if not isinstance(current, dict):
current = None
break
current = current.get(key)
if current not in (None, ""):
return current
if field in {"latitude", "longitude"}:
value = _get_coordinate_sequence_value(metadata, field)
if value not in (None, ""):
return value
return fallback
def _get_coordinate_sequence_value(metadata: Dict[str, Any], field: str) -> Any:
for key in ("coordinates", "coord", "coords"):
value = metadata.get(key)
if not isinstance(value, (list, tuple)) or len(value) < 2:
continue
# GeoJSON uses [longitude, latitude]. Most raw collector tuples in this
# codebase use explicit field names, so only sequence aliases are treated
# as GeoJSON-shaped to avoid guessing.
return value[1] if field == "latitude" else value[0]
return None
def build_dynamic_metadata(
metadata: Optional[Dict[str, Any]],
*,

View File

@@ -1,7 +1,6 @@
import os
import yaml
from functools import lru_cache
from typing import Optional
COLLECTOR_URL_KEYS = {
@@ -31,6 +30,8 @@ COLLECTOR_URL_KEYS = {
"opengeofeed_prefix_geo": "opengeofeed.public_csv_url",
"nro_delegated_prefix_geo": "nro.delegated_stats_url",
"news_live_streams": "news_live_streams.channels_url",
"barentswatch_vessels": "barentswatch_vessels.url",
"aisstream_vessels": "aisstream_vessels.url",
}
@@ -73,7 +74,7 @@ class DataSourcesConfig:
from app.models.datasource_config import DataSourceConfig
query = select(DataSourceConfig).where(
DataSourceConfig.name == collector_name, DataSourceConfig.is_active == True
DataSourceConfig.name == collector_name, DataSourceConfig.is_active
)
result = await db.execute(query)
db_config = result.scalar_one_or_none()

View File

@@ -94,3 +94,11 @@ news_live_streams:
streams_url: "https://iptv-org.github.io/api/streams.json"
# IPTV-org 台标 JSON
logos_url: "https://iptv-org.github.io/api/logos.json"
barentswatch_vessels:
# BarentsWatch Live AIS latest combined endpoint. Requires an AIS bearer token.
url: "https://live.ais.barentswatch.no/v1/latest/combined"
aisstream_vessels:
# AISStream realtime WebSocket endpoint. Requires an AISStream API key.
url: "wss://stream.aisstream.io/v0/stream"

View File

@@ -4,163 +4,258 @@ DEFAULT_DATASOURCES = {
"top500": {
"id": 1,
"name": "TOP500 Supercomputers",
"display_name": "TOP500 超算榜单",
"module": "L1",
"priority": "P0",
"frequency_minutes": 240,
"is_free": True,
"requires_credentials": False,
},
"epoch_ai_gpu": {
"id": 2,
"name": "Epoch AI GPU Clusters",
"display_name": "Epoch AI GPU 集群",
"module": "L1",
"priority": "P0",
"frequency_minutes": 360,
"is_free": True,
"requires_credentials": False,
},
"huggingface_models": {
"id": 3,
"name": "HuggingFace Models",
"display_name": "Hugging Face 模型",
"module": "L2",
"priority": "P1",
"frequency_minutes": 720,
"is_free": True,
"requires_credentials": False,
},
"huggingface_datasets": {
"id": 4,
"name": "HuggingFace Datasets",
"display_name": "Hugging Face 数据集",
"module": "L2",
"priority": "P1",
"frequency_minutes": 720,
"is_free": True,
"requires_credentials": False,
},
"huggingface_spaces": {
"id": 5,
"name": "HuggingFace Spaces",
"display_name": "Hugging Face Spaces",
"module": "L2",
"priority": "P2",
"frequency_minutes": 1440,
"is_free": True,
"requires_credentials": False,
},
"peeringdb_ixp": {
"id": 6,
"name": "PeeringDB IXP",
"display_name": "PeeringDB 交换中心",
"module": "L2",
"priority": "P1",
"frequency_minutes": 1440,
"is_free": True,
"requires_credentials": False,
},
"peeringdb_network": {
"id": 7,
"name": "PeeringDB Networks",
"display_name": "PeeringDB 网络",
"module": "L2",
"priority": "P2",
"frequency_minutes": 2880,
"is_free": True,
"requires_credentials": False,
},
"peeringdb_facility": {
"id": 8,
"name": "PeeringDB Facilities",
"display_name": "PeeringDB 设施",
"module": "L2",
"priority": "P2",
"frequency_minutes": 2880,
"is_free": True,
"requires_credentials": False,
},
"telegeography_cables": {
"id": 9,
"name": "Submarine Cables",
"display_name": "海底光缆",
"module": "L2",
"priority": "P1",
"frequency_minutes": 10080,
"is_free": True,
"requires_credentials": False,
},
"telegeography_landing": {
"id": 10,
"name": "Cable Landing Points",
"display_name": "光缆登陆点",
"module": "L2",
"priority": "P2",
"frequency_minutes": 10080,
"is_free": True,
"requires_credentials": False,
},
"telegeography_systems": {
"id": 11,
"name": "Cable Systems",
"display_name": "光缆系统",
"module": "L2",
"priority": "P2",
"frequency_minutes": 10080,
"is_free": True,
"requires_credentials": False,
},
"arcgis_cables": {
"id": 15,
"name": "ArcGIS Submarine Cables",
"display_name": "ArcGIS 海底光缆",
"module": "L2",
"priority": "P1",
"frequency_minutes": 10080,
"is_free": True,
"requires_credentials": False,
},
"arcgis_landing_points": {
"id": 16,
"name": "ArcGIS Landing Points",
"display_name": "ArcGIS 登陆点",
"module": "L2",
"priority": "P1",
"frequency_minutes": 10080,
"is_free": True,
"requires_credentials": False,
},
"arcgis_cable_landing_relation": {
"id": 17,
"name": "ArcGIS Cable-Landing Relations",
"display_name": "ArcGIS 光缆登陆关系",
"module": "L2",
"priority": "P1",
"frequency_minutes": 10080,
"is_free": True,
"requires_credentials": False,
},
"fao_landing_points": {
"id": 18,
"name": "FAO Landing Points",
"display_name": "FAO 登陆点",
"module": "L2",
"priority": "P1",
"frequency_minutes": 10080,
"is_free": True,
"requires_credentials": False,
},
"spacetrack_tle": {
"id": 19,
"name": "Space-Track TLE",
"display_name": "Space-Track 轨道根数",
"module": "L3",
"priority": "P2",
"frequency_minutes": 1440,
"is_free": True,
"requires_credentials": True,
"credential_provider": "spacetrack",
"credential_status": "planned",
},
"celestrak_tle": {
"id": 20,
"name": "CelesTrak TLE",
"display_name": "CelesTrak 轨道根数",
"module": "L3",
"priority": "P2",
"frequency_minutes": 1440,
"is_free": True,
"requires_credentials": False,
},
"ris_live_bgp": {
"id": 21,
"name": "RIPE RIS Live BGP",
"display_name": "RIPE RIS 实时 BGP",
"module": "L3",
"priority": "P1",
"frequency_minutes": 15,
"is_free": True,
"requires_credentials": False,
},
"bgpstream_bgp": {
"id": 22,
"name": "CAIDA BGPStream Backfill",
"display_name": "CAIDA BGPStream 回填",
"module": "L3",
"priority": "P1",
"frequency_minutes": 360,
"is_free": True,
"requires_credentials": False,
},
"iptoasn_prefix_geo": {
"id": 23,
"name": "IPtoASN Prefix Geography",
"display_name": "IPtoASN 前缀地理",
"module": "L3",
"priority": "P1",
"frequency_minutes": 1440,
"is_free": True,
"requires_credentials": False,
},
"opengeofeed_prefix_geo": {
"id": 24,
"name": "OpenGeoFeed Prefix Geography",
"display_name": "OpenGeoFeed 前缀地理",
"module": "L3",
"priority": "P1",
"frequency_minutes": 1440,
"is_free": True,
"requires_credentials": False,
},
"nro_delegated_prefix_geo": {
"id": 25,
"name": "NRO Delegated Prefix Geography",
"display_name": "NRO 分配前缀地理",
"module": "L3",
"priority": "P1",
"frequency_minutes": 1440,
"is_free": True,
"requires_credentials": False,
},
"news_live_streams": {
"id": 26,
"name": "News Live Streams",
"display_name": "新闻直播源",
"module": "L4",
"priority": "P2",
"frequency_minutes": 720,
"is_free": True,
"requires_credentials": False,
},
"barentswatch_vessels": {
"id": 27,
"name": "BarentsWatch AIS Vessels",
"display_name": "BarentsWatch AIS 船舶",
"module": "L4",
"priority": "P1",
"frequency_minutes": 1,
"is_free": True,
"requires_credentials": True,
"credential_provider": "barentswatch",
"credential_status": "supported",
},
"aisstream_vessels": {
"id": 28,
"name": "AISStream Vessels",
"display_name": "AISStream 实时船舶",
"module": "L4",
"priority": "P1",
"frequency_minutes": 1,
"is_free": True,
"requires_credentials": True,
"credential_provider": "aisstream",
"credential_status": "supported",
},
}

View File

@@ -105,7 +105,7 @@ async def get_current_user(
)
result = await db.execute(
text(
"SELECT id, username, email, password_hash, role, is_active FROM users WHERE id = :id"
"SELECT id, username, email, password_hash, role, is_active, gatekeeper_groups FROM users WHERE id = :id"
),
{"id": int(user_id)},
)
@@ -122,6 +122,7 @@ async def get_current_user(
user.password_hash = row[3]
user.role = row[4]
user.is_active = row[5]
user.gatekeeper_groups = row[6] or []
return user
@@ -144,7 +145,7 @@ async def get_current_user_refresh(
)
result = await db.execute(
text(
"SELECT id, username, email, password_hash, role, is_active FROM users WHERE id = :id"
"SELECT id, username, email, password_hash, role, is_active, gatekeeper_groups FROM users WHERE id = :id"
),
{"id": int(user_id)},
)
@@ -161,6 +162,7 @@ async def get_current_user_refresh(
user.password_hash = row[3]
user.role = row[4]
user.is_active = row[5]
user.gatekeeper_groups = row[6] or []
return user

View File

@@ -0,0 +1,157 @@
"""Registry of target schemas supported by mapped custom data sources."""
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime
from typing import Any
from pydantic import BaseModel, Field, ValidationError, field_validator
class VesselAISRecord(BaseModel):
mmsi: int = Field(ge=100000000, le=999999999)
lat: float = Field(ge=-90, le=90)
lon: float = Field(ge=-180, le=180)
sog: float | None = None
cog: float | None = Field(default=None, ge=0, le=360)
heading: int | None = Field(default=None, ge=0, le=511)
nav_status: int | None = None
name: str | None = None
callsign: str | None = None
vessel_type: str | int | None = None
vessel_type_name: str | None = None
received_at: datetime | None = None
class GeoPointRecord(BaseModel):
lat: float = Field(ge=-90, le=90)
lon: float = Field(ge=-180, le=180)
name: str | None = None
type: str | None = None
source_id: str | None = None
observed_at: datetime | None = None
metadata: dict[str, Any] = Field(default_factory=dict)
class GenericRecord(BaseModel):
data: dict[str, Any] = Field(default_factory=dict)
source_id: str | None = None
observed_at: datetime | None = None
@field_validator("data")
@classmethod
def require_payload(cls, value: dict[str, Any]) -> dict[str, Any]:
if not value:
raise ValueError("generic_records requires a non-empty data object")
return value
@dataclass(frozen=True)
class TargetField:
name: str
type: str
required: bool = False
description: str = ""
example: Any = None
def to_dict(self) -> dict[str, Any]:
return {
"name": self.name,
"type": self.type,
"required": self.required,
"description": self.description,
"example": self.example,
}
@dataclass(frozen=True)
class TargetSchema:
key: str
label: str
description: str
fields: tuple[TargetField, ...]
model: type[BaseModel]
destination: str
def to_dict(self) -> dict[str, Any]:
return {
"key": self.key,
"label": self.label,
"description": self.description,
"destination": self.destination,
"fields": [field.to_dict() for field in self.fields],
}
def validate_record(self, record: dict[str, Any]) -> tuple[dict[str, Any] | None, list[str]]:
try:
return self.model.model_validate(record).model_dump(mode="json"), []
except ValidationError as exc:
return None, [
".".join(str(part) for part in error["loc"]) + f": {error['msg']}"
for error in exc.errors()
]
TARGET_SCHEMAS: dict[str, TargetSchema] = {
"vessel_ais": TargetSchema(
key="vessel_ais",
label="船舶 AIS",
description="船只位置、航速、航向、MMSI 等 AIS 数据。",
destination="vessel_position",
model=VesselAISRecord,
fields=(
TargetField("mmsi", "integer", True, "MMSI 九位船舶标识", 257123000),
TargetField("lat", "float", True, "纬度", 59.91),
TargetField("lon", "float", True, "经度", 10.75),
TargetField("sog", "float", False, "对地航速,单位节", 12.4),
TargetField("cog", "float", False, "对地航向0-360 度", 184.5),
TargetField("heading", "integer", False, "船首向0-511", 186),
TargetField("nav_status", "integer", False, "导航状态码", 0),
TargetField("name", "string", False, "船名", "OSLO EXPRESS"),
TargetField("callsign", "string", False, "呼号", "LAAB"),
TargetField("vessel_type", "string", False, "船型代码", 70),
TargetField("vessel_type_name", "string", False, "船型名称", "Cargo"),
TargetField("received_at", "datetime", False, "数据接收时间", "2026-04-28T00:00:00Z"),
),
),
"geo_points": TargetSchema(
key="geo_points",
label="通用地理点",
description="带经纬度的通用实体或事件点位。",
destination="generic_geo_points",
model=GeoPointRecord,
fields=(
TargetField("lat", "float", True, "纬度", 1.3),
TargetField("lon", "float", True, "经度", 103.8),
TargetField("name", "string", False, "点位名称", "Singapore"),
TargetField("type", "string", False, "点位类型", "datacenter"),
TargetField("source_id", "string", False, "来源侧 ID", "sg-1"),
TargetField("observed_at", "datetime", False, "观测时间", "2026-04-28T00:00:00Z"),
TargetField("metadata", "object", False, "扩展字段", {"provider": "example"}),
),
),
"generic_records": TargetSchema(
key="generic_records",
label="通用结构化记录",
description="未知结构数据沉淀,不直接进入 Earth 图层。",
destination="collected_data",
model=GenericRecord,
fields=(
TargetField("data", "object", True, "结构化记录主体", {"raw": "value"}),
TargetField("source_id", "string", False, "来源侧 ID", "record-1"),
TargetField("observed_at", "datetime", False, "观测时间", "2026-04-28T00:00:00Z"),
),
),
}
def list_target_schemas() -> list[dict[str, Any]]:
return [schema.to_dict() for schema in TARGET_SCHEMAS.values()]
def get_target_schema(key: str) -> TargetSchema:
try:
return TARGET_SCHEMAS[key]
except KeyError as exc:
raise ValueError(f"Unsupported target schema: {key}") from exc

View File

@@ -75,7 +75,7 @@ class DataBroadcaster:
"timestamp": to_iso8601_utc(datetime.now(UTC)),
"payload": data,
},
channel=channel if channel in manager.active_connections else "all",
channel=channel,
)
async def broadcast_datasource_task_update(self, data: Dict[str, Any]):

View File

@@ -1,9 +1,6 @@
"""WebSocket Connection Manager"""
import json
import asyncio
from typing import Dict, Set, Optional
from datetime import datetime
from fastapi import WebSocket
import redis.asyncio as redis
@@ -15,6 +12,8 @@ class ConnectionManager:
def __init__(self):
self.active_connections: Dict[str, Set[WebSocket]] = {} # user_id -> connections
self.channel_subscriptions: Dict[str, Set[WebSocket]] = {}
self.websocket_channels: Dict[WebSocket, Set[str]] = {}
self.redis_client: Optional[redis.Redis] = None
async def connect(self, websocket: WebSocket, user_id: str):
@@ -40,6 +39,39 @@ class ConnectionManager:
self.active_connections[user_id].discard(websocket)
if not self.active_connections[user_id]:
del self.active_connections[user_id]
self.unsubscribe_all(websocket)
def subscribe(self, websocket: WebSocket, channels: list[str]):
normalized_channels = {
str(channel).strip()
for channel in channels
if str(channel).strip()
}
if not normalized_channels:
return
socket_channels = self.websocket_channels.setdefault(websocket, set())
for channel in normalized_channels:
self.channel_subscriptions.setdefault(channel, set()).add(websocket)
socket_channels.add(channel)
def unsubscribe(self, websocket: WebSocket, channels: list[str]):
for channel in {str(channel).strip() for channel in channels if str(channel).strip()}:
subscribers = self.channel_subscriptions.get(channel)
if subscribers is not None:
subscribers.discard(websocket)
if not subscribers:
del self.channel_subscriptions[channel]
socket_channels = self.websocket_channels.get(websocket)
if socket_channels is not None:
socket_channels.discard(channel)
if not socket_channels:
del self.websocket_channels[websocket]
def unsubscribe_all(self, websocket: WebSocket):
channels = list(self.websocket_channels.get(websocket, set()))
if channels:
self.unsubscribe(websocket, channels)
async def send_personal_message(self, message: dict, user_id: str):
if user_id in self.active_connections:
@@ -54,13 +86,19 @@ class ConnectionManager:
for user_id in self.active_connections:
await self.send_personal_message(message, user_id)
else:
await self.send_personal_message(message, channel)
for connection in list(self.channel_subscriptions.get(channel, set())):
try:
await connection.send_json(message)
except Exception:
self.unsubscribe_all(connection)
async def close_all(self):
for user_id in self.active_connections:
for connection in self.active_connections[user_id]:
await connection.close()
self.active_connections.clear()
self.channel_subscriptions.clear()
self.websocket_channels.clear()
manager = ConnectionManager()

View File

@@ -0,0 +1,328 @@
{
"_comment": "Seed payload for the bgp_collector_locations DB table. Coordinates were migrated from the legacy RIPE_RIS_COLLECTOR_COORDS table and default to city-center; seeded rows are unverified and should be upgraded in the database with source evidence when known.",
"locations": [
{
"canonical_name": "RIPE RIS rrc00",
"aliases": ["rrc00", "RIPE RIS rrc00", "AMS-IX"],
"operator": "RIPE NCC",
"site": "AMS-IX",
"city": "Amsterdam",
"country": "Netherlands",
"latitude": 52.3676,
"longitude": 4.9041,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc01",
"aliases": ["rrc01", "RIPE RIS rrc01", "LINX"],
"operator": "RIPE NCC",
"site": "LINX",
"city": "London",
"country": "United Kingdom",
"latitude": 51.5072,
"longitude": -0.1276,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc03",
"aliases": ["rrc03", "RIPE RIS rrc03", "AMS-IX"],
"operator": "RIPE NCC",
"site": "AMS-IX",
"city": "Amsterdam",
"country": "Netherlands",
"latitude": 52.3676,
"longitude": 4.9041,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc04",
"aliases": ["rrc04", "RIPE RIS rrc04", "CIXP", "CERN Internet Exchange Point"],
"operator": "RIPE NCC",
"site": "CIXP",
"city": "Geneva",
"country": "Switzerland",
"latitude": 46.2044,
"longitude": 6.1432,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc05",
"aliases": ["rrc05", "RIPE RIS rrc05", "VIX", "Vienna Internet Exchange"],
"operator": "RIPE NCC",
"site": "VIX",
"city": "Vienna",
"country": "Austria",
"latitude": 48.2082,
"longitude": 16.3738,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc06",
"aliases": ["rrc06", "RIPE RIS rrc06", "JPIX", "Otemachi"],
"operator": "RIPE NCC",
"site": "JPIX",
"city": "Otemachi",
"country": "Japan",
"latitude": 35.686,
"longitude": 139.7671,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc07",
"aliases": ["rrc07", "RIPE RIS rrc07", "Netnod", "Netnod Stockholm"],
"operator": "RIPE NCC",
"site": "Netnod Stockholm",
"city": "Stockholm",
"country": "Sweden",
"latitude": 59.3293,
"longitude": 18.0686,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc10",
"aliases": ["rrc10", "RIPE RIS rrc10", "MIX", "Milan Internet Exchange"],
"operator": "RIPE NCC",
"site": "MIX",
"city": "Milan",
"country": "Italy",
"latitude": 45.4642,
"longitude": 9.19,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc11",
"aliases": ["rrc11", "RIPE RIS rrc11", "NYIIX", "New York International Internet Exchange"],
"operator": "RIPE NCC",
"site": "NYIIX",
"city": "New York",
"country": "United States",
"latitude": 40.7128,
"longitude": -74.006,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc12",
"aliases": ["rrc12", "RIPE RIS rrc12", "DE-CIX", "DE-CIX Frankfurt"],
"operator": "RIPE NCC",
"site": "DE-CIX Frankfurt",
"city": "Frankfurt",
"country": "Germany",
"latitude": 50.1109,
"longitude": 8.6821,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc13",
"aliases": ["rrc13", "RIPE RIS rrc13", "MSK-IX"],
"operator": "RIPE NCC",
"site": "MSK-IX",
"city": "Moscow",
"country": "Russia",
"latitude": 55.7558,
"longitude": 37.6173,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc14",
"aliases": ["rrc14", "RIPE RIS rrc14", "PAIX", "Palo Alto Internet Exchange"],
"operator": "RIPE NCC",
"site": "PAIX",
"city": "Palo Alto",
"country": "United States",
"latitude": 37.4419,
"longitude": -122.143,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc15",
"aliases": ["rrc15", "RIPE RIS rrc15", "PTT.br Sao Paulo", "PTTMetro Sao Paulo"],
"operator": "RIPE NCC",
"site": "PTT.br",
"city": "Sao Paulo",
"country": "Brazil",
"latitude": -23.5558,
"longitude": -46.6396,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc16",
"aliases": ["rrc16", "RIPE RIS rrc16", "Equinix Miami", "NOTA Miami"],
"operator": "RIPE NCC",
"site": "Equinix Miami",
"city": "Miami",
"country": "United States",
"latitude": 25.7617,
"longitude": -80.1918,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc18",
"aliases": ["rrc18", "RIPE RIS rrc18", "CATNIX"],
"operator": "RIPE NCC",
"site": "CATNIX",
"city": "Barcelona",
"country": "Spain",
"latitude": 41.3874,
"longitude": 2.1686,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc19",
"aliases": ["rrc19", "RIPE RIS rrc19", "NAPAfrica", "JINX", "NAPAfrica Johannesburg"],
"operator": "RIPE NCC",
"site": "NAPAfrica Johannesburg",
"city": "Johannesburg",
"country": "South Africa",
"latitude": -26.2041,
"longitude": 28.0473,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc20",
"aliases": ["rrc20", "RIPE RIS rrc20", "SwissIX"],
"operator": "RIPE NCC",
"site": "SwissIX",
"city": "Zurich",
"country": "Switzerland",
"latitude": 47.3769,
"longitude": 8.5417,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc21",
"aliases": ["rrc21", "RIPE RIS rrc21", "France-IX Paris"],
"operator": "RIPE NCC",
"site": "France-IX Paris",
"city": "Paris",
"country": "France",
"latitude": 48.8566,
"longitude": 2.3522,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc22",
"aliases": ["rrc22", "RIPE RIS rrc22", "InterLAN Bucharest"],
"operator": "RIPE NCC",
"site": "InterLAN Bucharest",
"city": "Bucharest",
"country": "Romania",
"latitude": 44.4268,
"longitude": 26.1025,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc23",
"aliases": ["rrc23", "RIPE RIS rrc23", "Equinix Singapore"],
"operator": "RIPE NCC",
"site": "Equinix Singapore",
"city": "Singapore",
"country": "Singapore",
"latitude": 1.3521,
"longitude": 103.8198,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc24",
"aliases": ["rrc24", "RIPE RIS rrc24", "LACNIC Montevideo"],
"operator": "RIPE NCC",
"site": "LACNIC Montevideo",
"city": "Montevideo",
"country": "Uruguay",
"latitude": -34.9011,
"longitude": -56.1645,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc25",
"aliases": ["rrc25", "RIPE RIS rrc25", "AMS-IX"],
"operator": "RIPE NCC",
"site": "AMS-IX",
"city": "Amsterdam",
"country": "Netherlands",
"latitude": 52.3676,
"longitude": 4.9041,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
},
{
"canonical_name": "RIPE RIS rrc26",
"aliases": ["rrc26", "RIPE RIS rrc26", "UAE-IX"],
"operator": "RIPE NCC",
"site": "UAE-IX",
"city": "Dubai",
"country": "United Arab Emirates",
"latitude": 25.2048,
"longitude": 55.2708,
"precision": "city",
"confidence": 0.85,
"source_note": "Migrated from RIPE_RIS_COLLECTOR_COORDS legacy table",
"verified_at": null
}
],
"city_fallbacks": []
}

View File

@@ -103,13 +103,18 @@ async def init_db():
import app.models.datasource_config # noqa: F401
import app.models.alert # noqa: F401
import app.models.bgp_anomaly # noqa: F401
import app.models.bgp_collector_location # noqa: F401
import app.models.bgp_incident # noqa: F401
import app.models.bgp_observation # noqa: F401
import app.models.collected_data # noqa: F401
import app.models.compute_center_location # noqa: F401
import app.models.system_setting # noqa: F401
import app.models.playground_session # noqa: F401
import app.models.playground_message # noqa: F401
import app.models.system_log # noqa: F401
import app.models.vessel # noqa: F401
import app.models.vessel_enrichment # noqa: F401
import app.models.datasource_mapping # noqa: F401
logger.warning_event(
"Database pool settings active",
@@ -125,6 +130,14 @@ async def init_db():
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
await conn.execute(
text(
"""
ALTER TABLE users
ADD COLUMN IF NOT EXISTS gatekeeper_groups JSONB DEFAULT '[]'::jsonb
"""
)
)
await conn.execute(
text(
"""
@@ -144,7 +157,12 @@ async def init_db():
text(
"""
ALTER TABLE collection_tasks
ADD COLUMN IF NOT EXISTS phase VARCHAR(30) DEFAULT 'queued'
ADD COLUMN IF NOT EXISTS phase VARCHAR(30) DEFAULT 'queued',
ADD COLUMN IF NOT EXISTS phase_progress DOUBLE PRECISION,
ADD COLUMN IF NOT EXISTS phase_message VARCHAR(255),
ADD COLUMN IF NOT EXISTS phase_current BIGINT,
ADD COLUMN IF NOT EXISTS phase_total BIGINT,
ADD COLUMN IF NOT EXISTS phase_unit VARCHAR(30)
"""
)
)
@@ -156,6 +174,30 @@ async def init_db():
"""
)
)
await conn.execute(
text(
"""
CREATE INDEX IF NOT EXISTS idx_collected_data_source_current_id
ON collected_data (source, is_current, id)
"""
)
)
await conn.execute(
text(
"""
CREATE INDEX IF NOT EXISTS idx_collected_data_source_task_id
ON collected_data (source, task_id, id)
"""
)
)
await conn.execute(
text(
"""
CREATE INDEX IF NOT EXISTS idx_ais_raw_schema_observed_entity
ON ais_raw_observations (target_schema, observed_at, entity_key)
"""
)
)
await conn.execute(
text(
"""
@@ -176,5 +218,14 @@ async def init_db():
)
async with async_session_factory() as session:
from app.services.bgp_collector_locations import (
seed_default_bgp_collector_locations,
)
from app.services.compute_center_locations import (
seed_compute_center_locations_from_source_coords,
)
await seed_default_bgp_collector_locations(session)
await seed_compute_center_locations_from_source_coords(session)
await seed_default_datasources(session)
await ensure_default_admin_user(session)

View File

@@ -6,12 +6,16 @@ from app.models.datasource import DataSource
from app.models.datasource_config import DataSourceConfig
from app.models.alert import Alert, AlertSeverity, AlertStatus
from app.models.bgp_anomaly import BGPAnomaly
from app.models.bgp_collector_location import BGPCollectorLocation
from app.models.bgp_incident import BGPIncident
from app.models.bgp_observation import BGPObservation
from app.models.compute_center_location import ComputeCenterLocationRecord
from app.models.system_setting import SystemSetting
from app.models.playground_session import PlaygroundSession
from app.models.playground_message import PlaygroundMessage
from app.models.system_log import SystemLog, AuditLog
from app.models.vessel import AISConflictRecord, AISRawObservation, AISSourceHealth, VesselPosition, VesselStatic
from app.models.datasource_mapping import DataSourceMappingTemplate
__all__ = [
"User",
@@ -25,8 +29,18 @@ __all__ = [
"AlertSeverity",
"AlertStatus",
"BGPAnomaly",
"BGPCollectorLocation",
"BGPIncident",
"BGPObservation",
"ComputeCenterLocationRecord",
"SystemLog",
"AuditLog",
"PlaygroundSession",
"PlaygroundMessage",
"VesselPosition",
"VesselStatic",
"AISRawObservation",
"AISConflictRecord",
"AISSourceHealth",
"DataSourceMappingTemplate",
]

View File

@@ -0,0 +1,52 @@
"""Stored BGP route-collector locations."""
from sqlalchemy import Boolean, Column, DateTime, Float, Integer, JSON, String, Text
from sqlalchemy.sql import func
from app.core.time import to_iso8601_utc
from app.db.session import Base
class BGPCollectorLocation(Base):
"""Current known location for a BGP route collector."""
__tablename__ = "bgp_collector_locations"
id = Column(Integer, primary_key=True, autoincrement=True)
collector_id = Column(String(100), nullable=False, unique=True, index=True)
operator = Column(String(255), nullable=True)
site = Column(String(255), nullable=True)
city = Column(String(255), nullable=True)
country = Column(String(255), nullable=True)
latitude = Column(Float, nullable=True)
longitude = Column(Float, nullable=True)
precision = Column(String(30), nullable=False, default="city")
confidence = Column(Float, nullable=True)
source = Column(String(80), nullable=False, default="legacy_seed", index=True)
source_url = Column(String(500), nullable=True)
source_note = Column(Text, nullable=True)
raw_payload = Column(JSON, nullable=False, default=dict)
needs_confirmation = Column(Boolean, nullable=False, default=True, index=True)
verification_status = Column(String(30), nullable=False, default="unverified", index=True)
verified_at = Column(DateTime(timezone=True), nullable=True)
created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
def to_location_dict(self) -> dict:
return {
"city": self.city,
"country": self.country,
"latitude": self.latitude,
"longitude": self.longitude,
"precision": self.precision,
"source": self.source,
"needs_confirmation": self.needs_confirmation,
"matched_location_name": self.site or self.collector_id,
"verified_at": to_iso8601_utc(self.verified_at),
"confidence": self.confidence,
"operator": self.operator,
"site": self.site,
"verification_status": self.verification_status,
"source_note": self.source_note,
"source_url": self.source_url,
}

View File

@@ -48,6 +48,8 @@ class CollectedData(Base):
# Indexes for common queries
__table_args__ = (
Index("idx_collected_data_source_collected", "source", "collected_at"),
Index("idx_collected_data_source_current_id", "source", "is_current", "id"),
Index("idx_collected_data_source_task_id", "source", "task_id", "id"),
Index("idx_collected_data_source_type", "source", "data_type"),
Index("idx_collected_data_source_source_id", "source", "source_id"),
)

View File

@@ -0,0 +1,60 @@
"""Stored compute-center locations."""
from sqlalchemy import Boolean, Column, DateTime, Float, Integer, JSON, String, Text, UniqueConstraint
from sqlalchemy.sql import func
from app.core.time import to_iso8601_utc
from app.db.session import Base
class ComputeCenterLocationRecord(Base):
"""Current known location for a compute-center record."""
__tablename__ = "compute_center_locations"
__table_args__ = (
UniqueConstraint("source", "source_id", name="uq_compute_center_location_source_id"),
)
id = Column(Integer, primary_key=True, autoincrement=True)
source = Column(String(100), nullable=False, index=True)
source_id = Column(String(255), nullable=False, index=True)
name = Column(String(500), nullable=True)
operator = Column(String(255), nullable=True)
site = Column(String(255), nullable=True)
city = Column(String(255), nullable=True)
country = Column(String(255), nullable=True)
latitude = Column(Float, nullable=True)
longitude = Column(Float, nullable=True)
precision = Column(String(30), nullable=False, default="city")
confidence = Column(Float, nullable=True)
location_source = Column(String(80), nullable=False, default="stored_compute_center_location", index=True)
source_url = Column(String(500), nullable=True)
source_note = Column(Text, nullable=True)
raw_payload = Column(JSON, nullable=False, default=dict)
needs_confirmation = Column(Boolean, nullable=False, default=False, index=True)
verification_status = Column(String(30), nullable=False, default="verified", index=True)
verified_at = Column(DateTime(timezone=True), nullable=True)
created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
def to_location_dict(self) -> dict:
return {
"source": self.source,
"source_id": self.source_id,
"name": self.name,
"operator": self.operator,
"site": self.site,
"city": self.city,
"country": self.country,
"latitude": self.latitude,
"longitude": self.longitude,
"precision": self.precision,
"confidence": self.confidence,
"location_source": self.location_source,
"source_url": self.source_url,
"source_note": self.source_note,
"raw_payload": self.raw_payload or {},
"needs_confirmation": self.needs_confirmation,
"verification_status": self.verification_status,
"verified_at": to_iso8601_utc(self.verified_at),
}

View File

@@ -0,0 +1,32 @@
"""Mapping templates for user-defined data source payloads."""
from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Integer, JSON, String
from sqlalchemy.sql import func
from app.db.session import Base
class DataSourceMappingTemplate(Base):
__tablename__ = "datasource_mapping_templates"
id = Column(Integer, primary_key=True, autoincrement=True)
datasource_config_id = Column(
Integer,
ForeignKey("datasource_configs.id"),
nullable=False,
index=True,
)
target_schema = Column(String(80), nullable=False, index=True)
mapping_json = Column(JSON, nullable=False, default={})
sample_payload_hash = Column(String(64), nullable=True)
validation_status = Column(String(30), nullable=False, default="draft")
version = Column(Integer, nullable=False, default=1)
is_active = Column(Boolean, nullable=False, default=False, index=True)
created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
def __repr__(self):
return (
f"<DataSourceMappingTemplate {self.id}: "
f"{self.datasource_config_id}/{self.target_schema}/v{self.version}>"
)

View File

@@ -1,6 +1,6 @@
"""Collection Task model"""
from sqlalchemy import Column, DateTime, Integer, String, Text, Float
from sqlalchemy import BigInteger, Column, DateTime, Integer, String, Text, Float
from sqlalchemy.sql import func
from app.db.session import Base
@@ -13,6 +13,11 @@ class CollectionTask(Base):
datasource_id = Column(Integer, nullable=False, index=True)
status = Column(String(20), nullable=False) # pending, running, success, failed, cancelled
phase = Column(String(30), default="queued")
phase_progress = Column(Float)
phase_message = Column(String(255))
phase_current = Column(BigInteger)
phase_total = Column(BigInteger)
phase_unit = Column(String(30))
started_at = Column(DateTime(timezone=True))
completed_at = Column(DateTime(timezone=True))
records_processed = Column(Integer, default=0)

View File

@@ -1,4 +1,4 @@
from sqlalchemy import Boolean, Column, Integer, String, DateTime
from sqlalchemy import Boolean, Column, DateTime, Integer, JSON, String
from sqlalchemy.sql import func
from app.db.session import Base
@@ -12,6 +12,7 @@ class User(Base):
email = Column(String(255), unique=True, index=True, nullable=False)
password_hash = Column(String(255), nullable=False)
role = Column(String(20), default="viewer")
gatekeeper_groups = Column(JSON, default=list)
is_active = Column(Boolean, default=True)
last_login_at = Column(DateTime(timezone=True))
created_at = Column(DateTime(timezone=True), server_default=func.now())

View File

@@ -0,0 +1,186 @@
"""Vessel AIS models for live maritime tracking."""
from sqlalchemy import BigInteger, Column, DateTime, Float, Index, Integer, JSON, SmallInteger, String
from sqlalchemy.sql import func
from app.core.time import to_iso8601_utc
from app.db.session import Base
class VesselStatic(Base):
"""Slow-changing vessel identity and dimensions."""
__tablename__ = "vessel_static"
mmsi = Column(BigInteger, primary_key=True)
name = Column(String(128), nullable=True)
callsign = Column(String(16), nullable=True)
vessel_type = Column(SmallInteger, nullable=True, index=True)
vessel_type_name = Column(String(64), nullable=True, index=True)
flag = Column(String(4), nullable=True, index=True)
length = Column(Float, nullable=True)
width = Column(Float, nullable=True)
draught = Column(Float, nullable=True)
imo = Column(BigInteger, nullable=True)
updated_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now())
def to_dict(self) -> dict:
return {
"mmsi": self.mmsi,
"name": self.name,
"callsign": self.callsign,
"vessel_type": self.vessel_type,
"vessel_type_name": self.vessel_type_name,
"flag": self.flag,
"length": self.length,
"width": self.width,
"draught": self.draught,
"imo": self.imo,
"updated_at": to_iso8601_utc(self.updated_at),
}
class VesselPosition(Base):
"""Append-only AIS positions retained for short history windows."""
__tablename__ = "vessel_position"
id = Column(Integer, primary_key=True, autoincrement=True)
mmsi = Column(BigInteger, nullable=False, index=True)
lat = Column(Float, nullable=False)
lon = Column(Float, nullable=False)
sog = Column(Float, nullable=True)
cog = Column(Float, nullable=True)
heading = Column(SmallInteger, nullable=True)
nav_status = Column(SmallInteger, nullable=True, index=True)
received_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now(), index=True)
__table_args__ = (
Index("idx_vessel_pos_mmsi_time", "mmsi", "received_at"),
Index("idx_vessel_pos_time", "received_at"),
Index("idx_vessel_pos_lat_lon", "lat", "lon"),
)
def to_dict(self) -> dict:
return {
"id": self.id,
"mmsi": self.mmsi,
"lat": self.lat,
"lon": self.lon,
"sog": self.sog,
"cog": self.cog,
"heading": self.heading,
"nav_status": self.nav_status,
"received_at": to_iso8601_utc(self.received_at),
}
class AISRawObservation(Base):
"""Source-level AIS fact before aggregation and conflict resolution."""
__tablename__ = "ais_raw_observations"
id = Column(Integer, primary_key=True, autoincrement=True)
target_schema = Column(String(64), nullable=False, default="vessel_ais", index=True)
source = Column(String(100), nullable=False, index=True)
entity_key = Column(String(64), nullable=False, index=True)
delivery_mode = Column(String(32), nullable=False, index=True)
transport = Column(String(32), nullable=False, index=True)
message_type = Column(String(64), nullable=True, index=True)
source_message_id = Column(String(128), nullable=True, index=True)
observation_hash = Column(String(64), nullable=False, unique=True, index=True)
observed_at = Column(DateTime(timezone=True), nullable=False, index=True)
collected_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now(), index=True)
normalized_payload = Column(JSON, default=dict)
raw_payload = Column(JSON, default=dict)
quality_flags = Column(JSON, default=list)
__table_args__ = (
Index("idx_ais_raw_entity_observed", "target_schema", "entity_key", "observed_at"),
Index("idx_ais_raw_schema_observed_entity", "target_schema", "observed_at", "entity_key"),
Index("idx_ais_raw_source_entity", "source", "entity_key"),
)
def to_dict(self) -> dict:
return {
"id": self.id,
"target_schema": self.target_schema,
"source": self.source,
"entity_key": self.entity_key,
"delivery_mode": self.delivery_mode,
"transport": self.transport,
"message_type": self.message_type,
"source_message_id": self.source_message_id,
"observation_hash": self.observation_hash,
"observed_at": to_iso8601_utc(self.observed_at),
"collected_at": to_iso8601_utc(self.collected_at),
"normalized_payload": self.normalized_payload or {},
"raw_payload": self.raw_payload or {},
"quality_flags": self.quality_flags or [],
}
class AISConflictRecord(Base):
"""Recorded field-level disagreement between AIS sources."""
__tablename__ = "ais_conflict_records"
id = Column(Integer, primary_key=True, autoincrement=True)
target_schema = Column(String(64), nullable=False, default="vessel_ais", index=True)
entity_key = Column(String(64), nullable=False, index=True)
field = Column(String(64), nullable=False, index=True)
candidates = Column(JSON, default=dict)
selected_source = Column(String(100), nullable=True, index=True)
selected_value = Column(JSON, nullable=True)
selected_reason = Column(String(64), nullable=True, index=True)
resolved_by = Column(String(32), nullable=False, default="system", index=True)
status = Column(String(32), nullable=False, default="open", index=True)
created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now(), index=True)
updated_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now())
__table_args__ = (
Index("idx_ais_conflict_entity_field", "target_schema", "entity_key", "field"),
)
def to_dict(self) -> dict:
return {
"id": self.id,
"target_schema": self.target_schema,
"entity_key": self.entity_key,
"field": self.field,
"candidates": self.candidates or {},
"selected_source": self.selected_source,
"selected_value": self.selected_value,
"selected_reason": self.selected_reason,
"resolved_by": self.resolved_by,
"status": self.status,
"created_at": to_iso8601_utc(self.created_at),
"updated_at": to_iso8601_utc(self.updated_at),
}
class AISSourceHealth(Base):
"""Runtime health signal for an AIS collector source."""
__tablename__ = "ais_source_health"
source = Column(String(100), primary_key=True)
connection_state = Column(String(32), nullable=False, default="disconnected", index=True)
last_seen_at = Column(DateTime(timezone=True), nullable=True, index=True)
last_success_at = Column(DateTime(timezone=True), nullable=True, index=True)
last_error = Column(String(500), nullable=True)
message_rate = Column(Float, nullable=True)
lag_seconds = Column(Float, nullable=True)
updated_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now(), index=True)
def to_dict(self) -> dict:
return {
"source": self.source,
"connection_state": self.connection_state,
"last_seen_at": to_iso8601_utc(self.last_seen_at),
"last_success_at": to_iso8601_utc(self.last_success_at),
"last_error": self.last_error,
"message_rate": self.message_rate,
"lag_seconds": self.lag_seconds,
"updated_at": to_iso8601_utc(self.updated_at),
}

View File

@@ -0,0 +1,63 @@
"""Vessel enrichment cache tables (v5).
Profile and media enrichment are stored separately so cache TTLs can differ
and so the conflict-resolution + display layers can read either independently.
"""
from sqlalchemy import BigInteger, Column, DateTime, Float, JSON, String
from sqlalchemy.sql import func
from app.core.time import to_iso8601_utc
from app.db.session import Base
class VesselProfileEnrichment(Base):
"""Cached static vessel profile (type, flag, dimensions, operator, etc.)."""
__tablename__ = "vessel_profile_enrichment"
mmsi = Column(BigInteger, primary_key=True)
source = Column(String(100), nullable=False, default="system")
payload = Column(JSON, nullable=False, default=dict)
fetched_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now())
expires_at = Column(DateTime(timezone=True), nullable=True)
confidence = Column(Float, nullable=True)
reference_url = Column(String(500), nullable=True)
updated_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now())
def to_dict(self) -> dict:
return {
"mmsi": self.mmsi,
"source": self.source,
"payload": self.payload or {},
"fetched_at": to_iso8601_utc(self.fetched_at),
"expires_at": to_iso8601_utc(self.expires_at),
"confidence": self.confidence,
"reference_url": self.reference_url,
}
class VesselMediaEnrichment(Base):
"""Cached vessel imagery / external detail references."""
__tablename__ = "vessel_media_enrichment"
mmsi = Column(BigInteger, primary_key=True)
source = Column(String(100), nullable=False, default="system")
payload = Column(JSON, nullable=False, default=dict)
fetched_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now())
expires_at = Column(DateTime(timezone=True), nullable=True)
confidence = Column(Float, nullable=True)
reference_url = Column(String(500), nullable=True)
updated_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now())
def to_dict(self) -> dict:
return {
"mmsi": self.mmsi,
"source": self.source,
"payload": self.payload or {},
"fetched_at": to_iso8601_utc(self.fetched_at),
"expires_at": to_iso8601_utc(self.expires_at),
"confidence": self.confidence,
"reference_url": self.reference_url,
}

View File

@@ -12,17 +12,20 @@ class UserBase(BaseModel):
class UserCreate(UserBase):
password: str = Field(..., min_length=8)
role: str = "viewer"
gatekeeper_groups: list[str] = Field(default_factory=list)
class UserUpdate(BaseModel):
email: Optional[EmailStr] = None
role: Optional[str] = None
gatekeeper_groups: Optional[list[str]] = None
is_active: Optional[bool] = None
class UserInDB(UserBase):
id: int
role: str
gatekeeper_groups: list[str] = Field(default_factory=list)
is_active: bool
last_login_at: Optional[datetime]
created_at: datetime
@@ -34,6 +37,7 @@ class UserInDB(UserBase):
class UserResponse(UserBase):
id: int
role: str
gatekeeper_groups: list[str] = Field(default_factory=list)
is_active: bool
created_at: datetime

View File

@@ -3,9 +3,11 @@ from __future__ import annotations
import asyncio
import httpx
from fastapi import HTTPException, status
from fastapi import Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import settings
from app.db.session import get_db
from app.schemas.ai import (
AIProviderStatusResponse,
SituationalAnalysisRequest,
@@ -14,11 +16,27 @@ from app.schemas.ai import (
class AIProviderClient:
def __init__(self) -> None:
self.service_url = settings.AI_PROVIDER_SERVICE_URL.rstrip("/")
self.service_token = settings.AI_PROVIDER_SERVICE_TOKEN
self.timeout = settings.AI_PROVIDER_TIMEOUT_SECONDS
self.retry_attempts = max(settings.AI_PROVIDER_RETRY_ATTEMPTS, 1)
def __init__(
self,
*,
service_url: str | None = None,
service_token: str | None = None,
timeout: int | None = None,
retry_attempts: int | None = None,
llm_config: dict | None = None,
) -> None:
self.service_url = (
service_url if service_url is not None else settings.AI_PROVIDER_SERVICE_URL
).rstrip("/")
self.service_token = (
service_token if service_token is not None else settings.AI_PROVIDER_SERVICE_TOKEN
)
self.timeout = timeout if timeout is not None else settings.AI_PROVIDER_TIMEOUT_SECONDS
self.retry_attempts = max(
retry_attempts if retry_attempts is not None else settings.AI_PROVIDER_RETRY_ATTEMPTS,
1,
)
self.llm_config = llm_config or {}
def _headers(self, request_id: str | None = None) -> dict[str, str]:
headers = {"Content-Type": "application/json"}
@@ -26,6 +44,19 @@ class AIProviderClient:
headers["X-Provider-Token"] = self.service_token
if request_id:
headers["X-Request-ID"] = request_id
llm_header_map = {
"provider": "X-AI-Provider",
"provider_api": "X-AI-Provider-API",
"base_url": "X-AI-Base-URL",
"api_key": "X-AI-API-Key",
"model": "X-AI-Model",
"max_tokens": "X-AI-Max-Tokens",
"anthropic_version": "X-AI-Anthropic-Version",
}
for key, header_name in llm_header_map.items():
value = self.llm_config.get(key)
if value not in (None, ""):
headers[header_name] = str(value)
return headers
async def get_status(self, request_id: str | None = None) -> AIProviderStatusResponse:
@@ -105,5 +136,14 @@ class AIProviderClient:
)
def get_ai_provider_client() -> AIProviderClient:
return AIProviderClient()
async def get_ai_provider_client(db: AsyncSession = Depends(get_db)) -> AIProviderClient:
from app.api.v1.settings import get_runtime_ai_provider_config
runtime_config = await get_runtime_ai_provider_config(db)
return AIProviderClient(
service_url=runtime_config["service_url"],
service_token=runtime_config["service_token"],
timeout=runtime_config["timeout_seconds"],
retry_attempts=runtime_config["retry_attempts"],
llm_config=runtime_config.get("llm_config") or {},
)

View File

@@ -0,0 +1,7 @@
"""Backend-owned AI tool services.
The services in this package are business tools used by Planet's backend
orchestrators. They intentionally live outside ``aiprovider`` so model transport
stays separate from evidence collection and domain policy.
"""

View File

@@ -0,0 +1,48 @@
from __future__ import annotations
import hashlib
from typing import Iterable
from app.services.ai_tools.schemas import FetchedEvidence, SearchEvidence
def evidence_content_hash(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def normalize_search_evidence(items: Iterable[SearchEvidence], *, limit: int = 5) -> list[dict]:
normalized: list[dict] = []
seen_urls: set[str] = set()
for item in items:
if not item.url or item.url in seen_urls:
continue
seen_urls.add(item.url)
normalized.append(
{
"title": item.title,
"url": item.url,
"snippet": item.snippet,
"content": item.compact_text(),
"score": item.score,
"source_provider": item.source_provider,
"retrieved_at": item.retrieved_at.isoformat(),
"metadata": item.metadata,
}
)
if len(normalized) >= limit:
break
return normalized
def normalize_fetched_evidence(item: FetchedEvidence, *, text_limit: int = 1200) -> dict:
text = " ".join(item.text.split())[:text_limit]
return {
"title": item.title,
"url": item.final_url or item.url,
"text": text,
"content_hash": item.content_hash or evidence_content_hash(item.text),
"extractor": item.extractor,
"retrieved_at": item.retrieved_at.isoformat(),
"metadata": item.metadata,
}

View File

@@ -0,0 +1,63 @@
from __future__ import annotations
from datetime import UTC, datetime
from typing import Any
from pydantic import BaseModel, Field
class SearchEvidence(BaseModel):
title: str = ""
url: str = ""
snippet: str = ""
content: str = ""
score: float | None = None
source_provider: str = ""
retrieved_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
metadata: dict[str, Any] = Field(default_factory=dict)
def compact_text(self, limit: int = 700) -> str:
text = " ".join((self.content or self.snippet or "").split())
return text[:limit]
class FetchedEvidence(BaseModel):
url: str
final_url: str = ""
title: str = ""
text: str = ""
content_hash: str = ""
extractor: str = "basic_html"
retrieved_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
metadata: dict[str, Any] = Field(default_factory=dict)
class WebSearchProviderConfig(BaseModel):
provider: str = "tavily"
base_url: str = ""
api_key: str = ""
max_results: int = Field(default=5, ge=1, le=20)
timeout_seconds: int = Field(default=20, ge=3, le=120)
endpoint_path: str = ""
search_depth: str = "basic"
engine: str = "google"
include_answer: bool = False
include_raw_content: bool = False
include_text: bool = False
categories: str = "general"
engines: list[str] = Field(default_factory=list)
search_path: str = ""
scrape_path: str = ""
scrape_formats: list[str] = Field(default_factory=lambda: ["markdown"])
class WebSearchConfig(BaseModel):
enabled: bool = False
default_provider: str = "tavily"
provider: str = "tavily"
providers: dict[str, WebSearchProviderConfig] = Field(default_factory=dict)
@property
def active_provider_config(self) -> WebSearchProviderConfig:
return self.providers.get(self.default_provider) or self.providers.get(self.provider) or WebSearchProviderConfig(provider=self.default_provider or self.provider)

View File

@@ -0,0 +1,56 @@
from __future__ import annotations
import hashlib
import httpx
from bs4 import BeautifulSoup
from app.services.ai_tools.schemas import FetchedEvidence
class WebFetchError(RuntimeError):
pass
def _extract_title_and_text(html: str) -> tuple[str, str]:
soup = BeautifulSoup(html, "html.parser")
for tag in soup(["script", "style", "noscript", "svg"]):
tag.decompose()
title = soup.title.get_text(" ", strip=True) if soup.title else ""
main = soup.find("main") or soup.find("article") or soup.body or soup
text = main.get_text("\n", strip=True)
lines = [line.strip() for line in text.splitlines() if line.strip()]
return title, "\n".join(lines)
async def fetch_url_evidence(
url: str,
*,
timeout_seconds: int = 20,
max_bytes: int = 1_500_000,
) -> FetchedEvidence:
if not url:
raise WebFetchError("url is required")
try:
async with httpx.AsyncClient(
timeout=timeout_seconds,
follow_redirects=True,
headers={"User-Agent": "PlanetEvidenceFetcher/1.0"},
) as client:
response = await client.get(url)
response.raise_for_status()
content = response.content[:max_bytes]
except httpx.HTTPError as exc:
raise WebFetchError(f"failed to fetch page: {exc}") from exc
title, text = _extract_title_and_text(content.decode(response.encoding or "utf-8", errors="ignore"))
content_hash = hashlib.sha256(text.encode("utf-8")).hexdigest()
return FetchedEvidence(
url=url,
final_url=str(response.url),
title=title,
text=text,
content_hash=content_hash,
extractor="beautifulsoup_basic",
)

View File

@@ -0,0 +1,391 @@
from __future__ import annotations
from copy import deepcopy
from typing import Any
import httpx
from app.services.ai_tools.schemas import SearchEvidence, WebSearchConfig, WebSearchProviderConfig
WEB_SEARCH_PROVIDER_PRESETS: dict[str, dict[str, Any]] = {
"tavily": {
"provider": "tavily",
"label": "Tavily",
"api_key_env": "TAVILY_API_KEY",
"base_url": "https://api.tavily.com",
"endpoint_path": "/search",
"max_results": 5,
"timeout_seconds": 20,
"search_depth": "basic",
"include_answer": False,
"include_raw_content": False,
},
"brave": {
"provider": "brave",
"label": "Brave Search API",
"api_key_env": "BRAVE_SEARCH_API_KEY",
"base_url": "https://api.search.brave.com",
"endpoint_path": "/res/v1/web/search",
"max_results": 5,
"timeout_seconds": 20,
},
"serpapi": {
"provider": "serpapi",
"label": "SerpAPI",
"api_key_env": "SERPAPI_API_KEY",
"base_url": "https://serpapi.com",
"endpoint_path": "/search.json",
"engine": "google",
"max_results": 5,
"timeout_seconds": 20,
},
"exa": {
"provider": "exa",
"label": "Exa",
"api_key_env": "EXA_API_KEY",
"base_url": "https://api.exa.ai",
"endpoint_path": "/search",
"max_results": 5,
"timeout_seconds": 20,
"include_text": False,
},
"firecrawl": {
"provider": "firecrawl",
"label": "Firecrawl Search / Scrape",
"api_key_env": "FIRECRAWL_API_KEY",
"base_url": "https://api.firecrawl.dev",
"search_path": "/v2/search",
"scrape_path": "/v2/scrape",
"max_results": 5,
"timeout_seconds": 30,
"scrape_formats": ["markdown"],
},
"searxng": {
"provider": "searxng",
"label": "SearXNG",
"api_key_env": "SEARXNG_API_KEY",
"base_url": "http://localhost:8080",
"endpoint_path": "/",
"max_results": 5,
"timeout_seconds": 20,
"categories": "general",
"engines": [],
},
}
class WebSearchError(RuntimeError):
pass
class WebSearchConfigurationError(WebSearchError):
pass
def normalize_web_search_provider(provider: str | None) -> str:
return (provider or "tavily").strip().lower() or "tavily"
def get_web_search_provider_preset(provider: str) -> dict[str, Any]:
provider_id = normalize_web_search_provider(provider)
preset = WEB_SEARCH_PROVIDER_PRESETS.get(provider_id)
if not preset:
raise ValueError(f"Unsupported web search provider: {provider}")
return deepcopy(preset)
def list_web_search_provider_presets() -> list[dict[str, Any]]:
return [get_web_search_provider_preset(provider) for provider in WEB_SEARCH_PROVIDER_PRESETS]
def provider_defaults(provider: str) -> WebSearchProviderConfig:
preset = get_web_search_provider_preset(provider)
return WebSearchProviderConfig(**{
key: value
for key, value in preset.items()
if key in WebSearchProviderConfig.model_fields
})
class WebSearchClient:
def __init__(self, config: WebSearchConfig) -> None:
self.config = config
async def search(
self,
query: str,
*,
max_results: int | None = None,
domains: list[str] | None = None,
freshness_days: int | None = None,
) -> list[SearchEvidence]:
if not self.config.enabled:
raise WebSearchConfigurationError("WebSearch is disabled.")
provider_config = self.config.active_provider_config
provider = normalize_web_search_provider(provider_config.provider)
if provider != "searxng" and not provider_config.api_key:
raise WebSearchConfigurationError(f"{provider} API key is not configured.")
query = " ".join(str(query or "").split())
if not query:
raise WebSearchConfigurationError("search query is required.")
limit = max_results or provider_config.max_results
if provider == "tavily":
return await self._search_tavily(provider_config, query, limit, domains, freshness_days)
if provider == "brave":
return await self._search_brave(provider_config, query, limit, domains)
if provider == "serpapi":
return await self._search_serpapi(provider_config, query, limit)
if provider == "exa":
return await self._search_exa(provider_config, query, limit, domains)
if provider == "firecrawl":
return await self._search_firecrawl(provider_config, query, limit)
if provider == "searxng":
return await self._search_searxng(provider_config, query, limit, domains)
raise WebSearchConfigurationError(f"Unsupported web search provider: {provider}")
async def test_connection(self) -> list[SearchEvidence]:
return await self.search("Planet WebSearch connectivity test", max_results=1)
async def _request_json(
self,
method: str,
url: str,
*,
provider_config: WebSearchProviderConfig,
headers: dict[str, str] | None = None,
params: dict[str, Any] | None = None,
json: dict[str, Any] | None = None,
) -> dict[str, Any]:
try:
async with httpx.AsyncClient(timeout=provider_config.timeout_seconds) as client:
response = await client.request(
method,
url,
headers=headers,
params=params,
json=json,
)
response.raise_for_status()
data = response.json()
except httpx.HTTPStatusError as exc:
detail = exc.response.text or exc.response.reason_phrase
raise WebSearchError(f"{provider_config.provider} request failed: {detail}") from exc
except httpx.HTTPError as exc:
raise WebSearchError(f"{provider_config.provider} request failed: {exc}") from exc
except ValueError as exc:
raise WebSearchError(f"{provider_config.provider} returned invalid JSON") from exc
return data if isinstance(data, dict) else {}
async def _search_tavily(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
domains: list[str] | None,
freshness_days: int | None,
) -> list[SearchEvidence]:
body: dict[str, Any] = {
"api_key": config.api_key,
"query": query,
"max_results": max_results,
"search_depth": config.search_depth or "basic",
"include_answer": config.include_answer,
"include_raw_content": config.include_raw_content,
}
if domains:
body["include_domains"] = domains
if freshness_days:
body["days"] = freshness_days
data = await self._request_json(
"POST",
_join_url(config.base_url, config.endpoint_path or "/search"),
provider_config=config,
json=body,
)
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or ""),
snippet=str(item.get("content") or ""),
content=str(item.get("raw_content") or ""),
score=_float_or_none(item.get("score")),
source_provider="tavily",
metadata={"query": data.get("query") or query},
)
for item in data.get("results") or []
if isinstance(item, dict) and item.get("url")
]
async def _search_brave(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
domains: list[str] | None,
) -> list[SearchEvidence]:
search_query = query
if domains:
search_query = f"{query} " + " ".join(f"site:{domain}" for domain in domains)
data = await self._request_json(
"GET",
_join_url(config.base_url, config.endpoint_path or "/res/v1/web/search"),
provider_config=config,
headers={"X-Subscription-Token": config.api_key},
params={"q": search_query, "count": max_results},
)
results = (data.get("web") or {}).get("results") or []
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or ""),
snippet=str(item.get("description") or ""),
source_provider="brave",
metadata={"age": item.get("age")},
)
for item in results
if isinstance(item, dict) and item.get("url")
]
async def _search_serpapi(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
) -> list[SearchEvidence]:
data = await self._request_json(
"GET",
_join_url(config.base_url, config.endpoint_path or "/search.json"),
provider_config=config,
params={
"api_key": config.api_key,
"engine": config.engine or "google",
"q": query,
"num": max_results,
},
)
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("link") or ""),
snippet=str(item.get("snippet") or ""),
source_provider="serpapi",
metadata={"position": item.get("position")},
)
for item in data.get("organic_results") or []
if isinstance(item, dict) and item.get("link")
]
async def _search_exa(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
domains: list[str] | None,
) -> list[SearchEvidence]:
body: dict[str, Any] = {
"query": query,
"numResults": max_results,
}
if domains:
body["includeDomains"] = domains
if config.include_text:
body["contents"] = {"text": True}
data = await self._request_json(
"POST",
_join_url(config.base_url, config.endpoint_path or "/search"),
provider_config=config,
headers={"Authorization": f"Bearer {config.api_key}"},
json=body,
)
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or ""),
snippet=str(item.get("summary") or ""),
content=str(item.get("text") or ""),
score=_float_or_none(item.get("score")),
source_provider="exa",
metadata={"id": item.get("id")},
)
for item in data.get("results") or []
if isinstance(item, dict) and item.get("url")
]
async def _search_firecrawl(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
) -> list[SearchEvidence]:
data = await self._request_json(
"POST",
_join_url(config.base_url, config.search_path or "/v2/search"),
provider_config=config,
headers={"Authorization": f"Bearer {config.api_key}"},
json={"query": query, "limit": max_results},
)
raw_results = data.get("data") or data.get("results") or []
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or item.get("sourceURL") or ""),
snippet=str(item.get("description") or item.get("markdown") or ""),
source_provider="firecrawl",
metadata={"status": item.get("status")},
)
for item in raw_results
if isinstance(item, dict) and (item.get("url") or item.get("sourceURL"))
]
async def _search_searxng(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
domains: list[str] | None,
) -> list[SearchEvidence]:
search_query = query
if domains:
search_query = f"{query} " + " ".join(f"site:{domain}" for domain in domains)
params: dict[str, Any] = {
"q": search_query,
"format": "json",
"categories": config.categories or "general",
}
if config.engines:
params["engines"] = ",".join(config.engines)
headers = {"Authorization": f"Bearer {config.api_key}"} if config.api_key else None
data = await self._request_json(
"GET",
_join_url(config.base_url, config.endpoint_path or "/"),
provider_config=config,
headers=headers,
params=params,
)
results = data.get("results") or []
evidence = [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or ""),
snippet=str(item.get("content") or ""),
score=_float_or_none(item.get("score")),
source_provider="searxng",
metadata={"engine": item.get("engine")},
)
for item in results
if isinstance(item, dict) and item.get("url")
]
return evidence[:max_results]
def _join_url(base_url: str, path: str) -> str:
return f"{(base_url or '').rstrip('/')}/{(path or '').lstrip('/')}"
def _float_or_none(value: Any) -> float | None:
try:
return float(value)
except (TypeError, ValueError):
return None

View File

@@ -0,0 +1,209 @@
"""BarentsWatch AIS credential resolution and connectivity checks."""
from __future__ import annotations
import os
import shlex
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import httpx
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.data_sources import get_data_sources_config
from app.models.datasource_config import DataSourceConfig
BARENTSWATCH_LATEST_URL = "https://live.ais.barentswatch.no/v1/latest/combined"
BARENTSWATCH_TOKEN_URL = "https://id.barentswatch.no/connect/token"
BARENTSWATCH_DATASOURCE_NAME = "barentswatch_vessels"
@dataclass(frozen=True)
class BarentsWatchConfig:
endpoint: str
client_id: str
client_secret: str
credential_source: str
endpoint_source: str
def _read_zshrc_env(path: Path | None = None) -> dict[str, str]:
zshrc_path = path or Path.home() / ".zshrc"
if not zshrc_path.exists():
return {}
values: dict[str, str] = {}
for raw_line in zshrc_path.read_text(encoding="utf-8", errors="ignore").splitlines():
line = raw_line.strip()
if not line or line.startswith("#"):
continue
if line.startswith("export "):
line = line[len("export ") :].strip()
if "=" not in line:
continue
key, value = line.split("=", 1)
key = key.strip()
if not key or not key.replace("_", "").isalnum() or not key[0].isalpha():
continue
try:
parsed = shlex.split(value, comments=True, posix=True)
except ValueError:
parsed = [value.strip().strip("'\"")]
if parsed:
values[key] = parsed[0]
return values
def _first_env_value(zshrc_env: dict[str, str], *keys: str) -> tuple[str, str]:
for key in keys:
value = os.getenv(key)
if value:
return value, "environment"
for key in keys:
value = zshrc_env.get(key)
if value:
return value, "~/.zshrc"
return "", ""
async def get_barentswatch_datasource_record(db: AsyncSession) -> DataSourceConfig | None:
result = await db.execute(
select(DataSourceConfig)
.where(DataSourceConfig.name == BARENTSWATCH_DATASOURCE_NAME)
.where(DataSourceConfig.is_active.is_(True))
)
return result.scalar_one_or_none()
async def resolve_barentswatch_config(db: AsyncSession | None = None) -> BarentsWatchConfig:
record = await get_barentswatch_datasource_record(db) if db else None
auth_config = dict(record.auth_config or {}) if record else {}
config = dict(record.config or {}) if record else {}
zshrc_env = _read_zshrc_env()
env_client_id, env_source = _first_env_value(
zshrc_env,
"BARENTSWATCH_CLIENT_ID",
"BARRENTSWATCH_CLIENT_ID",
)
env_client_secret, secret_env_source = _first_env_value(
zshrc_env,
"BARENTSWATCH_CLIENT_SECRET",
"BARRENTSWATCH_CLIENT_SECRET",
)
client_id = auth_config.get("client_id") or config.get("client_id") or env_client_id
client_secret = (
auth_config.get("client_secret") or config.get("client_secret") or env_client_secret
)
credential_source = ""
if auth_config.get("client_id") or auth_config.get("client_secret"):
credential_source = "datasource_config"
elif config.get("client_id") or config.get("client_secret"):
credential_source = "datasource_runtime_config"
elif env_source or secret_env_source:
credential_source = env_source or secret_env_source
yaml_endpoint = get_data_sources_config().get_yaml_url(BARENTSWATCH_DATASOURCE_NAME)
endpoint = record.endpoint if record and record.endpoint else yaml_endpoint
return BarentsWatchConfig(
endpoint=endpoint or BARENTSWATCH_LATEST_URL,
client_id=str(client_id or ""),
client_secret=str(client_secret or ""),
credential_source=credential_source or "missing",
endpoint_source="datasource_config" if record and record.endpoint else "default",
)
async def fetch_barentswatch_access_token(
client: httpx.AsyncClient,
config: BarentsWatchConfig,
) -> str | None:
if not config.client_id or not config.client_secret:
return None
response = await client.post(
BARENTSWATCH_TOKEN_URL,
data={
"client_id": config.client_id,
"client_secret": config.client_secret,
"scope": "ais",
"grant_type": "client_credentials",
},
headers={"Content-Type": "application/x-www-form-urlencoded"},
)
response.raise_for_status()
payload = response.json()
token = payload.get("access_token")
return str(token) if token else None
async def check_barentswatch_connectivity(db: AsyncSession) -> dict[str, Any]:
config = await resolve_barentswatch_config(db)
return await check_barentswatch_config(config)
async def check_barentswatch_config(config: BarentsWatchConfig) -> dict[str, Any]:
if not config.client_id or not config.client_secret:
return {
"success": False,
"stage": "credentials",
"message": "未找到 BarentsWatch client id/client secret请先配置采集器凭证。",
"endpoint": config.endpoint,
"credential_source": config.credential_source,
"settings_tab": "collector_credentials",
}
try:
async with httpx.AsyncClient(timeout=20.0) as client:
token = await fetch_barentswatch_access_token(client, config)
if not token:
return {
"success": False,
"stage": "token",
"message": "BarentsWatch token 响应中没有 access_token请检查凭证。",
"endpoint": config.endpoint,
"credential_source": config.credential_source,
"settings_tab": "collector_credentials",
}
async with client.stream(
"GET",
config.endpoint,
headers={"Authorization": f"Bearer {token}"},
) as response:
response.raise_for_status()
return {
"success": True,
"stage": "endpoint",
"message": "BarentsWatch AIS token 和数据接口均可连通。",
"endpoint": config.endpoint,
"credential_source": config.credential_source,
"endpoint_source": config.endpoint_source,
}
except httpx.HTTPStatusError as exc:
status_code = exc.response.status_code
stage = "token" if str(exc.request.url) == BARENTSWATCH_TOKEN_URL else "endpoint"
return {
"success": False,
"stage": stage,
"message": f"BarentsWatch {stage} 请求返回 HTTP {status_code},请检查凭证或接口地址。",
"endpoint": config.endpoint,
"credential_source": config.credential_source,
"settings_tab": "collector_credentials",
}
except httpx.HTTPError as exc:
return {
"success": False,
"stage": "network",
"message": f"BarentsWatch 链路检查失败:{exc.__class__.__name__}",
"endpoint": config.endpoint,
"credential_source": config.credential_source,
"settings_tab": "collector_credentials",
}

View File

@@ -0,0 +1,324 @@
"""BGP route-collector location resolver.
Collector positions are stored in the ``bgp_collector_locations`` database
table. The old JSON registry is now only a seed payload used during database
initialization, not a runtime resolver or candidate source.
"""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any, Iterator
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.bgp_collector_location import BGPCollectorLocation
from app.services.location import (
LocationCandidate,
LocationPipeline,
LocationQuery,
NominatimResolver,
ResolutionResult,
ResolverOutput,
SourceCoordinatesResolver,
build_default_nominatim_geocoder,
coerce_str,
normalize_text,
)
SEED_PATH = (
Path(__file__).resolve().parents[1]
/ "data"
/ "seeds"
/ "ripe_ris_collector_locations_seed.json"
)
# ── Geocoder (kept at module level for monkeypatching + cache_clear) ──
_geocode_online = build_default_nominatim_geocoder()
# ── In-process compatibility cache ──────────────────────────────────
RIPE_RIS_COLLECTOR_COORDS: dict[str, dict[str, Any]] = {}
def _collector_record_to_dict(record: BGPCollectorLocation) -> dict[str, Any]:
return record.to_location_dict()
def set_bgp_collector_location_cache(
locations: dict[str, dict[str, Any]],
) -> None:
"""Replace the legacy compatibility cache in-place."""
RIPE_RIS_COLLECTOR_COORDS.clear()
RIPE_RIS_COLLECTOR_COORDS.update(
{coerce_str(key): dict(value) for key, value in locations.items()}
)
async def refresh_bgp_collector_location_cache(
session: AsyncSession,
) -> dict[str, dict[str, Any]]:
result = await session.execute(select(BGPCollectorLocation))
records = result.scalars().all()
cache = {
record.collector_id: _collector_record_to_dict(record)
for record in records
if record.collector_id
}
set_bgp_collector_location_cache(cache)
return cache
def _load_seed_payload() -> dict[str, Any]:
with SEED_PATH.open("r", encoding="utf-8") as handle:
return json.load(handle)
def _seed_entry_to_record_kwargs(entry: dict[str, Any], collector_id: str) -> dict[str, Any]:
return {
"collector_id": collector_id,
"operator": entry.get("operator") or "RIPE NCC",
"site": entry.get("site"),
"city": entry.get("city"),
"country": entry.get("country"),
"latitude": entry.get("latitude"),
"longitude": entry.get("longitude"),
"precision": entry.get("precision") or "city",
"confidence": entry.get("confidence"),
"source": "legacy_seed",
"source_url": None,
"source_note": entry.get("source_note")
or "Seeded from legacy RIPE RIS collector coordinates",
"raw_payload": entry,
"needs_confirmation": True,
"verification_status": "unverified",
"verified_at": None,
}
async def seed_default_bgp_collector_locations(session: AsyncSession) -> None:
"""Seed default RIPE RIS collector locations without overwriting users."""
payload = _load_seed_payload()
for entry in payload.get("locations", []):
aliases = entry.get("aliases") or []
collector_ids = [
coerce_str(alias)
for alias in aliases
if coerce_str(alias).startswith("rrc")
]
if not collector_ids:
continue
collector_id = collector_ids[0]
existing = await session.scalar(
select(BGPCollectorLocation).where(
BGPCollectorLocation.collector_id == collector_id
)
)
if existing:
continue
session.add(
BGPCollectorLocation(
**_seed_entry_to_record_kwargs(entry, collector_id)
)
)
await session.commit()
await refresh_bgp_collector_location_cache(session)
def get_bgp_collector_location_dict(collector_name: str) -> dict[str, Any]:
"""Return the current cached collector location dict, or ``{}`` if unknown."""
return dict(RIPE_RIS_COLLECTOR_COORDS.get(coerce_str(collector_name), {}))
def iter_known_collector_names() -> Iterator[str]:
"""Yield every collector technical name (rrcXX) known in the cache."""
return iter(sorted(RIPE_RIS_COLLECTOR_COORDS.keys()))
# ── Pipeline construction ──────────────────────────────────────────
class StoredCollectorLocationResolver:
"""Resolve a collector through the DB-backed compatibility cache."""
name = "stored_collector_location"
def resolve(self, query: LocationQuery) -> ResolverOutput:
collector = coerce_str(query.name)
if not collector:
for alias in query.aliases:
collector = coerce_str(alias)
if collector:
break
if not collector:
return ResolverOutput()
location = get_bgp_collector_location_dict(collector)
if not location:
return ResolverOutput()
latitude = location.get("latitude")
longitude = location.get("longitude")
if latitude in (None, 0.0) or longitude in (None, 0.0):
return ResolverOutput()
return ResolverOutput(
candidates=(
LocationCandidate(
latitude=float(latitude),
longitude=float(longitude),
display_name=location.get("matched_location_name") or collector,
precision=location.get("precision") or "city",
confidence=float(location.get("confidence") or 0.85),
query=f"stored_collector_location::{collector}",
source=location.get("source") or self.name,
source_note=location.get("source_note"),
matched_fields=("collector",),
needs_confirmation=bool(location.get("needs_confirmation")),
city=location.get("city"),
region=None,
country=location.get("country"),
matched_location_name=(
location.get("matched_location_name") or collector
),
location_verified_at=location.get("verified_at"),
suggested_registry_entry=None,
),
)
)
def _bgp_collector_query_plan(
query: LocationQuery,
) -> list[tuple[str, tuple[str, ...]]]:
"""Build the Nominatim query plan for a BGP collector."""
extra = query.extra or {}
site = str(extra.get("site") or "")
operator = str(extra.get("operator") or "")
city = query.city or ""
country = query.country or ""
plan: list[tuple[str, tuple[str, ...]]] = []
def add(parts: list[tuple[str, str]]) -> None:
non_empty = [(field, value) for field, value in parts if value]
if not non_empty:
return
seen: set[str] = set()
cleaned: list[str] = []
fields: list[str] = []
for field, value in non_empty:
key = normalize_text(value)
if not key or key in seen:
continue
seen.add(key)
cleaned.append(value)
fields.append(field)
if not cleaned:
return
composed = ", ".join(cleaned)
if not any(composed == existing for existing, _ in plan):
plan.append((composed, tuple(fields)))
add([("site", site), ("city", city), ("country", country)])
add([("site", site), ("country", country)])
add([("operator", operator), ("city", city), ("country", country)])
add([("city", city), ("country", country)])
return plan
BGP_COLLECTOR_PIPELINE = LocationPipeline(
[
SourceCoordinatesResolver(),
StoredCollectorLocationResolver(),
],
failure_reason=(
"Could not resolve BGP collector to renderable coordinates from"
" source coordinates or stored collector location."
),
)
BGP_COLLECTOR_COLLECTION_PIPELINE = LocationPipeline(
[
SourceCoordinatesResolver(),
NominatimResolver(
query_plan_builder=_bgp_collector_query_plan,
# Late-binding so tests can monkeypatch ``_geocode_online``.
geocoder=lambda q: _geocode_online(q),
),
],
failure_reason=(
"Could not resolve BGP collector to renderable coordinates from"
" source coordinates or online geocoding."
),
)
# ── Public API ─────────────────────────────────────────────────────
def resolve_bgp_collector_location(
collector_name: str,
*,
city: str | None = None,
country: str | None = None,
site: str | None = None,
operator: str | None = None,
) -> ResolutionResult:
"""Resolve a BGP collector to its best-known stored location."""
stored = get_bgp_collector_location_dict(collector_name)
name = coerce_str(collector_name) or None
query = LocationQuery(
name=name,
aliases=tuple(filter(None, (collector_name,))),
city=coerce_str(city or stored.get("city")) or None,
country=coerce_str(country or stored.get("country")) or None,
extra={
"site": coerce_str(site or stored.get("site")),
"operator": coerce_str(operator or stored.get("operator")) or "RIPE NCC",
},
)
return BGP_COLLECTOR_PIPELINE.resolve_best(query)
def collect_bgp_collector_location_candidates(
*,
collector: str | None = None,
city: str | None = None,
country: str | None = None,
site: str | None = None,
operator: str | None = None,
) -> tuple[list[LocationCandidate], list[str]]:
query = build_bgp_collector_location_query(
collector=collector,
city=city,
country=country,
site=site,
operator=operator,
)
return BGP_COLLECTOR_COLLECTION_PIPELINE.collect_candidates(query)
def build_bgp_collector_location_query(
*,
collector: str | None = None,
city: str | None = None,
country: str | None = None,
site: str | None = None,
operator: str | None = None,
) -> LocationQuery:
stored = get_bgp_collector_location_dict(collector or "")
name = coerce_str(collector) or None
return LocationQuery(
name=name,
aliases=tuple(filter(None, (collector,))),
city=coerce_str(city or stored.get("city")) or None,
country=coerce_str(country or stored.get("country")) or None,
extra={
"site": coerce_str(site or stored.get("site")),
"operator": coerce_str(operator or stored.get("operator")) or "RIPE NCC",
"collector": coerce_str(collector),
},
)

View File

@@ -0,0 +1,155 @@
"""BGP event location resolver.
A BGP event (announcement / withdrawal / RIB entry) is geographically tied to
the route collector that observed it. This module defines the pipeline that
turns an event payload into renderable coordinates.
Current resolver chain:
SourceCoordinates → event payload itself carries lat/lon (rare; some
enriched feeds do).
InheritFromCollector → look up the owning collector via
:func:`resolve_bgp_collector_location`.
Future plug-ins (no consumer changes required, just append to the list):
ASNFacilityResolver — origin/peer ASN → peeringdb facility.
PrefixGeoResolver — prefix → IP range geo lookup (iptoasn / opengeofeed).
"""
from __future__ import annotations
from typing import Any
from app.services.bgp_collector_locations import (
get_bgp_collector_location_dict,
)
from app.services.location import (
InheritFromAnotherEntityResolver,
LocationCandidate,
LocationPipeline,
LocationQuery,
ResolutionResult,
SourceCoordinatesResolver,
coerce_str,
)
def _inherit_from_owning_collector(
query: LocationQuery,
) -> LocationCandidate | None:
"""Look up the event's owning collector by exact name in the DB-backed cache."""
extra = query.extra or {}
collector_name = coerce_str(extra.get("collector"))
if not collector_name:
return None
legacy = get_bgp_collector_location_dict(collector_name)
if not legacy:
return None
latitude = legacy.get("latitude")
longitude = legacy.get("longitude")
if latitude in (None, 0.0) or longitude in (None, 0.0):
return None
return LocationCandidate(
latitude=float(latitude),
longitude=float(longitude),
display_name=legacy.get("matched_location_name") or collector_name,
precision=legacy.get("precision") or "city",
confidence=float(legacy.get("confidence") or 0.85),
query=f"inherit_from_collector::{collector_name}",
source="inherited_from_collector",
source_note=(
f"Inherited from owning collector {collector_name}"
),
matched_fields=("collector",),
needs_confirmation=bool(legacy.get("needs_confirmation")),
city=legacy.get("city"),
region=None,
country=legacy.get("country"),
matched_location_name=legacy.get("matched_location_name"),
location_verified_at=legacy.get("verified_at"),
suggested_registry_entry=None,
)
BGP_EVENT_PIPELINE = LocationPipeline(
[
SourceCoordinatesResolver(),
InheritFromAnotherEntityResolver(
source_lookup=_inherit_from_owning_collector,
name="inherited_from_collector",
),
# Plug new resolvers (peeringdb / ASN facility / prefix-geo) here.
],
failure_reason=(
"Could not resolve BGP event coordinates: no source coords, owning"
" collector unknown, and no fallback resolver matched."
),
)
def resolve_bgp_event_location(
*,
collector: str,
source_latitude: float | None = None,
source_longitude: float | None = None,
site: str | None = None,
operator: str | None = None,
peer_asn: int | None = None,
origin_asn: int | None = None,
prefix: str | None = None,
) -> ResolutionResult:
"""Resolve a BGP event to its renderable coordinates.
The ``peer_asn`` / ``origin_asn`` / ``prefix`` arguments are accepted
today so future resolvers (ASN→facility, prefix→geo) can consume them
without callers needing to change.
"""
query = LocationQuery(
name=collector or None,
aliases=tuple(filter(None, (collector,))),
source_latitude=source_latitude,
source_longitude=source_longitude,
extra={
"collector": collector or "",
"site": coerce_str(site),
"operator": coerce_str(operator),
"peer_asn": peer_asn,
"origin_asn": origin_asn,
"prefix": coerce_str(prefix),
},
)
return BGP_EVENT_PIPELINE.resolve_best(query)
def resolve_bgp_event_geo_dict(
collector: str,
*,
source_latitude: float | None = None,
source_longitude: float | None = None,
) -> dict[str, Any]:
"""Convenience wrapper returning the legacy ``collector_geo`` dict shape.
Preserves ``city``/``country``/``latitude``/``longitude`` keys (consumed
by existing detectors / enrichment / DB serialization) and adds
``precision``/``source``/``needs_confirmation`` for richer downstream use.
"""
result = resolve_bgp_event_location(
collector=collector,
source_latitude=source_latitude,
source_longitude=source_longitude,
)
candidate = result.location
if candidate is None:
return {}
return {
"city": candidate.city,
"country": candidate.country,
"latitude": candidate.latitude,
"longitude": candidate.longitude,
"precision": candidate.precision,
"source": candidate.source,
"needs_confirmation": candidate.needs_confirmation,
"matched_location_name": candidate.matched_location_name,
"confidence": candidate.confidence,
}

View File

@@ -36,6 +36,8 @@ from app.services.collectors.iptoasn import IPtoASNPrefixGeoCollector
from app.services.collectors.opengeofeed import OpenGeoFeedPrefixGeoCollector
from app.services.collectors.nro_delegated import NRODelegatedPrefixGeoCollector
from app.services.collectors.news_live_streams import NewsLiveStreamsCollector
from app.services.collectors.aisstream import AISStreamCollector
from app.services.collectors.vessel_ais import VesselAISCollector
collector_registry.register(TOP500Collector())
collector_registry.register(EpochAIGPUCollector())
@@ -63,3 +65,41 @@ collector_registry.register(IPtoASNPrefixGeoCollector())
collector_registry.register(OpenGeoFeedPrefixGeoCollector())
collector_registry.register(NRODelegatedPrefixGeoCollector())
collector_registry.register(NewsLiveStreamsCollector())
collector_registry.register(VesselAISCollector())
collector_registry.register(AISStreamCollector())
__all__ = [
"BaseCollector",
"HTTPCollector",
"IntervalCollector",
"collector_registry",
"CollectorRegistry",
"TOP500Collector",
"EpochAIGPUCollector",
"HuggingFaceModelCollector",
"HuggingFaceDatasetCollector",
"HuggingFaceSpacesCollector",
"PeeringDBIXPCollector",
"PeeringDBNetworkCollector",
"PeeringDBFacilityCollector",
"TeleGeographyCableCollector",
"TeleGeographyLandingPointCollector",
"TeleGeographyCableSystemCollector",
"CloudflareRadarDeviceCollector",
"CloudflareRadarTrafficCollector",
"CloudflareRadarTopASCollector",
"ArcGISCableCollector",
"FAOLandingPointCollector",
"ArcGISLandingPointCollector",
"ArcGISCableLandingRelationCollector",
"SpaceTrackTLECollector",
"CelesTrakTLECollector",
"RISLiveCollector",
"BGPStreamBackfillCollector",
"IPtoASNPrefixGeoCollector",
"OpenGeoFeedPrefixGeoCollector",
"NRODelegatedPrefixGeoCollector",
"NewsLiveStreamsCollector",
"VesselAISCollector",
"AISStreamCollector",
]

View File

@@ -0,0 +1,491 @@
"""AISStream WebSocket collector for realtime vessel AIS observations."""
from datetime import UTC, datetime
import asyncio
import json
import os
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.data_sources import get_data_sources_config
from app.core.time import to_iso8601_utc
from app.core.websocket.broadcaster import broadcaster
from app.models.datasource_config import DataSourceConfig
from app.models.task import CollectionTask
from app.services.collectors.base import BaseCollector
from app.services.vessel_ais_aggregation import (
AISSTREAM_DELIVERY_MODE,
AISSTREAM_TRANSPORT,
record_vessel_ais_observation,
update_ais_source_health,
)
from app.services.vessel_types import normalize_vessel_type_name
DEFAULT_AISSTREAM_URL = "wss://stream.aisstream.io/v0/stream"
DEFAULT_BOUNDING_BOXES = [[[-90, -180], [90, 180]]]
DEFAULT_MESSAGE_TYPES = ["PositionReport", "ShipStaticData"]
class AISStreamCollector(BaseCollector):
"""Collect AISStream WebSocket messages into the raw AIS observation layer."""
name = "aisstream_vessels"
priority = "P1"
module = "L4"
frequency_hours = 1
data_type = "vessel_ais"
fail_on_empty = False
async def _load_datasource_config(self) -> DataSourceConfig | None:
if self._db_session is None:
return None
result = await self._db_session.execute(
select(DataSourceConfig)
.where(DataSourceConfig.name == self.name)
.where(DataSourceConfig.is_active.is_(True))
)
return result.scalar_one_or_none()
async def _get_effective_config(self) -> dict[str, Any]:
datasource_config = await self._load_datasource_config()
config = dict(datasource_config.config or {}) if datasource_config else {}
auth_config = dict(datasource_config.auth_config or {}) if datasource_config else {}
endpoint = (
(datasource_config.endpoint if datasource_config else None)
or self._resolved_url
or get_data_sources_config().get_yaml_url(self.name)
or DEFAULT_AISSTREAM_URL
)
api_key = (
auth_config.get("api_key")
or config.get("api_key")
or os.getenv("AISSTREAM_API_KEY")
)
return {
"endpoint": endpoint,
"api_key": api_key,
"bounding_boxes": config.get("bounding_boxes") or DEFAULT_BOUNDING_BOXES,
"message_types": config.get("message_types") or DEFAULT_MESSAGE_TYPES,
"max_messages": int(config.get("max_messages") or 500),
"streaming_enabled": config.get("streaming_enabled", True) is not False,
"streaming_commit_interval": int(config.get("streaming_commit_interval") or 1),
"streaming_max_messages": int(config.get("streaming_max_messages") or 0),
"reconnect_delay_seconds": float(config.get("reconnect_delay_seconds") or 5),
"receive_timeout_seconds": float(config.get("receive_timeout_seconds") or 30),
}
def _build_subscription(self, config: dict[str, Any]) -> dict[str, Any]:
return {
"APIKey": config["api_key"],
"BoundingBoxes": config["bounding_boxes"],
"FilterMessageTypes": config["message_types"],
}
async def fetch(self) -> list[dict[str, Any]]:
config = await self._get_effective_config()
if not config["api_key"]:
raise RuntimeError("AISStream API key is not configured")
try:
import websockets
except ImportError as exc:
raise RuntimeError("Python package 'websockets' is required for AISStream") from exc
subscription = self._build_subscription(config)
messages: list[dict[str, Any]] = []
try:
async with websockets.connect(config["endpoint"]) as websocket:
await websocket.send(json.dumps(subscription))
while len(messages) < config["max_messages"]:
try:
raw_message = await asyncio.wait_for(
websocket.recv(),
timeout=config["receive_timeout_seconds"],
)
except TimeoutError:
break
payload = json.loads(raw_message)
if isinstance(payload, dict):
messages.append(payload)
except Exception as exc:
if self._db_session is not None:
await update_ais_source_health(
self._db_session,
source=self.name,
connection_state="disconnected",
last_error=f"{exc.__class__.__name__}: {exc}",
)
await self._db_session.commit()
raise
return messages
async def run(self, db: AsyncSession) -> dict[str, Any]:
"""Run AISStream as a long-lived streaming collector by default."""
config = await self._get_effective_config()
if not config.get("streaming_enabled", True):
return await super().run(db)
if not config["api_key"]:
return {"status": "failed", "error": "AISStream API key is not configured"}
from app.services.collectors.registry import collector_registry
if not collector_registry.is_active(self.name):
return {"status": "skipped", "reason": "Collector is disabled"}
try:
import websockets
except ImportError as exc:
return {"status": "failed", "error": "Python package 'websockets' is required for AISStream"}
start_time = datetime.now(UTC)
task = CollectionTask(
datasource_id=getattr(self, "_datasource_id", 1),
status="running",
phase="connecting",
phase_message="正在连接 AISStream 实时流",
phase_unit="messages",
started_at=start_time,
)
db.add(task)
await db.commit()
self._current_task = task
self._db_session = db
self._last_broadcast_progress = None
await self.resolve_url(db)
await self._publish_task_update(force=True)
records_added = 0
messages_seen = 0
unique_mmsi: set[str] = set()
reconnect_delay = config["reconnect_delay_seconds"]
try:
while True:
config = await self._get_effective_config()
subscription = self._build_subscription(config)
try:
await update_ais_source_health(
db,
source=self.name,
connection_state="connecting",
)
await self.set_phase("connecting", message="正在连接 AISStream 实时流")
await db.commit()
async with websockets.connect(config["endpoint"]) as websocket:
await websocket.send(json.dumps(subscription))
await update_ais_source_health(
db,
source=self.name,
connection_state="connected",
last_success_at=datetime.now(UTC),
)
await self.set_phase(
"streaming",
message="正在接收 AISStream 实时消息",
reset_progress=False,
)
await db.commit()
while True:
try:
raw_message = await asyncio.wait_for(
websocket.recv(),
timeout=config["receive_timeout_seconds"],
)
except TimeoutError:
await update_ais_source_health(
db,
source=self.name,
connection_state="connected",
last_success_at=datetime.now(UTC),
)
await db.commit()
continue
payload = json.loads(raw_message)
if not isinstance(payload, dict):
continue
messages_seen += 1
record = self._normalize_message(payload)
if not record:
continue
unique_mmsi.add(str(record["mmsi"]))
created = await self._save_stream_record(db, record)
if created:
records_added += 1
task.records_processed = messages_seen
task.total_records = None
task.progress = None
task.phase = "streaming"
task.phase_message = "正在接收 AISStream 实时消息"
task.phase_current = messages_seen
task.phase_total = None
task.phase_unit = "messages"
await self._publish_task_update(force=True)
if config["streaming_max_messages"] and messages_seen >= config["streaming_max_messages"]:
task.status = "success"
task.phase = "stopped"
task.phase_message = "AISStream 测试流已停止"
task.completed_at = datetime.now(UTC)
await db.commit()
await self._publish_task_update(force=True)
return {
"status": "success",
"task_id": task.id,
"records_processed": records_added,
"messages_seen": messages_seen,
"unique_mmsi": len(unique_mmsi),
"execution_time_seconds": (datetime.now(UTC) - start_time).total_seconds(),
}
except asyncio.CancelledError:
raise
except Exception as exc:
await update_ais_source_health(
db,
source=self.name,
connection_state="reconnecting",
last_error=f"{exc.__class__.__name__}: {exc}",
)
task.phase = "reconnecting"
task.phase_message = "AISStream 连接中断,正在重连"
task.error_message = f"{exc.__class__.__name__}: {exc}"
await db.commit()
await self._publish_task_update(force=True)
await asyncio.sleep(reconnect_delay)
except asyncio.CancelledError:
task.status = "cancelled"
task.phase = "stopped"
task.phase_message = "AISStream 实时流已停止"
task.completed_at = datetime.now(UTC)
await update_ais_source_health(
db,
source=self.name,
connection_state="disconnected",
last_error=None,
)
await db.commit()
await self._publish_task_update(force=True)
raise
def transform(self, raw_data: list[dict[str, Any]]) -> list[dict[str, Any]]:
records = []
for item in raw_data:
record = self._normalize_message(item)
if record:
records.append(record)
return records
async def _save_data(
self,
db: AsyncSession,
data: list[dict[str, Any]],
task_id: int | None = None,
snapshot_id: int | None = None,
) -> int:
now = datetime.now(UTC)
records_added = 0
latest_observed_at = now
for index, item in enumerate(data):
observed_at = item.get("received_at") or now
observation = await record_vessel_ais_observation(
db,
source=self.name,
normalized_payload=item,
raw_payload=item.get("_raw_payload") or item,
delivery_mode=AISSTREAM_DELIVERY_MODE,
transport=AISSTREAM_TRANSPORT,
message_type=item.get("_message_type") or "PositionReport",
source_message_id=item.get("_source_message_id"),
observed_at=observed_at,
collected_at=now,
)
if observation is not None:
records_added += 1
if isinstance(observed_at, datetime) and observed_at > latest_observed_at:
latest_observed_at = observed_at
if (index + 1) % 1000 == 0:
await self.update_progress(index + 1, commit=True)
await update_ais_source_health(
db,
source=self.name,
connection_state="connected",
observed_count=len(data),
last_seen_at=latest_observed_at,
last_success_at=now if data else None,
lag_seconds=max((now - latest_observed_at).total_seconds(), 0),
)
await db.commit()
await self.update_progress(records_added, force=True)
return records_added
async def _save_stream_record(self, db: AsyncSession, item: dict[str, Any]) -> bool:
now = datetime.now(UTC)
observed_at = item.get("received_at") or now
observation = await record_vessel_ais_observation(
db,
source=self.name,
normalized_payload=item,
raw_payload=item.get("_raw_payload") or item,
delivery_mode=AISSTREAM_DELIVERY_MODE,
transport=AISSTREAM_TRANSPORT,
message_type=item.get("_message_type") or "PositionReport",
source_message_id=item.get("_source_message_id"),
observed_at=observed_at,
collected_at=now,
)
await update_ais_source_health(
db,
source=self.name,
connection_state="connected",
observed_count=1,
last_seen_at=observed_at if isinstance(observed_at, datetime) else now,
last_success_at=now,
lag_seconds=max((now - observed_at).total_seconds(), 0) if isinstance(observed_at, datetime) else None,
)
await db.commit()
await self._broadcast_vessel_delta(item, created=observation is not None)
return observation is not None
async def _broadcast_vessel_delta(self, item: dict[str, Any], *, created: bool) -> None:
await broadcaster.broadcast_custom(
"vessels",
{
"action": "upsert",
"source": self.name,
"created": created,
"vessels": [
{
"mmsi": item.get("mmsi"),
"mmsi_display": str(item.get("mmsi")) if item.get("mmsi") is not None else None,
"name": item.get("name"),
"lat": item.get("lat"),
"lon": item.get("lon"),
"sog": item.get("sog"),
"cog": item.get("cog"),
"heading": item.get("heading"),
"nav_status": item.get("nav_status"),
"vessel_type": item.get("vessel_type"),
"vessel_type_name": item.get("vessel_type_name"),
"received_at": to_iso8601_utc(item.get("received_at")),
}
],
},
)
def _normalize_message(self, item: dict[str, Any]) -> dict[str, Any] | None:
message_type = str(item.get("MessageType") or item.get("message_type") or "")
metadata = item.get("MetaData") if isinstance(item.get("MetaData"), dict) else {}
message = item.get("Message") if isinstance(item.get("Message"), dict) else {}
body = message.get(message_type) if isinstance(message.get(message_type), dict) else message
if not isinstance(body, dict):
body = {}
mmsi = _as_int(_pick(metadata, "MMSI", "mmsi") or _pick(body, "MMSI", "mmsi"))
if mmsi is None:
return None
received_at = _parse_datetime(
_pick(metadata, "time_utc", "Time_UTC", "timestamp")
or _pick(body, "Timestamp", "timestamp", "time")
)
ship_name = _clean_text(
_pick(body, "Name", "ShipName", "name")
or _pick(metadata, "ShipName", "ship_name", "name")
)
record: dict[str, Any] = {
"mmsi": mmsi,
"received_at": received_at,
"_message_type": message_type or None,
"_source_message_id": item.get("MessageID") or item.get("message_id"),
"_raw_payload": item,
}
lat = _as_float(_pick(body, "Latitude", "lat", "latitude"))
lon = _as_float(_pick(body, "Longitude", "lon", "lng", "longitude"))
if lat is not None and lon is not None:
if not (-90 <= lat <= 90 and -180 <= lon <= 180):
return None
record.update(
{
"lat": lat,
"lon": lon,
"sog": _as_float(_pick(body, "Sog", "SOG", "speedOverGround")),
"cog": _as_float(_pick(body, "Cog", "COG", "courseOverGround")),
"heading": _as_int(_pick(body, "TrueHeading", "Heading", "heading")),
"nav_status": _as_int(_pick(body, "NavigationalStatus", "nav_status")),
}
)
vessel_type = _as_int(_pick(body, "Type", "ShipType", "vessel_type"))
record.update(
{
"name": ship_name,
"callsign": _pick(body, "CallSign", "callsign"),
"imo": _as_int(_pick(body, "ImoNumber", "IMO", "imo")),
"vessel_type": vessel_type,
"vessel_type_name": _pick(body, "TypeName", "ShipTypeName", "vessel_type_name")
or normalize_vessel_type_name(vessel_type),
"length": _as_float(_pick(body, "DimensionToBow", "Length", "length")),
"width": _as_float(_pick(body, "DimensionToPort", "Width", "width")),
}
)
return record
def _pick(item: dict[str, Any], *keys: str) -> Any:
for key in keys:
if key in item and item[key] not in (None, ""):
return item[key]
return None
def _clean_text(value: Any) -> str | None:
if value in (None, ""):
return None
text = str(value).strip()
return text or None
def _as_float(value: Any) -> float | None:
try:
if value in (None, ""):
return None
return float(value)
except (TypeError, ValueError):
return None
def _as_int(value: Any) -> int | None:
try:
if value in (None, ""):
return None
return int(float(value))
except (TypeError, ValueError):
return None
def _parse_datetime(value: Any) -> datetime | None:
if isinstance(value, datetime):
return value if value.tzinfo else value.replace(tzinfo=UTC)
if not value:
return None
if isinstance(value, (int, float)):
timestamp = float(value)
if timestamp > 10_000_000_000:
timestamp /= 1000
return datetime.fromtimestamp(timestamp, UTC)
if isinstance(value, str):
try:
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
return parsed if parsed.tzinfo else parsed.replace(tzinfo=UTC)
except ValueError:
return None
return None

View File

@@ -54,6 +54,11 @@ class BaseCollector(ABC):
"task_id": self._current_task.id,
"status": self._current_task.status,
"phase": self._current_task.phase,
"phase_progress": self._current_task.phase_progress,
"phase_message": self._current_task.phase_message,
"phase_current": self._current_task.phase_current,
"phase_total": self._current_task.phase_total,
"phase_unit": self._current_task.phase_unit,
"progress": progress,
"records_processed": self._current_task.records_processed,
"total_records": self._current_task.total_records,
@@ -80,12 +85,52 @@ class BaseCollector(ABC):
await self._publish_task_update(force=force)
async def set_phase(self, phase: str):
async def set_phase(self, phase: str, *, message: str | None = None, reset_progress: bool = True):
if self._current_task and self._db_session:
self._current_task.phase = phase
self._current_task.phase_message = message
if reset_progress:
self._current_task.phase_progress = None
self._current_task.phase_current = None
self._current_task.phase_total = None
self._current_task.phase_unit = None
await self._db_session.commit()
await self._publish_task_update(force=True)
async def update_phase_progress(
self,
*,
current: int | None = None,
total: int | None = None,
unit: str | None = None,
message: str | None = None,
progress: float | None = None,
commit: bool = False,
force: bool = False,
):
"""Update progress for the current phase without changing task totals."""
if not self._current_task or not self._db_session:
return
if progress is None and current is not None and total and total > 0:
progress = (current / total) * 100
if progress is not None:
self._current_task.phase_progress = max(0.0, min(float(progress), 100.0))
if current is not None:
self._current_task.phase_current = max(0, int(current))
if total is not None:
self._current_task.phase_total = max(0, int(total))
if unit is not None:
self._current_task.phase_unit = unit
if message is not None:
self._current_task.phase_message = message
if commit:
await self._db_session.commit()
await self._publish_task_update(force=force)
@abstractmethod
async def fetch(self) -> List[Dict[str, Any]]:
"""Fetch raw data from source"""
@@ -251,7 +296,7 @@ class BaseCollector(ABC):
await self._publish_task_update(force=True)
try:
await self.set_phase("fetching")
await self.set_phase("fetching", message="正在拉取原始数据")
raw_data = await self.fetch()
task.total_records = len(raw_data)
await db.commit()
@@ -260,15 +305,20 @@ class BaseCollector(ABC):
if self.fail_on_empty and not raw_data:
raise RuntimeError(f"Collector {self.name} returned no data")
await self.set_phase("transforming")
await self.set_phase("transforming", message="正在转换采集数据")
data = self.transform(raw_data)
snapshot_id = await self._create_snapshot(db, task_id, data, start_time)
await self.set_phase("saving")
await self.set_phase("saving", message="正在保存采集数据")
records_count = await self._save_data(db, data, task_id=task_id, snapshot_id=snapshot_id)
task.status = "success"
task.phase = "completed"
task.phase_progress = 100.0
task.phase_message = "采集完成"
task.phase_current = records_count
task.phase_total = records_count
task.phase_unit = "records"
task.records_processed = records_count
task.progress = 100.0
task.completed_at = datetime.now(UTC)
@@ -285,6 +335,7 @@ class BaseCollector(ABC):
await db.rollback()
task.status = "cancelled"
task.phase = "cancelled"
task.phase_message = "采集已取消"
task.error_message = "Collection cancelled by operator and rolled back"
task.completed_at = datetime.now(UTC)
if snapshot_id is not None:
@@ -301,6 +352,7 @@ class BaseCollector(ABC):
await db.rollback()
task.status = "failed"
task.phase = "failed"
task.phase_message = str(e)
task.error_message = str(e)
task.completed_at = datetime.now(UTC)
if snapshot_id is not None:

View File

@@ -13,6 +13,11 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.models.bgp_anomaly import BGPAnomaly
from app.models.bgp_observation import BGPObservation
from app.models.collected_data import CollectedData
from app.services.bgp_collector_locations import (
RIPE_RIS_COLLECTOR_COORDS,
get_bgp_collector_location_dict,
)
from app.services.bgp_event_locations import resolve_bgp_event_geo_dict
from app.services.bgp_incidents import create_bgp_incidents_for_anomalies
from app.services.bgp_detectors import (
detect_mass_withdrawal_anomalies,
@@ -23,32 +28,17 @@ from app.services.bgp_detectors import (
)
from app.services.bgp_enrichment import enrich_bgp_events_for_batch, extract_bgp_network_fields
RIPE_RIS_COLLECTOR_COORDS: dict[str, dict[str, Any]] = {
"rrc00": {"city": "Amsterdam", "country": "Netherlands", "latitude": 52.3676, "longitude": 4.9041},
"rrc01": {"city": "London", "country": "United Kingdom", "latitude": 51.5072, "longitude": -0.1276},
"rrc03": {"city": "Amsterdam", "country": "Netherlands", "latitude": 52.3676, "longitude": 4.9041},
"rrc04": {"city": "Geneva", "country": "Switzerland", "latitude": 46.2044, "longitude": 6.1432},
"rrc05": {"city": "Vienna", "country": "Austria", "latitude": 48.2082, "longitude": 16.3738},
"rrc06": {"city": "Otemachi", "country": "Japan", "latitude": 35.686, "longitude": 139.7671},
"rrc07": {"city": "Stockholm", "country": "Sweden", "latitude": 59.3293, "longitude": 18.0686},
"rrc10": {"city": "Milan", "country": "Italy", "latitude": 45.4642, "longitude": 9.19},
"rrc11": {"city": "New York", "country": "United States", "latitude": 40.7128, "longitude": -74.006},
"rrc12": {"city": "Frankfurt", "country": "Germany", "latitude": 50.1109, "longitude": 8.6821},
"rrc13": {"city": "Moscow", "country": "Russia", "latitude": 55.7558, "longitude": 37.6173},
"rrc14": {"city": "Palo Alto", "country": "United States", "latitude": 37.4419, "longitude": -122.143},
"rrc15": {"city": "Sao Paulo", "country": "Brazil", "latitude": -23.5558, "longitude": -46.6396},
"rrc16": {"city": "Miami", "country": "United States", "latitude": 25.7617, "longitude": -80.1918},
"rrc18": {"city": "Barcelona", "country": "Spain", "latitude": 41.3874, "longitude": 2.1686},
"rrc19": {"city": "Johannesburg", "country": "South Africa", "latitude": -26.2041, "longitude": 28.0473},
"rrc20": {"city": "Zurich", "country": "Switzerland", "latitude": 47.3769, "longitude": 8.5417},
"rrc21": {"city": "Paris", "country": "France", "latitude": 48.8566, "longitude": 2.3522},
"rrc22": {"city": "Bucharest", "country": "Romania", "latitude": 44.4268, "longitude": 26.1025},
"rrc23": {"city": "Singapore", "country": "Singapore", "latitude": 1.3521, "longitude": 103.8198},
"rrc24": {"city": "Montevideo", "country": "Uruguay", "latitude": -34.9011, "longitude": -56.1645},
"rrc25": {"city": "Amsterdam", "country": "Netherlands", "latitude": 52.3676, "longitude": 4.9041},
"rrc26": {"city": "Dubai", "country": "United Arab Emirates", "latitude": 25.2048, "longitude": 55.2708},
}
# Re-exported for backward compatibility with anything that imports
# ``RIPE_RIS_COLLECTOR_COORDS`` from this module. New code should call
# ``app.services.bgp_collector_locations.get_bgp_collector_location_dict()``
# or ``resolve_bgp_collector_location()`` instead — those use the DB-backed
# collector-location cache.
__all__ = [
"RIPE_RIS_COLLECTOR_COORDS",
"normalize_bgp_event",
"save_bgp_observations_for_batch",
"create_bgp_anomalies_for_batch",
]
def _safe_int(value: Any) -> int | None:
@@ -131,7 +121,19 @@ def normalize_bgp_event(payload: dict[str, Any], *, project: str) -> dict[str, A
)
source_id = hashlib.sha1(source_material.encode("utf-8")).hexdigest()[:24]
collector_location = RIPE_RIS_COLLECTOR_COORDS.get(collector, {})
# Routes through the BGP event pipeline: source coords (if any) →
# collector inheritance. Returned dict keeps the legacy
# {city, country, latitude, longitude} keys plus richer
# {precision, source, needs_confirmation, matched_location_name, confidence}.
collector_location = resolve_bgp_event_geo_dict(
collector,
source_latitude=payload.get("latitude"),
source_longitude=payload.get("longitude"),
)
# Empty result (unknown collector & no source coords) — keep the
# downstream-expected dict shape so detectors / serializers don't crash.
if not collector_location:
collector_location = get_bgp_collector_location_dict(collector)
network_fields = extract_bgp_network_fields(prefix)
metadata = {
"project": project,

View File

@@ -108,6 +108,11 @@ class IPtoASNPrefixGeoCollector(BaseCollector):
self._current_task.total_records = total_expected
self._current_task.records_processed = 0
self._current_task.progress = 0.0
self._current_task.phase_progress = 0.0
self._current_task.phase_message = "正在下载 IPtoASN 数据"
self._current_task.phase_current = 0
self._current_task.phase_total = total_expected
self._current_task.phase_unit = "bytes"
await self._db_session.commit()
await self._publish_task_update(force=True)
@@ -135,7 +140,14 @@ class IPtoASNPrefixGeoCollector(BaseCollector):
return
last_emit["value"] = aggregated
last_emit["t"] = now
await self.update_progress(min(aggregated, total_expected), commit=True)
current = min(aggregated, total_expected)
await self.update_phase_progress(
current=current,
total=total_expected,
unit="bytes",
message="正在下载 IPtoASN 数据",
)
await self.update_progress(current, commit=True)
batches = await asyncio.gather(
*(
@@ -148,6 +160,12 @@ class IPtoASNPrefixGeoCollector(BaseCollector):
)
)
if total_expected > 0:
await self.update_phase_progress(
current=total_expected,
total=total_expected,
unit="bytes",
message="IPtoASN 数据下载完成",
)
await self.update_progress(total_expected, commit=True, force=True)
rows: list[dict[str, Any]] = []

View File

@@ -39,12 +39,23 @@ class NRODelegatedPrefixGeoCollector(BaseCollector):
self._current_task.total_records = total_expected
self._current_task.records_processed = 0
self._current_task.progress = 0.0
self._current_task.phase_progress = 0.0
self._current_task.phase_message = "正在下载 NRO delegated 数据"
self._current_task.phase_current = 0
self._current_task.phase_total = total_expected
self._current_task.phase_unit = "bytes"
await self._db_session.commit()
await self._publish_task_update(force=True)
async def on_progress(downloaded: int, total: int | None) -> None:
if not total or total <= 0:
return
await self.update_phase_progress(
current=min(downloaded, total),
total=total,
unit="bytes",
message="正在下载 NRO delegated 数据",
)
await self.update_progress(min(downloaded, total), commit=True)
body_path = await self._downloader.download_file(

View File

@@ -40,12 +40,23 @@ class OpenGeoFeedPrefixGeoCollector(BaseCollector):
self._current_task.total_records = total_expected
self._current_task.records_processed = 0
self._current_task.progress = 0.0
self._current_task.phase_progress = 0.0
self._current_task.phase_message = "正在下载 OpenGeoFeed 数据"
self._current_task.phase_current = 0
self._current_task.phase_total = total_expected
self._current_task.phase_unit = "bytes"
await self._db_session.commit()
await self._publish_task_update(force=True)
async def on_progress(downloaded: int, total: int | None) -> None:
if not total or total <= 0:
return
await self.update_phase_progress(
current=min(downloaded, total),
total=total,
unit="bytes",
message="正在下载 OpenGeoFeed 数据",
)
await self.update_progress(min(downloaded, total), commit=True)
body_path = await self._downloader.download_file(

View File

@@ -0,0 +1,279 @@
"""BarentsWatch AIS collector for vessel tracking."""
from datetime import UTC, datetime
from typing import Any
import httpx
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.time import to_iso8601_utc
from app.core.websocket.broadcaster import broadcaster
from app.services.barentswatch import (
BARENTSWATCH_LATEST_URL,
fetch_barentswatch_access_token,
resolve_barentswatch_config,
)
from app.services.collectors.base import BaseCollector
from app.services.vessel_ais_aggregation import (
BARENTSWATCH_DELIVERY_MODE,
BARENTSWATCH_TRANSPORT,
record_vessel_ais_observation,
update_ais_source_health,
)
from app.services.vessel_types import normalize_vessel_type_name
class VesselAISCollector(BaseCollector):
"""Collect latest AIS positions and append them to vessel tables."""
name = "barentswatch_vessels"
priority = "P1"
module = "L4"
frequency_hours = 1
data_type = "vessel_ais"
@property
def base_url(self) -> str:
return self._resolved_url or BARENTSWATCH_LATEST_URL
async def _get_access_token(self, client: httpx.AsyncClient) -> str | None:
config = await resolve_barentswatch_config(self._db_session)
return await fetch_barentswatch_access_token(client, config)
async def fetch(self) -> list[dict[str, Any]]:
async with httpx.AsyncClient(timeout=60.0) as client:
headers: dict[str, str] = {}
token = await self._get_access_token(client)
if token:
headers["Authorization"] = f"Bearer {token}"
response = await client.get(self.base_url, headers=headers)
if response.status_code == 401 and not token:
return self._get_sample_data()
response.raise_for_status()
payload = response.json()
if isinstance(payload, list):
return [item for item in payload if isinstance(item, dict)]
if isinstance(payload, dict):
for key in ("features", "data", "items", "vessels"):
value = payload.get(key)
if isinstance(value, list):
if key == "features":
return [
{
**(item.get("properties") or {}),
"geometry": item.get("geometry"),
}
for item in value
if isinstance(item, dict)
]
return [item for item in value if isinstance(item, dict)]
return self._get_sample_data()
def transform(self, raw_data: list[dict[str, Any]]) -> list[dict[str, Any]]:
transformed = []
for item in raw_data:
record = self._normalize_record(item)
if record:
transformed.append(record)
return transformed
async def _save_data(
self,
db: AsyncSession,
data: list[dict[str, Any]],
task_id: int | None = None,
snapshot_id: int | None = None,
) -> int:
now = datetime.now(UTC)
records_added = 0
for index, item in enumerate(data):
observed_at = item.get("received_at") or now
await record_vessel_ais_observation(
db,
source=self.name,
normalized_payload=item,
raw_payload=item,
delivery_mode=BARENTSWATCH_DELIVERY_MODE,
transport=BARENTSWATCH_TRANSPORT,
observed_at=observed_at,
collected_at=now,
)
records_added += 1
if (index + 1) % 1000 == 0:
await self.update_progress(index + 1, commit=True)
latest_observed_at = max(
(item.get("received_at") for item in data if item.get("received_at")),
default=now,
)
await update_ais_source_health(
db,
source=self.name,
connection_state="connected",
observed_count=len(data),
last_seen_at=latest_observed_at,
last_success_at=now if data else None,
lag_seconds=max((now - latest_observed_at).total_seconds(), 0),
)
await db.commit()
await self._broadcast_vessel_snapshot(data)
await self.update_progress(records_added, force=True)
return records_added
async def _broadcast_vessel_snapshot(self, data: list[dict[str, Any]]) -> None:
"""Push REST collector updates through the same realtime vessel channel."""
if not data:
return
batch_size = 500
for offset in range(0, len(data), batch_size):
batch = data[offset : offset + batch_size]
await broadcaster.broadcast_custom(
"vessels",
{
"action": "upsert",
"source": self.name,
"created": True,
"vessels": [
{
"mmsi": item.get("mmsi"),
"mmsi_display": str(item.get("mmsi")) if item.get("mmsi") is not None else None,
"name": item.get("name"),
"callsign": item.get("callsign"),
"lat": item.get("lat"),
"lon": item.get("lon"),
"sog": item.get("sog"),
"cog": item.get("cog"),
"heading": item.get("heading"),
"nav_status": item.get("nav_status"),
"vessel_type": item.get("vessel_type"),
"vessel_type_name": item.get("vessel_type_name"),
"received_at": to_iso8601_utc(item.get("received_at")),
}
for item in batch
],
},
)
def _normalize_record(self, item: dict[str, Any]) -> dict[str, Any] | None:
mmsi = _as_int(_pick(item, "mmsi", "MMSI", "Mmsi"))
lat = _as_float(_pick(item, "lat", "latitude", "Latitude"))
lon = _as_float(_pick(item, "lon", "lng", "longitude", "Longitude"))
geometry = item.get("geometry")
coordinates = geometry.get("coordinates") if isinstance(geometry, dict) else None
if (lat is None or lon is None) and isinstance(coordinates, list) and len(coordinates) >= 2:
lon = _as_float(coordinates[0])
lat = _as_float(coordinates[1])
if mmsi is None or lat is None or lon is None:
return None
if not (-90 <= lat <= 90 and -180 <= lon <= 180):
return None
vessel_type = _as_int(_pick(item, "vessel_type", "shipType", "ship_type", "ShipType"))
vessel_type_name = (
_pick(item, "vessel_type_name", "shipTypeName", "ship_type_name", "VesselTypeName")
or normalize_vessel_type_name(vessel_type)
)
received_at = _parse_datetime(_pick(item, "received_at", "timestamp", "time", "msgtime"))
return {
"mmsi": mmsi,
"name": _pick(item, "name", "shipName", "ship_name", "Name"),
"callsign": _pick(item, "callsign", "callSign", "CallSign"),
"vessel_type": vessel_type,
"vessel_type_name": vessel_type_name,
"flag": _pick(item, "flag", "country", "Flag"),
"length": _as_float(_pick(item, "length", "shipLength", "Length")),
"width": _as_float(_pick(item, "width", "shipWidth", "Width")),
"draught": _as_float(_pick(item, "draught", "draft", "Draught")),
"imo": _as_int(_pick(item, "imo", "IMO", "imoNumber")),
"lat": lat,
"lon": lon,
"sog": _as_float(_pick(item, "sog", "speedOverGround", "SOG")),
"cog": _as_float(_pick(item, "cog", "courseOverGround", "COG")),
"heading": _as_int(_pick(item, "heading", "trueHeading", "Heading")),
"nav_status": _as_int(_pick(item, "nav_status", "navStatus", "NavigationalStatus")),
"received_at": received_at,
}
def _get_sample_data(self) -> list[dict[str, Any]]:
return [
{
"mmsi": 257123000,
"name": "OSLO TRADER",
"lat": 59.91,
"lon": 10.73,
"sog": 12.4,
"cog": 214,
"heading": 215,
"nav_status": 0,
"vessel_type": 70,
"vessel_type_name": "Cargo",
"flag": "NO",
"length": 185,
},
{
"mmsi": 257456000,
"name": "NORDIC FJORD",
"lat": 60.39,
"lon": 5.32,
"sog": 0.2,
"cog": 82,
"heading": 80,
"nav_status": 1,
"vessel_type": 60,
"vessel_type_name": "Passenger",
"flag": "NO",
"length": 126,
},
]
def _pick(item: dict[str, Any], *keys: str) -> Any:
for key in keys:
if key in item and item[key] not in (None, ""):
return item[key]
return None
def _as_float(value: Any) -> float | None:
try:
if value in (None, ""):
return None
return float(value)
except (TypeError, ValueError):
return None
def _as_int(value: Any) -> int | None:
try:
if value in (None, ""):
return None
return int(float(value))
except (TypeError, ValueError):
return None
def _parse_datetime(value: Any) -> datetime | None:
if isinstance(value, datetime):
return value if value.tzinfo else value.replace(tzinfo=UTC)
if not value:
return None
if isinstance(value, (int, float)):
timestamp = float(value)
if timestamp > 10_000_000_000:
timestamp /= 1000
return datetime.fromtimestamp(timestamp, UTC)
if isinstance(value, str):
try:
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
return parsed if parsed.tzinfo else parsed.replace(tzinfo=UTC)
except ValueError:
return None
return None

View File

@@ -0,0 +1,886 @@
"""Compute-center location resolver, built on the shared location pipeline.
This module is a thin domain wrapper that wires up
:mod:`app.services.location` for compute centers:
SourceCoordinates
The online Nominatim step is intentionally reserved for the user-triggered
``collect-location`` flow. The regular GeoJSON endpoint runs during Earth
startup, so it must stay local and deterministic.
For the full design and the reason behind the abstraction (compute centers,
BGP collectors, BGP events, and future entities all share one pipeline),
see ``docs/plans/location-resolver-shared-pipeline-plan.md``.
The ``ComputeCenterLocation`` dataclass and the public function signatures are
preserved verbatim so existing callers and tests do not need to change.
"""
from __future__ import annotations
from dataclasses import dataclass
from datetime import UTC, datetime
from functools import lru_cache
from typing import Any
import httpx
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.collected_data_fields import get_record_field
from app.models.collected_data import CollectedData
from app.models.compute_center_location import ComputeCenterLocationRecord
from app.services.location import (
LocationCandidate,
LocationPipeline,
LocationQuery,
NominatimResolver,
ResolverOutput,
SourceCoordinatesResolver,
build_default_nominatim_geocoder,
coerce_str,
normalize_country_text,
normalize_text,
parse_float,
)
ROR_SEARCH_URL = "https://api.ror.org/v2/organizations"
DEFAULT_ROR_USER_AGENT = "planet-earth-location-resolver/1.0"
DEFAULT_ROR_TIMEOUT_SECONDS = 8.0
RENDERABLE_PRECISIONS: tuple[str, ...] = ("precise", "site", "city")
FORBIDDEN_PRECISIONS: tuple[str, ...] = (
"country",
"estimated_country",
"country_major_compute_city",
"region",
"unknown",
)
# ── Public dataclasses ──────────────────────────────────────────────
@dataclass(frozen=True)
class ComputeCenterLocation:
latitude: float | None
longitude: float | None
location_precision: str
geography_mode: str
is_estimated: bool
estimated_reason: str | None = None
location_confidence: float | None = None
location_source: str | None = None
location_source_note: str | None = None
location_verified_at: str | None = None
matched_location_name: str | None = None
needs_confirmation: bool = False
city: str | None = None
region: str | None = None
country: str | None = None
@property
def is_renderable(self) -> bool:
if self.latitude in (None, 0.0) or self.longitude in (None, 0.0):
return False
return self.location_precision in RENDERABLE_PRECISIONS
def to_geojson_properties(self) -> dict[str, Any]:
return {
"latitude": self.latitude,
"longitude": self.longitude,
"location_precision": self.location_precision,
"geography_mode": self.geography_mode,
"is_estimated": self.is_estimated,
"estimated_reason": self.estimated_reason,
"location_confidence": self.location_confidence,
"location_source": self.location_source,
"location_source_note": self.location_source_note,
"location_verified_at": self.location_verified_at,
"matched_location_name": self.matched_location_name,
"needs_confirmation": self.needs_confirmation,
}
@dataclass(frozen=True)
class ResolutionDiagnostic:
failure_reason: str
attempted_queries: tuple[str, ...] = ()
record_id: int | None = None
source: str | None = None
source_id: str | None = None
name: str | None = None
country: str | None = None
city: str | None = None
site: str | None = None
operator: str | None = None
def to_dict(self) -> dict[str, Any]:
return {
"failure_reason": self.failure_reason,
"attempted_queries": list(self.attempted_queries),
"record_id": self.record_id,
"source": self.source,
"source_id": self.source_id,
"name": self.name,
"country": self.country,
"city": self.city,
"site": self.site,
"operator": self.operator,
}
@dataclass(frozen=True)
class ResolutionResult:
location: ComputeCenterLocation | None
diagnostic: ResolutionDiagnostic | None
@property
def is_resolved(self) -> bool:
return bool(self.location and self.location.is_renderable)
# ── Geocoder (kept at module level so tests can monkeypatch + cache_clear) ──
_geocode_online = build_default_nominatim_geocoder()
# ── Stored location cache ───────────────────────────────────────────
COMPUTE_CENTER_LOCATION_CACHE: dict[str, dict[str, Any]] = {}
def _cache_key(source: str | None, source_id: str | None) -> str:
return f"{coerce_str(source)}:{coerce_str(source_id)}"
def set_compute_center_location_cache(
locations: dict[str, dict[str, Any]],
) -> None:
COMPUTE_CENTER_LOCATION_CACHE.clear()
COMPUTE_CENTER_LOCATION_CACHE.update(
{coerce_str(key): dict(value) for key, value in locations.items()}
)
async def refresh_compute_center_location_cache(
session: AsyncSession,
) -> dict[str, dict[str, Any]]:
result = await session.execute(select(ComputeCenterLocationRecord))
records = result.scalars().all()
cache = {}
for record in records:
if not hasattr(record, "to_location_dict"):
continue
if not record.source or not record.source_id:
continue
cache[_cache_key(record.source, record.source_id)] = record.to_location_dict()
set_compute_center_location_cache(cache)
return cache
def get_compute_center_location_dict(
source: str | None,
source_id: str | None,
) -> dict[str, Any]:
return dict(COMPUTE_CENTER_LOCATION_CACHE.get(_cache_key(source, source_id), {}))
# ── Pipeline construction ──────────────────────────────────────────
@lru_cache(maxsize=512)
def _lookup_ror_organization(query: str) -> dict[str, Any] | None:
"""Lookup a research organization in ROR for user-triggered candidates."""
if not query:
return None
response = httpx.get(
ROR_SEARCH_URL,
params={"query": query},
headers={"User-Agent": DEFAULT_ROR_USER_AGENT},
timeout=DEFAULT_ROR_TIMEOUT_SECONDS,
)
response.raise_for_status()
payload = response.json()
items = payload.get("items") if isinstance(payload, dict) else None
if not isinstance(items, list) or not items:
return None
first = items[0]
if not isinstance(first, dict):
return None
organization = first.get("organization")
if isinstance(organization, dict):
return organization
return first
def _compute_center_ror_query_plan(
query: LocationQuery,
) -> list[tuple[str, tuple[str, ...]]]:
extra = query.extra or {}
raw_parts: list[tuple[str, str]] = [
("site", coerce_str(extra.get("site"))),
("operator", coerce_str(extra.get("operator"))),
("organization", coerce_str(extra.get("organization"))),
]
for field, value in tuple(raw_parts):
if "/" not in value:
continue
raw_parts.extend(
(field, part.strip())
for part in value.split("/")
if len(part.strip()) >= 3
)
plan: list[tuple[str, tuple[str, ...]]] = []
seen: set[str] = set()
for field, value in raw_parts:
key = normalize_text(value)
if not key or key in seen:
continue
seen.add(key)
plan.append((value, (field,)))
return plan
def _organization_label(organization: dict[str, Any], fallback: str) -> str:
names = organization.get("names")
if isinstance(names, list):
for name in names:
if not isinstance(name, dict):
continue
types = name.get("types")
if isinstance(types, list) and "ror_display" in types:
value = coerce_str(name.get("value"))
if value:
return value
for name in names:
if isinstance(name, dict):
value = coerce_str(name.get("value"))
if value:
return value
return fallback
class ROROrganizationResolver:
"""Resolve source-provided organization/site text through the open ROR API."""
name = "ror_organization_registry"
def __init__(
self,
*,
query_plan_builder=_compute_center_ror_query_plan,
lookup=lambda q: _lookup_ror_organization(q),
confidence: float = 0.68,
) -> None:
self._query_plan_builder = query_plan_builder
self._lookup = lookup
self._confidence = confidence
def resolve(self, query: LocationQuery):
from app.services.location import ResolverOutput
from app.services.location.text import parse_float
attempted: list[str] = []
candidates: list[LocationCandidate] = []
context_country = normalize_text(normalize_country_text(query.country))
for ror_query, matched_fields in self._query_plan_builder(query):
attempted.append(f"ror:{ror_query}")
try:
organization = self._lookup(ror_query)
except Exception:
continue
if not isinstance(organization, dict):
continue
locations = organization.get("locations")
if not isinstance(locations, list) or not locations:
continue
location = locations[0]
if not isinstance(location, dict):
continue
details = location.get("geonames_details")
if not isinstance(details, dict):
continue
latitude = parse_float(details.get("lat"))
longitude = parse_float(details.get("lng"))
if latitude in (None, 0.0) or longitude in (None, 0.0):
continue
country = normalize_country_text(details.get("country_name"))
if context_country and normalize_text(country) != context_country:
continue
city = coerce_str(details.get("name")) or None
region = coerce_str(details.get("country_subdivision_name")) or None
display_name = _organization_label(organization, ror_query)
ror_id = coerce_str(organization.get("id"))
geonames_id = location.get("geonames_id")
source_note = (
f"ROR organization match: {display_name}"
+ (f" ({ror_id})" if ror_id else "")
+ (f"; GeoNames {geonames_id}" if geonames_id else "")
)
candidates.append(
LocationCandidate(
latitude=latitude,
longitude=longitude,
display_name=display_name,
precision="city",
confidence=self._confidence,
query=ror_query,
source=self.name,
source_note=source_note,
matched_fields=matched_fields,
needs_confirmation=True,
city=city,
region=region,
country=country or query.country,
matched_location_name=display_name,
location_verified_at=None,
suggested_registry_entry=None,
)
)
return ResolverOutput(
candidates=tuple(candidates),
attempted_queries=tuple(attempted),
)
class StoredComputeCenterLocationResolver:
"""Resolve a compute center through the DB-backed current-location cache."""
name = "stored_compute_center_location"
def resolve(self, query: LocationQuery) -> ResolverOutput:
extra = query.extra or {}
stored = get_compute_center_location_dict(
coerce_str(extra.get("source")),
coerce_str(extra.get("source_id")),
)
if not stored:
return ResolverOutput()
latitude = parse_float(stored.get("latitude"))
longitude = parse_float(stored.get("longitude"))
if latitude in (None, 0.0) or longitude in (None, 0.0):
return ResolverOutput()
return ResolverOutput(
candidates=(
LocationCandidate(
latitude=latitude,
longitude=longitude,
display_name=stored.get("name") or query.name or "Compute center",
precision=stored.get("precision") or "city",
confidence=float(stored.get("confidence") or 0.85),
query=f"stored_compute_center_location::{stored.get('source')}:{stored.get('source_id')}",
source=self.name,
source_note=stored.get("source_note"),
matched_fields=("source", "source_id"),
needs_confirmation=bool(stored.get("needs_confirmation")),
city=stored.get("city") or query.city,
region=None,
country=stored.get("country") or query.country,
matched_location_name=stored.get("site") or stored.get("name") or query.name,
location_verified_at=stored.get("verified_at"),
suggested_registry_entry=None,
),
)
)
def _short_system_name(name: Any) -> str:
"""Strip vendor/system suffix from TOP500 names like ``"El Capitan - HPE Cray ..."``."""
text = coerce_str(name)
if not text:
return ""
head = text.split(" - ", 1)[0].strip()
return head or text
def _record_context(record: Any, metadata: dict[str, Any]) -> dict[str, str]:
name = coerce_str(getattr(record, "name", None))
return {
"source": coerce_str(getattr(record, "source", None)),
"source_id": coerce_str(getattr(record, "source_id", None)),
"name": name,
"name_short": _short_system_name(name),
"city": coerce_str(get_record_field(record, "city")),
"country": coerce_str(get_record_field(record, "country")),
"site": coerce_str(metadata.get("site") or metadata.get("organization")),
"operator": coerce_str(
metadata.get("operator")
or metadata.get("organization")
or metadata.get("owner")
or metadata.get("manufacturer")
),
"organization": coerce_str(metadata.get("organization")),
}
def _context_to_query(
context: dict[str, str],
*,
source_lat: float | None = None,
source_lon: float | None = None,
) -> LocationQuery:
name = context.get("name") or None
name_short = context.get("name_short") or ""
aliases: tuple[str, ...] = ()
if name_short and name_short != name:
aliases = (name_short,)
return LocationQuery(
name=name,
aliases=aliases,
city=context.get("city") or None,
country=context.get("country") or None,
source_latitude=source_lat,
source_longitude=source_lon,
extra={
"source": context.get("source") or "",
"source_id": context.get("source_id") or "",
"site": context.get("site") or "",
"operator": context.get("operator") or "",
"organization": context.get("organization") or "",
},
)
def _compute_center_query_plan(
query: LocationQuery,
) -> list[tuple[str, tuple[str, ...]]]:
"""Build the Nominatim query plan for a compute-center query.
Mirrors the legacy ``_build_online_query_plan`` ordering exactly.
"""
name = query.name or ""
name_short = (query.aliases[0] if query.aliases else "") or name
extra = query.extra or {}
site = str(extra.get("site") or "")
operator = str(extra.get("operator") or "")
city = query.city or ""
country = query.country or ""
plan: list[tuple[str, tuple[str, ...]]] = []
def add(parts: list[tuple[str, str]]) -> None:
non_empty = [(field, value) for field, value in parts if value]
if not non_empty:
return
seen: set[str] = set()
cleaned: list[str] = []
fields: list[str] = []
for field, value in non_empty:
key = normalize_text(value)
if not key or key in seen:
continue
seen.add(key)
cleaned.append(value)
fields.append(field)
if not cleaned:
return
composed = ", ".join(cleaned)
if not any(composed == existing for existing, _ in plan):
plan.append((composed, tuple(fields)))
add([("site", site), ("country", country)])
add([("operator", operator), ("city", city), ("country", country)])
add([("name", name_short), ("operator", operator), ("country", country)])
add([("name", name_short), ("site", site)])
add([("name", name_short), ("country", country)])
add([("name", name_short), ("city", city), ("country", country)])
add([("city", city), ("country", country)])
if name and name != name_short:
add([("name", name), ("country", country)])
return plan
COMPUTE_CENTER_PIPELINE = LocationPipeline(
[
SourceCoordinatesResolver(),
StoredComputeCenterLocationResolver(),
],
failure_reason=(
"Could not resolve to city-level coordinates from source coords"
" or stored compute-center location."
),
)
COMPUTE_CENTER_COLLECTION_PIPELINE = LocationPipeline(
[
SourceCoordinatesResolver(),
ROROrganizationResolver(),
NominatimResolver(
query_plan_builder=_compute_center_query_plan,
# Late-binding so test monkeypatching of ``_geocode_online`` works.
geocoder=lambda q: _geocode_online(q),
),
],
failure_reason=(
"Could not resolve to city-level coordinates from source coords"
", ROR organization lookup, or online geocoding."
),
)
# ── Candidate → ComputeCenterLocation conversion ───────────────────
_GEOGRAPHY_MODE_BY_SOURCE = {
"source_coordinates": "source_coordinates",
"stored_compute_center_location": "stored_compute_center_location",
"ror_organization_registry": "ror_organization",
"nominatim_online_geocode": "online_geocode",
}
def _candidate_to_location(
candidate: LocationCandidate,
*,
context: dict[str, str],
) -> ComputeCenterLocation:
geography_mode = _GEOGRAPHY_MODE_BY_SOURCE.get(candidate.source, "online_geocode")
is_estimated = candidate.needs_confirmation or candidate.source.startswith(
"nominatim"
)
estimated_reason: str | None
if candidate.source == "source_coordinates":
estimated_reason = None
elif candidate.source == "stored_compute_center_location":
estimated_reason = candidate.source_note
elif candidate.source == "ror_organization_registry":
fields_summary = ", ".join(candidate.matched_fields) or "organization"
estimated_reason = (
f"Resolved by ROR organization lookup '{candidate.query}' "
f"(matched fields: {fields_summary})"
)
elif candidate.source == "nominatim_online_geocode":
fields_summary = ", ".join(candidate.matched_fields) or "name"
estimated_reason = (
f"Resolved by online geocoding query '{candidate.query}' "
f"(matched fields: {fields_summary})"
)
else:
estimated_reason = candidate.source_note
country = (
candidate.country
or normalize_country_text(context.get("country"))
or context.get("country")
or None
)
return ComputeCenterLocation(
latitude=candidate.latitude,
longitude=candidate.longitude,
location_precision=candidate.precision,
geography_mode=geography_mode,
is_estimated=is_estimated,
estimated_reason=estimated_reason,
location_confidence=candidate.confidence,
location_source=candidate.source,
location_source_note=candidate.source_note,
location_verified_at=candidate.location_verified_at,
matched_location_name=candidate.matched_location_name
or context.get("name")
or None,
needs_confirmation=candidate.needs_confirmation,
city=candidate.city or context.get("city") or None,
region=candidate.region,
country=country,
)
def _diagnostic_for(
record: Any,
context: dict[str, str],
*,
failure_reason: str,
attempted_queries: tuple[str, ...] = (),
) -> ResolutionDiagnostic:
return ResolutionDiagnostic(
failure_reason=failure_reason,
attempted_queries=attempted_queries,
record_id=getattr(record, "id", None),
source=getattr(record, "source", None),
source_id=getattr(record, "source_id", None),
name=context.get("name") or getattr(record, "name", None),
country=context.get("country") or None,
city=context.get("city") or None,
site=context.get("site") or None,
operator=context.get("operator") or None,
)
# ── Public API ─────────────────────────────────────────────────────
def resolve_compute_center_location(
record: Any,
metadata: dict[str, Any] | None = None,
) -> ComputeCenterLocation:
"""Backwards-compatible thin wrapper returning the renderable location only.
Records that cannot be resolved to city-level get a placeholder
:class:`ComputeCenterLocation` with ``location_precision='unknown'``.
Callers should generally prefer :func:`resolve_compute_center_location_full`.
"""
full = resolve_compute_center_location_full(record, metadata)
return full.location or ComputeCenterLocation(
latitude=None,
longitude=None,
location_precision="unknown",
geography_mode="unresolved",
is_estimated=True,
estimated_reason="No resolvable location hints",
location_confidence=0.0,
location_source="unknown",
location_source_note=(
"No source coordinates, ROR organization match, or online"
" geocoding result."
),
matched_location_name=None,
needs_confirmation=False,
)
def resolve_compute_center_location_full(
record: Any,
metadata: dict[str, Any] | None = None,
*,
allow_online: bool = False,
) -> ResolutionResult:
metadata = metadata or {}
context = _record_context(record, metadata)
from app.services.location.text import parse_float as _parse_float
source_lat = _parse_float(get_record_field(record, "latitude"))
source_lon = _parse_float(get_record_field(record, "longitude"))
if source_lat in (None, 0.0):
source_lat = None
if source_lon in (None, 0.0):
source_lon = None
query = _context_to_query(
context, source_lat=source_lat, source_lon=source_lon
)
pipeline = (
COMPUTE_CENTER_COLLECTION_PIPELINE
if allow_online
else COMPUTE_CENTER_PIPELINE
)
pipeline_result = pipeline.resolve_best(query)
if pipeline_result.location and pipeline_result.location.precision in RENDERABLE_PRECISIONS:
location = _candidate_to_location(pipeline_result.location, context=context)
return ResolutionResult(location=location, diagnostic=None)
return ResolutionResult(
location=None,
diagnostic=_diagnostic_for(
record,
context,
failure_reason=(
"Could not resolve to city-level coordinates from source coords"
", ROR organization lookup, or online geocoding."
if allow_online
else (
"Could not resolve to city-level coordinates from source coords"
" or stored compute-center location."
)
),
attempted_queries=pipeline_result.attempted_queries,
),
)
def collect_location_candidates(
*,
name: str | None = None,
source: str | None = None,
source_id: str | None = None,
operator: str | None = None,
site: str | None = None,
city: str | None = None,
country: str | None = None,
organization: str | None = None,
record_id: int | None = None,
) -> tuple[list[LocationCandidate], list[str]]:
"""Run the full resolution chain and return ranked candidates with attempted queries.
The unused ``source`` / ``source_id`` / ``record_id`` arguments are kept
for backward compatibility with the API handler that calls this function.
"""
query = build_compute_center_location_query(
name=name,
source=source,
source_id=source_id,
operator=operator,
site=site,
city=city,
country=country,
organization=organization,
)
return COMPUTE_CENTER_COLLECTION_PIPELINE.collect_candidates(query)
def build_compute_center_location_query(
*,
name: str | None = None,
source: str | None = None,
source_id: str | None = None,
operator: str | None = None,
site: str | None = None,
city: str | None = None,
country: str | None = None,
organization: str | None = None,
) -> LocationQuery:
name_value = coerce_str(name)
context: dict[str, str] = {
"source": coerce_str(source),
"source_id": coerce_str(source_id),
"name": name_value,
"name_short": _short_system_name(name_value),
"city": coerce_str(city),
"country": coerce_str(country),
"site": coerce_str(site or organization),
"operator": coerce_str(operator or organization),
"organization": coerce_str(organization),
}
return _context_to_query(context)
def _record_operator(metadata: dict[str, Any]) -> str | None:
return coerce_str(
metadata.get("operator")
or metadata.get("organization")
or metadata.get("owner")
or metadata.get("manufacturer")
) or None
async def seed_compute_center_locations_from_source_coords(
session: AsyncSession,
) -> None:
"""Seed stored compute-center locations only from real source coordinates."""
stmt = (
select(CollectedData)
.where(CollectedData.source.in_(["top500", "epoch_ai_gpu"]))
.where(CollectedData.is_current.is_(True))
)
result = await session.execute(stmt)
records = result.scalars().all()
changed = False
for record in records:
source_value = coerce_str(getattr(record, "source", None))
source_id = coerce_str(getattr(record, "source_id", None))
if not source_value or not source_id:
continue
latitude = parse_float(get_record_field(record, "latitude"))
longitude = parse_float(get_record_field(record, "longitude"))
if latitude in (None, 0.0) or longitude in (None, 0.0):
continue
existing = await session.scalar(
select(ComputeCenterLocationRecord)
.where(ComputeCenterLocationRecord.source == source_value)
.where(ComputeCenterLocationRecord.source_id == source_id)
)
if existing:
continue
metadata = record.extra_data or {}
session.add(
ComputeCenterLocationRecord(
source=source_value,
source_id=source_id,
name=getattr(record, "name", None),
operator=_record_operator(metadata),
site=coerce_str(metadata.get("site") or metadata.get("organization")) or None,
city=coerce_str(get_record_field(record, "city")) or None,
country=coerce_str(get_record_field(record, "country")) or None,
latitude=latitude,
longitude=longitude,
precision="precise",
confidence=1.0,
location_source="source_coordinates",
source_note="Seeded from source-provided compute-center coordinates",
raw_payload={
"record_id": getattr(record, "id", None),
"source": source_value,
"source_id": source_id,
},
needs_confirmation=False,
verification_status="source_provided",
verified_at=None,
)
)
changed = True
if changed:
await session.commit()
await refresh_compute_center_location_cache(session)
async def upsert_compute_center_location(
session: AsyncSession,
*,
source: str,
source_id: str,
name: str | None = None,
operator: str | None = None,
site: str | None = None,
city: str | None = None,
country: str | None = None,
latitude: float,
longitude: float,
precision: str = "city",
confidence: float | None = None,
location_source: str = "manual_selection",
source_url: str | None = None,
source_note: str | None = None,
raw_payload: dict[str, Any] | None = None,
needs_confirmation: bool = False,
verification_status: str = "verified",
) -> ComputeCenterLocationRecord:
existing = await session.scalar(
select(ComputeCenterLocationRecord)
.where(ComputeCenterLocationRecord.source == source)
.where(ComputeCenterLocationRecord.source_id == source_id)
)
verified_at = None if needs_confirmation else datetime.now(UTC)
values = {
"name": name,
"operator": operator,
"site": site,
"city": city,
"country": country,
"latitude": latitude,
"longitude": longitude,
"precision": precision,
"confidence": confidence,
"location_source": location_source,
"source_url": source_url,
"source_note": source_note,
"raw_payload": raw_payload or {},
"needs_confirmation": needs_confirmation,
"verification_status": verification_status,
"verified_at": verified_at,
}
if existing:
for key, value in values.items():
setattr(existing, key, value)
record = existing
else:
record = ComputeCenterLocationRecord(
source=source,
source_id=source_id,
**values,
)
session.add(record)
await session.commit()
await session.refresh(record)
await refresh_compute_center_location_cache(session)
return record

View File

@@ -0,0 +1,290 @@
"""Credential setup guides for collector integrations."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
from sqlalchemy import select
from app.models.system_setting import SystemSetting
from app.schemas.ai import SituationalAnalysisRequest
from app.services.ai_client import AIProviderClient
from app.services.ai_tools.evidence_store import normalize_search_evidence
from app.services.ai_tools.web_search import WebSearchClient, WebSearchError
CREDENTIAL_GUIDES_CATEGORY = "collector_credential_guides"
@dataclass(frozen=True)
class CredentialGuideDefault:
provider: str
title: str
prompt: str
markdown: str
BARENTSWATCH_DEFAULT_GUIDE = CredentialGuideDefault(
provider="barentswatch",
title="BarentsWatch AIS 凭证获取教程",
prompt=(
"请生成一份中文教程,指导开发者获取 BarentsWatch Live AIS API 的 "
"OAuth client credentials。教程要面向已经有本地开发环境的人包含注册/登录、"
"创建 client、申请或确认 ais scope、复制 client id 和 client secret、"
"在系统设置中填写并验证连接、常见失败排查。不要编造具体页面按钮文案,"
"必须参考官方 tutorialhttps://developer.barentswatch.no/docs/tutorial 。"
"必须强调 Live AIS 要选择 AIS-client / AIS - API而不是普通 API-client。"
"如果步骤可能变化,要提醒以 BarentsWatch developer portal 当前页面为准。"
),
markdown="""## BarentsWatch AIS 凭证获取
官方教程https://developer.barentswatch.no/docs/tutorial
1. 先打开上面的 BarentsWatch 官方 tutorial按官方流程登录或注册开发者账号。
2. 在 Developer access 页面选择 `AIS - API`,不要选择普通的 `BarentsWatch - API`。
3. 在 `AIS - API` 下创建用于 Planet 的 AIS client。
4. 创建时记下你设置的 password / client secret。
5. 回到 My Page 复制完整 `Client ID`。它通常长得像 `your.email@example.com:client-name`。
6. 回到 Planet 的 `设置 -> 采集器设置 -> BarentsWatch AIS`,填入 `Client ID` 和 `Client Secret`。
7. 点击 `连接` 验证 token 和 AIS endpoint 是否可访问。
8. 连接成功后保存凭证。
### 请求规则
- Token 地址:`https://id.barentswatch.no/connect/token`
- 请求方式:`POST`
- Content-Type`application/x-www-form-urlencoded`
- Body 必须包含:`grant_type=client_credentials`、`client_id`、`client_secret`、`scope=ais`
- `client_id`、`client_secret`、`scope`、`grant_type` 都要放在 body不要放在 header。
- AIS 数据请求使用 header`Authorization: Bearer <access_token>`
### 常见排查
- `未找到凭证`:确认 `Client ID` 和 `Client Secret` 已填写,或已经写入 `~/.zshrc`。
- `HTTP 401/403`:通常是选成了普通 `BarentsWatch - API` client、client secret 错误,或 token 请求没有使用 `scope=ais`。
- `network` 错误:检查本机是否能访问 `id.barentswatch.no` 和 `live.ais.barentswatch.no`。
- Endpoint 建议保持默认:`https://live.ais.barentswatch.no/v1/latest/combined`。
""",
)
AISSTREAM_DEFAULT_GUIDE = CredentialGuideDefault(
provider="aisstream",
title="AISStream API Key 获取教程",
prompt=(
"请生成一份中文教程,指导开发者获取 AISStream 的 API Key 并配置到 Planet。"
"教程要面向已经有本地开发环境的人,包含注册/登录 AISStream、获取 API Key、"
"理解免费额度和订阅范围、在 Planet 设置中心填写 API Key、配置 bounding boxes "
"和 message types、验证连接、常见失败排查。必须提醒用户以 AISStream 当前官网和"
"服务条款为准,不要编造具体页面按钮文案。"
),
markdown="""## AISStream API Key 获取
官方入口https://aisstream.io/
1. 打开 AISStream 官网,按当前页面指引注册或登录账号。
2. 在账号/API 管理页面创建或复制你的 API Key。
3. 先确认当前账号额度、使用条款和可订阅区域。实时 AIS 流量可能很大,不建议一开始订阅全球范围。
4. 回到 Planet 的 `设置 -> 采集器设置 -> AISStream 实时船舶`。
5. 在 `AISStream 凭证` 中填入 API Key。
6. Endpoint 通常保持默认:`wss://stream.aisstream.io/v0/stream`。
7. 按需配置 `Bounding Boxes JSON` 和 `消息类型`。
8. 点击连接测试,确认系统能读取凭证且 WebSocket endpoint 格式有效。
9. 保存采集器设置后再运行 `aisstream_vessels` collector。
### 推荐配置
默认消息类型:
```json
["PositionReport", "ShipStaticData"]
```
默认 Bounding Boxes 示例:
```json
[[[-90, -180], [90, 180]]]
```
这个示例表示全球范围。实际使用时建议先改成较小区域,降低消息量和处理压力。
### 请求规则
- Endpoint`wss://stream.aisstream.io/v0/stream`
- 传输方式WebSocket
- API Key 放在订阅 payload 中,不放在 HTTP header。
- Planet 会把 AISStream 标记为 `delivery_mode = realtime_stream`、`transport = websocket`。
- AISStream collector 只写入 AIS raw observations不直接覆盖最终船只展示表。
### 常见排查
- `未找到凭证`:确认 API Key 已保存到采集器设置,或设置了 `AISSTREAM_API_KEY` 环境变量 / `~/.zshrc`。
- `endpoint 必须是 ws:// 或 wss://`AISStream 是 WebSocket 流接口,不要填普通 `https://` API 地址。
- 采集量过大:缩小 `Bounding Boxes JSON`,减少 `message_types`,或降低单次最大消息数。
- 没有船只数据:确认订阅区域内确实有 AIS 活动,并检查 API Key 当前额度和权限。
- 连接中断实时流可能受网络和上游限流影响collector 会记录源健康状态供聚合服务回退。
""",
)
DEFAULT_CREDENTIAL_GUIDES = {
BARENTSWATCH_DEFAULT_GUIDE.provider: BARENTSWATCH_DEFAULT_GUIDE,
AISSTREAM_DEFAULT_GUIDE.provider: AISSTREAM_DEFAULT_GUIDE,
}
async def _get_guide_store(db) -> tuple[SystemSetting | None, dict[str, Any]]:
result = await db.execute(
select(SystemSetting).where(SystemSetting.category == CREDENTIAL_GUIDES_CATEGORY)
)
record = result.scalar_one_or_none()
payload = dict(record.payload or {}) if record and isinstance(record.payload, dict) else {}
return record, payload
async def get_credential_guide(db, provider: str) -> dict[str, Any]:
default = DEFAULT_CREDENTIAL_GUIDES.get(provider)
if default is None:
raise ValueError(f"Unsupported credential guide provider: {provider}")
_record, store = await _get_guide_store(db)
custom = store.get(provider) if isinstance(store.get(provider), dict) else None
return {
"provider": provider,
"title": custom.get("title") if custom else default.title,
"markdown": custom.get("markdown") if custom else default.markdown,
"prompt": default.prompt,
"source": "ai" if custom else "default",
"sources": custom.get("sources", []) if custom else [],
"verification_status": (
custom.get("verification_status", "verified_with_search_evidence")
if custom
else "default_unverified"
),
"verification_error": custom.get("verification_error") if custom else None,
}
async def save_credential_guide(
db,
provider: str,
title: str,
markdown: str,
*,
sources: list[dict[str, Any]] | None = None,
verification_status: str = "verified_with_search_evidence",
verification_error: str | None = None,
) -> dict[str, Any]:
default = DEFAULT_CREDENTIAL_GUIDES.get(provider)
if default is None:
raise ValueError(f"Unsupported credential guide provider: {provider}")
record, store = await _get_guide_store(db)
store[provider] = {
"title": title or default.title,
"markdown": markdown,
"sources": sources or [],
"verification_status": verification_status,
"verification_error": verification_error,
}
if record is None:
db.add(SystemSetting(category=CREDENTIAL_GUIDES_CATEGORY, payload=store))
else:
record.payload = store
await db.commit()
return await get_credential_guide(db, provider)
async def reset_credential_guide(db, provider: str) -> dict[str, Any]:
default = DEFAULT_CREDENTIAL_GUIDES.get(provider)
if default is None:
raise ValueError(f"Unsupported credential guide provider: {provider}")
record, store = await _get_guide_store(db)
if provider in store:
store.pop(provider, None)
if record is not None:
record.payload = store
await db.commit()
return await get_credential_guide(db, provider)
async def generate_credential_guide(
db,
provider: str,
ai_client: AIProviderClient,
web_search_client: WebSearchClient | None = None,
) -> dict[str, Any]:
default = DEFAULT_CREDENTIAL_GUIDES.get(provider)
if default is None:
raise ValueError(f"Unsupported credential guide provider: {provider}")
search_evidence: list[dict[str, Any]] = []
search_error: str | None = None
if web_search_client is not None:
try:
evidence = await web_search_client.search(
_credential_guide_search_query(default),
max_results=5,
)
search_evidence = normalize_search_evidence(evidence, limit=5)
except WebSearchError as exc:
search_error = str(exc)
except Exception as exc:
search_error = f"WebSearch unavailable: {exc}"
if not search_evidence:
guide = await get_credential_guide(db, provider)
guide["verification_status"] = "unverified_no_search_evidence"
guide["verification_error"] = search_error
guide["sources"] = []
return guide
response = await ai_client.analyze(
SituationalAnalysisRequest(
title=f"Generate credential guide for {provider}",
objective=(
default.prompt
+ "\n只能根据 context.search_evidence 中的来源生成教程;"
+ "如果证据不足,明确说明需要以官方页面为准。"
),
context={
"provider": provider,
"current_default_guide": default.markdown,
"product_context": "Planet collector credential settings",
"search_evidence": search_evidence,
},
observations=[
"Use concise Chinese markdown.",
"Prefer stable concepts over brittle UI labels.",
"Include verification and troubleshooting steps.",
"Include a short sources section with the provided URLs.",
],
constraints=[
"Do not ask the user for secrets.",
"Do not include fabricated screenshots.",
"Do not invent source URLs or product UI labels.",
"Use only the provided search_evidence as factual support.",
"Return markdown only.",
],
)
)
markdown = response.content.strip()
if not markdown:
markdown = default.markdown
return await save_credential_guide(
db,
provider,
default.title,
markdown,
sources=search_evidence,
verification_status="verified_with_search_evidence",
)
def _credential_guide_search_query(default: CredentialGuideDefault) -> str:
if default.provider == "barentswatch":
return "BarentsWatch developer tutorial AIS API OAuth client credentials"
if default.provider == "aisstream":
return "AISStream API key documentation websocket stream"
return f"{default.provider} API credentials documentation"

View File

@@ -0,0 +1,391 @@
"""Runtime helpers for mapped custom data sources."""
from __future__ import annotations
import asyncio
import base64
import json
from datetime import UTC, datetime
from typing import Any
import httpx
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.target_schema_registry import TARGET_SCHEMAS
from app.db.session import async_session_factory
from app.models.datasource_config import DataSourceConfig
from app.models.datasource_mapping import DataSourceMappingTemplate
from app.services.datasource_mapping import (
MappingError,
execute_mapping,
extract_path,
persist_mapped_records,
)
DEFAULT_MAPPING_TEMPLATES: dict[str, dict[str, Any]] = {
"vessel_ais": {
"source": {"items_path": "$"},
"fields": {
"mmsi": {"path": "$.mmsi", "type": "integer"},
"name": {"path": "$.name", "type": "string", "default": None},
"lat": {"path": "$.lat", "type": "float"},
"lon": {"path": "$.lon", "type": "float"},
"sog": {"path": "$.sog", "type": "float", "default": None},
"cog": {"path": "$.cog", "type": "float", "default": None},
"heading": {"path": "$.heading", "type": "integer", "default": None},
"nav_status": {"path": "$.nav_status", "type": "integer", "default": None},
"callsign": {"path": "$.callsign", "type": "string", "default": None},
"vessel_type": {"path": "$.vessel_type", "type": "string", "default": None},
"vessel_type_name": {"path": "$.vessel_type_name", "type": "string", "default": None},
"received_at": {"path": "$.received_at", "type": "datetime", "default": None},
},
"meta": {"generated_by": "default_template", "requires_review": False},
},
}
RUNNING_CUSTOM_STREAM_TASKS: dict[int, asyncio.Task[Any]] = {}
class CustomDatasourceRuntimeError(RuntimeError):
"""Raised when a custom datasource cannot run."""
def build_request_headers(auth_type: str, auth_config: dict, headers: dict) -> dict[str, str]:
request_headers = {str(key): str(value) for key, value in (headers or {}).items()}
auth_type = str(auth_type or "none").lower()
auth_config = auth_config or {}
if auth_type == "bearer" and auth_config.get("token"):
request_headers["Authorization"] = f"Bearer {auth_config['token']}"
elif auth_type == "api_key" and auth_config.get("api_key"):
location = str(auth_config.get("in") or auth_config.get("location") or "header").lower()
if location != "query":
key_name = auth_config.get("key_name", "X-API-Key")
request_headers[str(key_name)] = str(auth_config["api_key"])
elif auth_type == "basic":
username = auth_config.get("username", "")
password = auth_config.get("password", "")
credentials = f"{username}:{password}"
encoded = base64.b64encode(credentials.encode()).decode()
request_headers["Authorization"] = f"Basic {encoded}"
return request_headers
def build_query_params(auth_type: str, auth_config: dict, config: dict) -> dict[str, Any]:
params: dict[str, Any] = {}
candidate = (config or {}).get("params") or (config or {}).get("query_params")
if isinstance(candidate, dict):
params.update(candidate)
auth_type = str(auth_type or "none").lower()
auth_config = auth_config or {}
if auth_type == "api_key" and auth_config.get("api_key"):
location = str(auth_config.get("in") or auth_config.get("location") or "header").lower()
if location == "query":
key_name = auth_config.get("key_name") or auth_config.get("param_name") or "api_key"
params[str(key_name)] = auth_config["api_key"]
return params
async def load_active_mapping(
db: AsyncSession,
datasource_config_id: int,
) -> DataSourceMappingTemplate:
result = await db.execute(
select(DataSourceMappingTemplate)
.where(DataSourceMappingTemplate.datasource_config_id == datasource_config_id)
.where(DataSourceMappingTemplate.is_active.is_(True))
.order_by(DataSourceMappingTemplate.version.desc())
.limit(1)
)
mapping = result.scalar_one_or_none()
if mapping is not None:
return mapping
datasource = await db.get(DataSourceConfig, datasource_config_id)
if datasource is None:
raise CustomDatasourceRuntimeError("Configuration not found")
target_schema = (datasource.config or {}).get("target_schema")
template_body = DEFAULT_MAPPING_TEMPLATES.get(str(target_schema or "")) if target_schema else None
if not template_body or target_schema not in TARGET_SCHEMAS:
raise CustomDatasourceRuntimeError(
"No active mapping template found and no default template available for this target schema"
)
mapping = DataSourceMappingTemplate(
datasource_config_id=datasource_config_id,
target_schema=str(target_schema),
mapping_json=template_body,
sample_payload_hash=None,
validation_status="valid",
version=1,
is_active=True,
)
db.add(mapping)
await db.commit()
await db.refresh(mapping)
return mapping
async def fetch_rest_payload(config: DataSourceConfig, limit_bytes: int) -> Any:
request_config = config.config or {}
method = str(request_config.get("method") or request_config.get("request_method") or "GET").upper()
if method not in {"GET", "POST"}:
raise CustomDatasourceRuntimeError("Only GET and POST sample requests are supported.")
headers = build_request_headers(config.auth_type, config.auth_config or {}, config.headers or {})
params = build_query_params(config.auth_type, config.auth_config or {}, request_config)
timeout = float(request_config.get("timeout", 30))
json_body = request_config.get("json_body")
if json_body is None and str(request_config.get("body_type") or "").lower() in {"json", ""}:
candidate = request_config.get("body")
if isinstance(candidate, (dict, list)):
json_body = candidate
async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client:
response = await client.request(
method,
config.endpoint,
headers=headers,
params=params or None,
json=json_body,
)
response.raise_for_status()
content = response.content[:limit_bytes]
if "application/json" in response.headers.get("content-type", ""):
return json.loads(content.decode(response.encoding or "utf-8"))
return {"text": content.decode(response.encoding or "utf-8", errors="replace")}
async def run_mapped_rest_config(
db: AsyncSession,
datasource: DataSourceConfig,
) -> dict[str, Any]:
mapping = await load_active_mapping(db, datasource.id)
sample = await fetch_rest_payload(datasource, 5_000_000)
mapped = execute_mapping(sample, mapping.mapping_json, mapping.target_schema)
if mapped["failed_count"] > 0:
return {
"status": "failed",
"datasource_config_id": datasource.id,
"mapping_id": mapping.id,
"mapping_version": mapping.version,
"target_schema": mapping.target_schema,
"mapped_count": mapped["mapped_count"],
"failed_count": mapped["failed_count"],
"errors": mapped["errors"][:20],
}
request_config = datasource.config or {}
written_count = await persist_mapped_records(
db,
datasource_name=datasource.name,
datasource_config_id=datasource.id,
target_schema=mapping.target_schema,
records=mapped["records"],
mapping_version=mapping.version,
delivery_mode=request_config.get("delivery_mode") or "polling",
transport="http",
)
return {
"status": "success",
"datasource_config_id": datasource.id,
"mapping_id": mapping.id,
"mapping_version": mapping.version,
"target_schema": mapping.target_schema,
"fetched_count": mapped["total_items"],
"mapped_count": mapped["mapped_count"],
"written_count": written_count,
}
def _items_from_ws_message(payload: Any, config: dict) -> Any:
message_path = config.get("ws_message_path")
items_path = config.get("ws_items_path")
value = extract_path(payload, message_path) if message_path else payload
return extract_path(value, items_path) if items_path else value
async def _connect_websocket(endpoint: str, headers: dict[str, str]):
import websockets
try:
return await websockets.connect(endpoint, additional_headers=headers or None)
except TypeError:
return await websockets.connect(endpoint, extra_headers=headers or None)
async def test_websocket_config(config: DataSourceConfig) -> dict[str, Any]:
if not str(config.endpoint or "").startswith(("ws://", "wss://")):
raise CustomDatasourceRuntimeError("WebSocket datasource endpoint must start with ws:// or wss://")
runtime_config = config.config or {}
headers = build_request_headers(config.auth_type, config.auth_config or {}, config.headers or {})
receive_timeout = float(runtime_config.get("receive_timeout_seconds") or runtime_config.get("timeout") or 10)
async with await _connect_websocket(config.endpoint, headers) as websocket:
subscribe_message = runtime_config.get("ws_subscribe_message")
if isinstance(subscribe_message, (dict, list)):
await websocket.send(json.dumps(subscribe_message))
elif isinstance(subscribe_message, str) and subscribe_message.strip():
await websocket.send(subscribe_message)
raw_message = await asyncio.wait_for(websocket.recv(), timeout=receive_timeout)
return {
"success": True,
"message_preview": raw_message[:1000] if isinstance(raw_message, str) else str(raw_message)[:1000],
}
async def run_mapped_websocket_config(
db: AsyncSession,
datasource: DataSourceConfig,
*,
debug_max_messages: int | None = None,
use_config_debug_max_messages: bool = True,
) -> dict[str, Any]:
if not str(datasource.endpoint or "").startswith(("ws://", "wss://")):
raise CustomDatasourceRuntimeError("WebSocket datasource endpoint must start with ws:// or wss://")
mapping = await load_active_mapping(db, datasource.id)
runtime_config = datasource.config or {}
max_messages = debug_max_messages
if max_messages is None and use_config_debug_max_messages:
max_messages = runtime_config.get("debug_max_messages")
max_messages = int(max_messages) if max_messages else None
receive_timeout = float(runtime_config.get("receive_timeout_seconds") or runtime_config.get("timeout") or 30)
reconnect = bool(runtime_config.get("ws_reconnect", True))
reconnect_delay = float(runtime_config.get("reconnect_delay_seconds") or 3)
headers = build_request_headers(datasource.auth_type, datasource.auth_config or {}, datasource.headers or {})
messages_seen = 0
mapped_count = 0
failed_count = 0
written_count = 0
errors: list[dict[str, Any]] = []
started_at = datetime.now(UTC)
while True:
try:
async with await _connect_websocket(datasource.endpoint, headers) as websocket:
subscribe_message = runtime_config.get("ws_subscribe_message")
if isinstance(subscribe_message, (dict, list)):
await websocket.send(json.dumps(subscribe_message))
elif isinstance(subscribe_message, str) and subscribe_message.strip():
await websocket.send(subscribe_message)
while True:
raw_message = await asyncio.wait_for(websocket.recv(), timeout=receive_timeout)
messages_seen += 1
try:
payload = json.loads(raw_message)
except json.JSONDecodeError as exc:
failed_count += 1
errors.append({"message": "invalid_json", "error": str(exc)})
continue
extracted = _items_from_ws_message(payload, runtime_config)
try:
mapped = execute_mapping(extracted, mapping.mapping_json, mapping.target_schema)
except (MappingError, ValueError) as exc:
failed_count += 1
errors.append({"message": "mapping_failed", "error": str(exc)})
continue
mapped_count += mapped["mapped_count"]
failed_count += mapped["failed_count"]
if mapped["errors"]:
errors.extend(mapped["errors"][:5])
if mapped["records"]:
written_count += await persist_mapped_records(
db,
datasource_name=datasource.name,
datasource_config_id=datasource.id,
target_schema=mapping.target_schema,
records=mapped["records"],
mapping_version=mapping.version,
delivery_mode=runtime_config.get("delivery_mode") or "realtime_stream",
transport="websocket",
)
if max_messages and messages_seen >= max_messages:
return {
"status": "success",
"datasource_config_id": datasource.id,
"mapping_id": mapping.id,
"mapping_version": mapping.version,
"target_schema": mapping.target_schema,
"messages_seen": messages_seen,
"mapped_count": mapped_count,
"failed_count": failed_count,
"written_count": written_count,
"errors": errors[:20],
"execution_time_seconds": (datetime.now(UTC) - started_at).total_seconds(),
}
except asyncio.CancelledError:
raise
except Exception as exc:
failed_count += 1
errors.append({"message": "websocket_error", "error": f"{exc.__class__.__name__}: {exc}"})
if not reconnect or max_messages:
return {
"status": "failed" if written_count == 0 else "partial",
"datasource_config_id": datasource.id,
"mapping_id": mapping.id,
"mapping_version": mapping.version,
"target_schema": mapping.target_schema,
"messages_seen": messages_seen,
"mapped_count": mapped_count,
"failed_count": failed_count,
"written_count": written_count,
"errors": errors[:20],
}
await asyncio.sleep(reconnect_delay)
async def run_custom_stream_by_id(config_id: int) -> dict[str, Any]:
async with async_session_factory() as db:
datasource = await db.get(DataSourceConfig, config_id)
if not datasource:
raise CustomDatasourceRuntimeError("Configuration not found")
return await run_mapped_websocket_config(
db,
datasource,
use_config_debug_max_messages=False,
)
def start_custom_stream(config_id: int) -> bool:
existing = RUNNING_CUSTOM_STREAM_TASKS.get(config_id)
if existing is not None and not existing.done():
return False
task = asyncio.create_task(run_custom_stream_by_id(config_id), name=f"custom-stream:{config_id}")
RUNNING_CUSTOM_STREAM_TASKS[config_id] = task
def _cleanup(done_task: asyncio.Task[Any]) -> None:
if RUNNING_CUSTOM_STREAM_TASKS.get(config_id) is done_task:
RUNNING_CUSTOM_STREAM_TASKS.pop(config_id, None)
task.add_done_callback(_cleanup)
return True
async def stop_custom_stream(config_id: int) -> bool:
task = RUNNING_CUSTOM_STREAM_TASKS.get(config_id)
if task is None or task.done():
RUNNING_CUSTOM_STREAM_TASKS.pop(config_id, None)
return False
task.cancel()
try:
await task
except asyncio.CancelledError:
return True
return task.cancelled()
def get_custom_stream_status(config_id: int) -> dict[str, Any]:
task = RUNNING_CUSTOM_STREAM_TASKS.get(config_id)
return {
"config_id": config_id,
"running": bool(task and not task.done()),
"done": bool(task and task.done()),
}

View File

@@ -0,0 +1,437 @@
"""Connectivity validation helpers for built-in datasource overrides."""
from __future__ import annotations
from datetime import UTC, datetime
import hashlib
import json
import os
from typing import Any
import httpx
from sqlalchemy import func, select
from app.core.data_sources import get_data_sources_config
from app.core.datasource_defaults import DEFAULT_DATASOURCES
from app.models.collected_data import CollectedData
from app.models.datasource import DataSource
from app.models.datasource_config import DataSourceConfig
from app.models.system_setting import SystemSetting
from app.services.barentswatch import (
_read_zshrc_env,
fetch_barentswatch_access_token,
resolve_barentswatch_config,
)
CONNECTIVITY_VALIDATION_KEY = "connectivity_validation"
CONNECTIVITY_STORE_CATEGORY = "datasource_connectivity_validations"
SUPPORTED_CREDENTIAL_PROVIDERS = {"barentswatch", "spacetrack", "aisstream"}
def _sha256_json(payload: Any) -> str:
encoded = json.dumps(payload, ensure_ascii=False, sort_keys=True, default=str).encode("utf-8")
return hashlib.sha256(encoded).hexdigest()
def _resolve_spacetrack_credentials() -> tuple[str, str, str]:
zshrc_env = _read_zshrc_env()
username = os.getenv("SPACETRACK_USERNAME") or zshrc_env.get("SPACETRACK_USERNAME") or ""
password = os.getenv("SPACETRACK_PASSWORD") or zshrc_env.get("SPACETRACK_PASSWORD") or ""
source = "environment" if os.getenv("SPACETRACK_USERNAME") or os.getenv("SPACETRACK_PASSWORD") else ""
if not source and (username or password):
source = "~/.zshrc"
return username, password, source or "missing"
async def _resolve_aisstream_api_key(
db=None,
credential_override: dict[str, str] | None = None,
) -> tuple[str, str]:
if credential_override and credential_override.get("api_key"):
return str(credential_override["api_key"]), "draft"
env_key = os.getenv("AISSTREAM_API_KEY")
zshrc_key = _read_zshrc_env().get("AISSTREAM_API_KEY")
if db is not None:
result = await db.execute(
select(DataSourceConfig)
.where(DataSourceConfig.name == "aisstream_vessels")
.where(DataSourceConfig.is_active.is_(True))
)
record = result.scalar_one_or_none()
if record:
auth_config = record.auth_config or {}
runtime_config = record.config or {}
api_key = auth_config.get("api_key") or runtime_config.get("api_key")
if api_key:
return str(api_key), "datasource_config"
if env_key:
return env_key, "environment"
if zshrc_key:
return zshrc_key, "~/.zshrc"
return "", "missing"
def strip_connectivity_validation(config: dict | None) -> dict:
cleaned = dict(config or {})
cleaned.pop(CONNECTIVITY_VALIDATION_KEY, None)
return cleaned
def merge_connectivity_validation(existing_config: dict | None, next_config: dict | None) -> dict:
merged = strip_connectivity_validation(next_config)
validation = (existing_config or {}).get(CONNECTIVITY_VALIDATION_KEY)
if validation:
merged[CONNECTIVITY_VALIDATION_KEY] = validation
return merged
def get_connectivity_validation(config: DataSourceConfig | None) -> dict | None:
validation = (config.config or {}).get(CONNECTIVITY_VALIDATION_KEY) if config else None
return validation if isinstance(validation, dict) else None
async def build_builtin_connectivity_checksum(
source: str,
endpoint: str,
auth_type: str,
headers: dict | None,
config: dict | None,
db=None,
credential_override: dict[str, str] | None = None,
) -> tuple[str, dict[str, Any]]:
defaults = DEFAULT_DATASOURCES.get(source, {})
credential_provider = defaults.get("credential_provider")
credential_fingerprint = ""
credential_source = "none"
has_credentials = not defaults.get("requires_credentials", False)
if credential_provider == "barentswatch":
if credential_override:
client_id = credential_override.get("client_id", "")
client_secret = credential_override.get("client_secret", "")
credential_source = "draft"
else:
barentswatch_config = await resolve_barentswatch_config(db)
client_id = barentswatch_config.client_id
client_secret = barentswatch_config.client_secret
credential_source = barentswatch_config.credential_source
has_credentials = bool(client_id and client_secret)
credential_fingerprint = _sha256_json(
{
"client_id": client_id,
"client_secret": client_secret,
}
)
elif credential_provider == "spacetrack":
username, password, credential_source = _resolve_spacetrack_credentials()
has_credentials = bool(username and password)
credential_fingerprint = _sha256_json(
{
"username": username,
"password": password,
}
)
elif credential_provider == "aisstream":
api_key, credential_source = await _resolve_aisstream_api_key(db, credential_override)
has_credentials = bool(api_key)
credential_fingerprint = _sha256_json({"api_key": api_key})
elif defaults.get("requires_credentials"):
credential_source = str(credential_provider or "unsupported")
checksum_payload = {
"source": source,
"endpoint": endpoint,
"auth_type": "none",
"headers": headers or {},
"config": strip_connectivity_validation(config),
"credential_provider": credential_provider or "none",
"credential_fingerprint": credential_fingerprint,
}
return _sha256_json(checksum_payload), {
"requires_credentials": bool(defaults.get("requires_credentials", False)),
"credential_provider": credential_provider,
"credential_source": credential_source,
"has_credentials": has_credentials,
}
async def test_builtin_connectivity(
source: str,
endpoint: str,
auth_type: str,
headers: dict | None,
config: dict | None,
db=None,
credential_override: dict[str, str] | None = None,
) -> dict[str, Any]:
defaults = DEFAULT_DATASOURCES.get(source)
if not defaults:
return {
"success": False,
"message": "未知内置采集器,无法执行连接校验。",
}
checksum, credential_context = await build_builtin_connectivity_checksum(
source,
endpoint,
auth_type,
headers,
config,
db,
credential_override,
)
if credential_context["requires_credentials"] and not credential_context["has_credentials"]:
return {
"success": False,
"checksum": checksum,
"stage": "credentials",
"message": "该采集器需要凭证,请先到采集器凭证设置中配置。",
"settings_tab": "collector_credentials",
**credential_context,
}
if (
credential_context["requires_credentials"]
and credential_context["credential_provider"] not in SUPPORTED_CREDENTIAL_PROVIDERS
):
return {
"success": False,
"checksum": checksum,
"stage": "credentials",
"message": "该采集器的凭证链路尚未接入,暂时无法完成连接校验。",
"settings_tab": "collector_credentials",
**credential_context,
}
request_headers = {str(key): str(value) for key, value in (headers or {}).items()}
request_config = strip_connectivity_validation(config)
timeout = float(request_config.get("timeout") or 30)
request_endpoint = endpoint
if credential_context["credential_provider"] == "aisstream":
if not str(request_endpoint).startswith(("ws://", "wss://")):
return {
"success": False,
"checksum": checksum,
"stage": "endpoint",
"message": "AISStream endpoint 必须是 ws:// 或 wss:// WebSocket 地址。",
**credential_context,
}
return {
"success": True,
"checksum": checksum,
"stage": "credentials",
"message": "AISStream 凭证已配置WebSocket endpoint 格式有效。",
**credential_context,
}
try:
async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client:
if credential_context["credential_provider"] == "barentswatch":
barentswatch_config = await resolve_barentswatch_config(db)
token = await fetch_barentswatch_access_token(client, barentswatch_config)
if not token:
return {
"success": False,
"checksum": checksum,
"stage": "token",
"message": "凭证可读取,但 token 响应中没有 access_token。",
"settings_tab": "collector_credentials",
**credential_context,
}
request_headers["Authorization"] = f"Bearer {token}"
elif credential_context["credential_provider"] == "spacetrack":
username, password, _source = _resolve_spacetrack_credentials()
login_url = "https://www.space-track.org/ajaxauth/login"
login_response = await client.post(
login_url,
data={
"identity": username,
"password": password,
},
)
login_response.raise_for_status()
started = datetime.now(UTC)
async with client.stream("GET", request_endpoint, headers=request_headers) as response:
response.raise_for_status()
status_code = response.status_code
elapsed_ms = (datetime.now(UTC) - started).total_seconds() * 1000
return {
"success": True,
"checksum": checksum,
"stage": "endpoint",
"message": "连接验证成功。",
"status_code": status_code,
"response_time_ms": elapsed_ms,
**credential_context,
}
except httpx.HTTPStatusError as exc:
return {
"success": False,
"checksum": checksum,
"stage": "endpoint",
"message": f"连接验证失败HTTP {exc.response.status_code}",
"error": f"HTTP Error: {exc.response.status_code}",
**credential_context,
}
except httpx.HTTPError as exc:
return {
"success": False,
"checksum": checksum,
"stage": "network",
"message": f"连接验证失败:{exc.__class__.__name__}",
"error": str(exc),
**credential_context,
}
def make_success_validation(checksum: str, result: dict[str, Any]) -> dict[str, Any]:
return {
"checksum": checksum,
"status": "success",
"validated_at": datetime.now(UTC).isoformat(),
"status_code": result.get("status_code"),
"credential_source": result.get("credential_source"),
}
def is_builtin_validation_current(config: DataSourceConfig | None, checksum: str) -> bool:
validation = get_connectivity_validation(config)
return bool(
validation
and validation.get("status") == "success"
and validation.get("checksum") == checksum
)
async def get_connectivity_store(db) -> dict[str, Any]:
result = await db.execute(
select(SystemSetting).where(SystemSetting.category == CONNECTIVITY_STORE_CATEGORY)
)
record = result.scalar_one_or_none()
return dict(record.payload or {}) if record and isinstance(record.payload, dict) else {}
async def save_connectivity_success(
db,
source: str,
checksum: str,
result: dict[str, Any],
*,
connected_by: str,
) -> dict[str, Any]:
store = await get_connectivity_store(db)
validation = {
**make_success_validation(checksum, result),
"connected_by": connected_by,
}
store[source] = validation
existing = await db.execute(
select(SystemSetting).where(SystemSetting.category == CONNECTIVITY_STORE_CATEGORY)
)
record = existing.scalar_one_or_none()
if record is None:
db.add(SystemSetting(category=CONNECTIVITY_STORE_CATEGORY, payload=store))
else:
record.payload = store
return validation
async def load_builtin_override_config(db, source: str) -> DataSourceConfig | None:
result = await db.execute(
select(DataSourceConfig)
.where(DataSourceConfig.name == source)
.where(DataSourceConfig.is_active.is_(True))
)
return result.scalar_one_or_none()
async def get_builtin_effective_candidate(db, source: str) -> dict[str, Any]:
override = await load_builtin_override_config(db, source)
default_endpoint = get_data_sources_config().get_yaml_url(source)
return {
"name": source,
"endpoint": (override.endpoint if override and override.endpoint else default_endpoint) or "",
"auth_type": override.auth_type if override else "none",
"headers": override.headers if override else {},
"config": strip_connectivity_validation(override.config if override else {}),
}
async def has_collected_data(db, source: str) -> bool:
result = await db.execute(select(func.count(CollectedData.id)).where(CollectedData.source == source))
if (result.scalar() or 0) > 0:
return True
datasource_result = await db.execute(select(DataSource).where(DataSource.source == source))
datasource = datasource_result.scalar_one_or_none()
return bool(datasource and datasource.last_status == "success")
async def get_builtin_connection_status(
db,
source: str,
endpoint: str,
auth_type: str,
headers: dict | None,
config: dict | None,
) -> dict[str, Any]:
checksum, credential_context = await build_builtin_connectivity_checksum(
source,
endpoint,
auth_type,
headers,
config,
db,
)
store = await get_connectivity_store(db)
validation = store.get(source)
if isinstance(validation, dict) and validation.get("status") == "success":
if validation.get("checksum") == checksum:
return {
"connected": True,
"checksum": checksum,
"connected_by": validation.get("connected_by") or "connection_button",
"message": "当前配置已完成连接验证。",
**credential_context,
}
effective = await get_builtin_effective_candidate(db, source)
effective_checksum, _ = await build_builtin_connectivity_checksum(
source,
effective["endpoint"],
effective["auth_type"],
effective["headers"],
effective["config"],
db,
)
if checksum == effective_checksum and await has_collected_data(db, source):
return {
"connected": True,
"checksum": checksum,
"connected_by": "collection",
"message": "当前配置已有成功采集数据,视为已连接。",
**credential_context,
}
if isinstance(validation, dict) and validation.get("status") == "success":
return {
"connected": False,
"checksum": checksum,
"connected_by": None,
"message": "接口地址或凭证指纹已变化,请重新点击连接验证。",
**credential_context,
}
return {
"connected": False,
"checksum": checksum,
"connected_by": None,
"message": "当前配置尚未连接,请点击连接验证。",
**credential_context,
}

View File

@@ -0,0 +1,410 @@
"""Deterministic mapping support for custom data sources."""
from __future__ import annotations
import hashlib
import json
import re
from datetime import UTC, datetime
from typing import Any
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.target_schema_registry import TargetSchema, get_target_schema
SECRET_KEY_PATTERN = re.compile(
r"(token|secret|password|passwd|authorization|api[_-]?key|client[_-]?secret)",
re.IGNORECASE,
)
class MappingError(ValueError):
"""Raised when a mapping definition cannot be executed."""
def stable_payload_hash(payload: Any) -> str:
encoded = json.dumps(payload, ensure_ascii=False, sort_keys=True, default=str).encode()
return hashlib.sha256(encoded).hexdigest()
def redact_for_llm(value: Any) -> Any:
if isinstance(value, dict):
redacted = {}
for key, item in value.items():
if SECRET_KEY_PATTERN.search(str(key)):
redacted[key] = "[REDACTED]"
else:
redacted[key] = redact_for_llm(item)
return redacted
if isinstance(value, list):
return [redact_for_llm(item) for item in value[:20]]
return value
def extract_path(payload: Any, path: str | None) -> Any:
if not path or path == "$":
return payload
normalized = path.strip()
if normalized.startswith("$."):
normalized = normalized[2:]
elif normalized.startswith("$"):
normalized = normalized[1:]
normalized = normalized.strip(".")
if not normalized:
return payload
current = payload
for raw_segment in normalized.split("."):
segment = raw_segment.strip()
if not segment:
continue
list_all = segment.endswith("[*]")
if list_all:
segment = segment[:-3]
index = None
match = re.fullmatch(r"(.+)\[(\d+)\]", segment)
if match:
segment = match.group(1)
index = int(match.group(2))
if segment:
if isinstance(current, dict):
current = current.get(segment)
else:
return None
if list_all:
return current if isinstance(current, list) else []
if index is not None:
if not isinstance(current, list) or index >= len(current):
return None
current = current[index]
return current
def _convert_value(value: Any, target_type: str | None) -> Any:
if value is None or target_type in (None, "", "any"):
return value
if target_type == "string":
return str(value)
if target_type == "integer":
return int(value)
if target_type == "float":
return float(value)
if target_type == "boolean":
if isinstance(value, bool):
return value
if isinstance(value, str):
return value.strip().lower() in {"1", "true", "yes", "y", "on"}
return bool(value)
if target_type == "datetime":
if isinstance(value, datetime):
return value
if isinstance(value, (int, float)):
return datetime.fromtimestamp(value)
if isinstance(value, str):
return datetime.fromisoformat(value.replace("Z", "+00:00"))
return value
if target_type == "object":
if isinstance(value, dict):
return value
raise ValueError("expected object")
if target_type == "array":
if isinstance(value, list):
return value
raise ValueError("expected array")
return value
def _apply_enum(value: Any, enum_map: Any) -> Any:
if not isinstance(enum_map, dict):
return value
key = str(value)
return enum_map.get(key, enum_map.get(value, value))
def _map_one(item: Any, field_mapping: dict[str, Any]) -> tuple[dict[str, Any], list[str]]:
output: dict[str, Any] = {}
errors: list[str] = []
for field_name, rule in field_mapping.items():
if isinstance(rule, str):
rule = {"path": rule}
if not isinstance(rule, dict):
errors.append(f"{field_name}: mapping rule must be an object or path string")
continue
value = extract_path(item, rule.get("path"))
if value is None and "default" in rule:
value = rule.get("default")
value = _apply_enum(value, rule.get("enum"))
try:
value = _convert_value(value, rule.get("type"))
except (TypeError, ValueError) as exc:
errors.append(f"{field_name}: failed to convert value {value!r}: {exc}")
continue
if value is not None or rule.get("include_null", False):
output[field_name] = value
return output, errors
def execute_mapping(
payload: Any,
mapping_json: dict[str, Any],
target_schema: str | TargetSchema,
*,
limit: int | None = None,
) -> dict[str, Any]:
schema = get_target_schema(target_schema) if isinstance(target_schema, str) else target_schema
source = mapping_json.get("source") or {}
fields = mapping_json.get("fields")
if not isinstance(fields, dict) or not fields:
raise MappingError("mapping_json.fields must be a non-empty object")
items_path = source.get("items_path") or mapping_json.get("items_path") or "$"
items = extract_path(payload, items_path)
if isinstance(items, dict):
items = [items]
elif not isinstance(items, list):
items = []
if limit is not None:
items = items[:limit]
mapped_records: list[dict[str, Any]] = []
errors: list[dict[str, Any]] = []
for index, item in enumerate(items):
mapped, mapping_errors = _map_one(item, fields)
validated, validation_errors = schema.validate_record(mapped)
all_errors = mapping_errors + validation_errors
if all_errors:
errors.append({"index": index, "errors": all_errors, "record": mapped})
continue
if validated is not None:
mapped_records.append(validated)
return {
"target_schema": schema.key,
"total_items": len(items),
"mapped_count": len(mapped_records),
"failed_count": len(errors),
"records": mapped_records,
"errors": errors,
}
def build_heuristic_mapping(sample_payload: Any, target_schema_key: str) -> dict[str, Any]:
schema = get_target_schema(target_schema_key)
items_path = "$"
sample_item = sample_payload
if isinstance(sample_payload, dict):
for key in ("data", "items", "results", "features", "vessels"):
candidate = sample_payload.get(key)
if isinstance(candidate, list) and candidate:
items_path = f"$.{key}[*]"
sample_item = candidate[0]
break
elif isinstance(sample_payload, list) and sample_payload:
items_path = "$"
sample_item = sample_payload[0]
available = _flatten_keys(sample_item if isinstance(sample_item, dict) else {})
fields: dict[str, Any] = {}
for field in schema.fields:
candidate = _best_field_match(field.name, available)
if candidate:
fields[field.name] = {"path": f"$.{candidate}", "type": field.type}
elif field.name == "data" and target_schema_key == "generic_records":
fields[field.name] = {"path": "$", "type": "object"}
elif not field.required:
fields[field.name] = {"path": f"$.{field.name}", "type": field.type, "default": None}
return {
"source": {"items_path": items_path},
"fields": fields,
"meta": {
"generated_by": "heuristic",
"requires_review": True,
},
}
def _flatten_keys(payload: dict[str, Any], prefix: str = "") -> list[str]:
keys: list[str] = []
for key, value in payload.items():
dotted = f"{prefix}.{key}" if prefix else str(key)
keys.append(dotted)
if isinstance(value, dict):
keys.extend(_flatten_keys(value, dotted))
return keys
def _best_field_match(field_name: str, candidates: list[str]) -> str | None:
aliases = {
"lat": ("lat", "latitude", "y"),
"lon": ("lon", "lng", "longitude", "x"),
"mmsi": ("mmsi",),
"sog": ("sog", "speed", "speedOverGround"),
"cog": ("cog", "course", "courseOverGround"),
"received_at": ("received_at", "timestamp", "time", "updated_at"),
"observed_at": ("observed_at", "timestamp", "time", "updated_at"),
"source_id": ("id", "source_id", "uuid"),
}.get(field_name, (field_name,))
lowered = {candidate.lower(): candidate for candidate in candidates}
for alias in aliases:
if alias.lower() in lowered:
return lowered[alias.lower()]
for candidate in candidates:
tail = candidate.split(".")[-1].lower()
if tail in {alias.lower() for alias in aliases}:
return candidate
return None
def _parse_datetime(value: Any) -> datetime | None:
if value is None:
return None
if isinstance(value, datetime):
return value
if isinstance(value, str):
return datetime.fromisoformat(value.replace("Z", "+00:00"))
return None
async def persist_mapped_records(
db: AsyncSession,
*,
datasource_name: str,
datasource_config_id: int,
target_schema: str,
records: list[dict[str, Any]],
mapping_version: int,
delivery_mode: str | None = None,
transport: str | None = None,
) -> int:
"""Persist validated mapped records to the destination for a target schema."""
if target_schema == "vessel_ais":
from app.core.time import to_iso8601_utc
from app.core.websocket.broadcaster import broadcaster
from app.services.vessel_ais_aggregation import (
record_vessel_ais_observation,
update_ais_source_health,
)
now = datetime.now(UTC)
latest_observed_at = now
written_count = 0
for record in records:
observed_at = _parse_datetime(record.get("received_at")) or now
observation = await record_vessel_ais_observation(
db,
source=datasource_name,
normalized_payload=record,
raw_payload=record,
delivery_mode=delivery_mode or "polling",
transport=transport or "http",
message_type="PositionReport",
observed_at=observed_at,
collected_at=now,
)
if observation is not None:
written_count += 1
if observed_at > latest_observed_at:
latest_observed_at = observed_at
await update_ais_source_health(
db,
source=datasource_name,
connection_state="connected",
observed_count=len(records),
last_seen_at=latest_observed_at,
last_success_at=now if records else None,
lag_seconds=max((now - latest_observed_at).total_seconds(), 0),
)
await db.commit()
if records:
await broadcaster.broadcast_custom(
"vessels",
{
"action": "upsert",
"source": datasource_name,
"created": True,
"vessels": [
{
"mmsi": record.get("mmsi"),
"mmsi_display": str(record.get("mmsi")) if record.get("mmsi") is not None else None,
"name": record.get("name"),
"callsign": record.get("callsign"),
"lat": record.get("lat"),
"lon": record.get("lon"),
"sog": record.get("sog"),
"cog": record.get("cog"),
"heading": record.get("heading"),
"nav_status": record.get("nav_status"),
"vessel_type": record.get("vessel_type"),
"vessel_type_name": record.get("vessel_type_name"),
"received_at": to_iso8601_utc(_parse_datetime(record.get("received_at"))),
}
for record in records
],
},
)
return written_count
from app.models.collected_data import CollectedData
collected_at = datetime.now(UTC)
for index, record in enumerate(records):
if target_schema == "geo_points":
source_id = record.get("source_id") or f"{datasource_config_id}:{index}"
name = record.get("name")
metadata = {
"latitude": record.get("lat"),
"longitude": record.get("lon"),
"type": record.get("type"),
"mapping_version": mapping_version,
"target_schema": target_schema,
**(record.get("metadata") or {}),
}
reference_date = _parse_datetime(record.get("observed_at"))
else:
source_id = record.get("source_id") or f"{datasource_config_id}:{index}"
name = None
metadata = {
"data": record.get("data") or {},
"mapping_version": mapping_version,
"target_schema": target_schema,
}
reference_date = _parse_datetime(record.get("observed_at"))
db.add(
CollectedData(
source=datasource_name,
source_id=str(source_id),
entity_key=f"{datasource_name}:{source_id}",
data_type=target_schema,
name=name,
title=name,
extra_data=metadata,
collected_at=collected_at,
reference_date=reference_date,
is_valid=1,
is_current=True,
change_type="created",
change_summary={},
)
)
await db.commit()
return len(records)

View File

@@ -0,0 +1,120 @@
"""Server-side Docs metadata and Gatekeeper authorization helpers."""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import Literal
from app.models.user import User
DocsAccess = Literal["public", "docs_user", "docs_developer", "docs_admin"]
DocsLang = Literal["zh", "en"]
VALID_DOCS_LANGS = {"zh", "en"}
DOCS_README_FILENAME = "README.md"
DEFAULT_DOCS_SLUG = "overview"
REPO_ROOT = Path(__file__).resolve().parents[3]
TECHNICAL_DOCS_ROOT = REPO_ROOT / "docs" / "technical"
@dataclass(frozen=True)
class DocsMetadata:
filename: str
slug: str
access: DocsAccess
group: str
order: int
zh_title: str
en_title: str
DOCS_METADATA: tuple[DocsMetadata, ...] = (
DocsMetadata(DOCS_README_FILENAME, DEFAULT_DOCS_SLUG, "public", "Overview", 0, "技术文档", "Technical Docs"),
DocsMetadata("quickstart.md", "quickstart", "public", "Manual", 1, "快速开始", "Quickstart"),
DocsMetadata("manual.md", "manual", "public", "Manual", 2, "Planet 使用手册", "Planet Manual"),
DocsMetadata("faq.md", "faq", "public", "Manual", 3, "常见问题", "FAQ"),
DocsMetadata("location-pipeline-user.md", "location-pipeline-user", "public", "Manual", 4, "Earth 位置候选采集使用手册", "Earth Location Candidate Collection User Guide"),
DocsMetadata("earth-frontend-context.md", "earth-frontend-context", "docs_developer", "Earth", 10, "Earth 前端结构", "Earth Frontend Context"),
DocsMetadata("earth-layer-style-reference.md", "earth-layer-style-reference", "docs_developer", "Earth", 11, "Earth 图层样式属性索引", "Earth Layer Style Reference"),
DocsMetadata("earth-render-layer-order.md", "earth-render-layer-order", "docs_developer", "Earth", 12, "Earth 渲染图层顺序", "Earth Render Layer Order"),
DocsMetadata("earth-satellite-footprint-policy.md", "earth-satellite-footprint-policy", "docs_developer", "Earth", 13, "Earth 卫星覆盖策略", "Earth Satellite Footprint Policy"),
DocsMetadata("earth-bgp-context.md", "earth-bgp-context", "docs_developer", "Earth", 14, "BGP 态势上下文", "BGP Context"),
DocsMetadata("earth-news-live-streams-collector-format.md", "earth-news-live-streams-collector-format", "docs_developer", "Earth", 15, "新闻直播采集格式", "News Live Streams Collector Format"),
DocsMetadata("earth-interactable-usage.md", "earth-interactable-usage", "docs_developer", "Earth", 16, "Earth 可交互图标接入", "Earth Interactable Usage"),
DocsMetadata("earth-toolbar-overlay-coordination.md", "earth-toolbar-overlay-coordination", "docs_developer", "Earth", 17, "Earth 工具栏与浮层协同", "Earth Toolbar and Overlay Coordination"),
DocsMetadata("frontend-admin-frontend-context.md", "frontend-admin-frontend-context", "docs_developer", "Frontend", 20, "控制台前端结构", "Admin Frontend Context"),
DocsMetadata("frontend-layout-guidelines.md", "frontend-layout-guidelines", "docs_developer", "Frontend", 21, "前端布局指南", "Frontend Layout Guidelines"),
DocsMetadata("docs-gatekeeper-development.md", "docs-gatekeeper-development", "docs_developer", "Frontend", 22, "Docs Gatekeeper 开发说明", "Docs Gatekeeper Development Guide"),
DocsMetadata("backend-collectors.md", "backend-collectors", "docs_developer", "Backend", 30, "数据采集系统", "Data Collectors"),
DocsMetadata("backend-system-service-control.md", "backend-system-service-control", "docs_admin", "Backend", 31, "系统服务控制", "System Service Control"),
DocsMetadata("datasource-collector-settings-connectivity.md", "datasource-collector-settings-connectivity", "docs_developer", "Backend", 32, "数据源、采集器设置与连接验证", "Datasource Collector Settings and Connectivity"),
DocsMetadata("backend-datasources-api-performance.md", "backend-datasources-api-performance", "docs_developer", "Backend", 33, "数据源 API 性能", "Datasource API Performance"),
DocsMetadata("location-pipeline-development.md", "location-pipeline-development", "docs_developer", "Backend", 34, "通用位置估算管线开发说明", "Shared Location Resolution Pipeline Development Guide"),
DocsMetadata("agents-aiprovider.md", "agents-aiprovider", "docs_developer", "Agents", 40, "AI Provider 指南", "AI Provider Guide"),
DocsMetadata("ops-docker-compose-buildx-upgrade.md", "ops-docker-compose-buildx-upgrade", "docs_admin", "Ops", 50, "Docker + Compose + Buildx 升级", "Docker + Compose + Buildx Upgrade"),
DocsMetadata("ops-planet-sh-startup.md", "ops-planet-sh-startup", "docs_admin", "Ops", 51, "planet.sh 启动机制", "planet.sh Startup"),
)
DOCS_BY_SLUG = {entry.slug: entry for entry in DOCS_METADATA}
def get_user_gatekeeper_groups(user: User | None) -> set[str]:
if user is None:
return set()
role = user.role.value if hasattr(user.role, "value") else str(user.role or "")
if role == "super_admin":
return {"docs_user", "docs_developer", "docs_admin"}
if role == "admin":
return {"docs_user", "docs_developer", "docs_admin"}
groups = set()
raw_groups = user.gatekeeper_groups or []
if isinstance(raw_groups, list):
groups.update(str(group) for group in raw_groups)
if "docs_admin" in groups:
groups.update({"docs_developer", "docs_user"})
if "docs_developer" in groups:
groups.add("docs_user")
return groups
def can_read_doc(entry: DocsMetadata, user: User | None) -> bool:
if entry.access == "public":
return True
return entry.access in get_user_gatekeeper_groups(user)
def doc_path_for(entry: DocsMetadata, lang: str) -> Path:
if lang not in VALID_DOCS_LANGS:
raise ValueError("Unsupported docs language")
return TECHNICAL_DOCS_ROOT / lang / entry.filename
def title_for(entry: DocsMetadata, lang: str) -> str:
return entry.zh_title if lang == "zh" else entry.en_title
def catalog_for_user(user: User | None) -> list[dict]:
items: list[dict] = []
for entry in DOCS_METADATA:
if not can_read_doc(entry, user):
continue
for lang in sorted(VALID_DOCS_LANGS):
if not doc_path_for(entry, lang).exists():
continue
items.append(
{
"slug": entry.slug,
"filename": entry.filename,
"lang": lang,
"title": title_for(entry, lang),
"group": entry.group,
"order": entry.order,
"access": entry.access,
}
)
return sorted(items, key=lambda item: (item["lang"], item["order"], item["title"]))

View File

@@ -0,0 +1,150 @@
"""LLM provider presets used by Settings and the runtime AI provider bridge."""
from __future__ import annotations
from typing import Any
import httpx
MODELS_DEV_URL = "https://models.dev/api.json"
FALLBACK_LLM_PROVIDER_PRESETS: dict[str, dict[str, Any]] = {
"minimax": {
"provider": "minimax",
"label": "MiniMax",
"provider_api": "anthropic-messages",
"base_url": "https://api.minimaxi.com/anthropic",
"model": "MiniMax-M2.7",
"models": ["MiniMax-M2.7", "MiniMax-M2.7-highspeed", "MiniMax-M2.5", "MiniMax-M2"],
"api_key_env": "MINIMAX_API_KEY",
"source": "fallback",
},
"openai": {
"provider": "openai",
"label": "OpenAI",
"provider_api": "openai-completions",
"base_url": "https://api.openai.com/v1",
"model": "gpt-5.1",
"models": ["gpt-5.1", "gpt-5.1-codex", "gpt-4.1", "gpt-4o"],
"api_key_env": "OPENAI_API_KEY",
"source": "fallback",
},
"anthropic": {
"provider": "anthropic",
"label": "Anthropic",
"provider_api": "anthropic-messages",
"base_url": "https://api.anthropic.com/v1",
"model": "claude-sonnet-4-6",
"models": ["claude-sonnet-4-6", "claude-opus-4-5", "claude-3-5-haiku-20241022"],
"api_key_env": "ANTHROPIC_API_KEY",
"source": "fallback",
},
"deepseek": {
"provider": "deepseek",
"label": "DeepSeek",
"provider_api": "openai-completions",
"base_url": "https://api.deepseek.com/v1",
"model": "deepseek-chat",
"models": ["deepseek-chat", "deepseek-reasoner"],
"api_key_env": "DEEPSEEK_API_KEY",
"source": "fallback",
},
"alibaba": {
"provider": "alibaba",
"label": "Alibaba Qwen / DashScope",
"provider_api": "openai-completions",
"base_url": "https://dashscope.aliyuncs.com/compatible-mode/v1",
"model": "qwen3-max",
"models": ["qwen3-max", "qwen3.5-plus", "qwen-max", "qwen-plus"],
"api_key_env": "DASHSCOPE_API_KEY",
"source": "fallback",
},
"moonshotai": {
"provider": "moonshotai",
"label": "Moonshot AI / Kimi",
"provider_api": "openai-completions",
"base_url": "https://api.moonshot.ai/v1",
"model": "kimi-k2.5",
"models": ["kimi-k2.5", "kimi-k2-thinking", "kimi-k2-turbo-preview"],
"api_key_env": "MOONSHOT_API_KEY",
"source": "fallback",
},
"openrouter": {
"provider": "openrouter",
"label": "OpenRouter",
"provider_api": "openai-completions",
"base_url": "https://openrouter.ai/api/v1",
"model": "openai/gpt-5.1",
"models": ["openai/gpt-5.1", "anthropic/claude-sonnet-4.5", "qwen/qwen3-max"],
"api_key_env": "OPENROUTER_API_KEY",
"source": "fallback",
},
"ollama": {
"provider": "ollama",
"label": "Ollama Local",
"provider_api": "ollama-generate",
"base_url": "http://127.0.0.1:11434",
"model": "qwen2.5:7b",
"models": ["qwen2.5:7b", "llama3.1:8b", "mistral:7b"],
"api_key_env": "",
"source": "fallback",
},
}
MODELS_DEV_PROVIDER_KEYS = {
"minimax": "minimax",
"openai": "openai",
"anthropic": "anthropic",
"deepseek": "deepseek",
"alibaba": "alibaba",
"moonshotai": "moonshotai",
"openrouter": "openrouter",
}
def list_fallback_llm_provider_presets() -> list[dict[str, Any]]:
return [dict(value) for value in FALLBACK_LLM_PROVIDER_PRESETS.values()]
def get_fallback_llm_provider_preset(provider: str) -> dict[str, Any]:
key = provider.strip().lower()
if key not in FALLBACK_LLM_PROVIDER_PRESETS:
raise ValueError(f"Unsupported LLM provider preset: {provider}")
return dict(FALLBACK_LLM_PROVIDER_PRESETS[key])
async def refresh_llm_provider_preset(provider: str) -> dict[str, Any]:
fallback = get_fallback_llm_provider_preset(provider)
models_dev_key = MODELS_DEV_PROVIDER_KEYS.get(fallback["provider"])
if not models_dev_key:
return fallback
async with httpx.AsyncClient(timeout=15.0, follow_redirects=True) as client:
response = await client.get(
MODELS_DEV_URL,
headers={"User-Agent": "Planet/1.0"},
)
response.raise_for_status()
catalog = response.json()
upstream = catalog.get(models_dev_key)
if not isinstance(upstream, dict):
return fallback
upstream_models = upstream.get("models") if isinstance(upstream.get("models"), dict) else {}
model_ids = list(upstream_models.keys())[:80]
base_url = upstream.get("api") or fallback["base_url"]
if fallback["provider"] == "deepseek" and base_url == "https://api.deepseek.com":
base_url = "https://api.deepseek.com/v1"
refreshed = {
**fallback,
"label": upstream.get("name") or fallback["label"],
"base_url": base_url,
"model": model_ids[0] if model_ids else fallback["model"],
"models": model_ids or fallback["models"],
"api_key_env": (upstream.get("env") or [fallback["api_key_env"]])[0],
"source": MODELS_DEV_URL,
}
return refreshed

View File

@@ -0,0 +1,57 @@
"""Shared location-resolution pipeline.
A reusable abstraction for "given a record, decide its lat/lon" — used by
compute centers, BGP collectors, BGP events, and any future entity that needs
location estimation.
Each domain wires its own :class:`LocationPipeline` from a sequence of
:class:`LocationResolver` instances. Future algorithms (peeringdb, IXP tables,
user-confirmed coordinates, …) plug in by implementing the protocol — no
changes needed to consumers.
"""
from .models import (
LocationCandidate,
LocationQuery,
ResolutionDiagnostic,
ResolutionResult,
ResolverOutput,
)
from .pipeline import LocationPipeline, LocationResolver
from .resolvers.inherit import InheritFromAnotherEntityResolver
from .resolvers.nominatim import (
NominatimResolver,
build_default_nominatim_geocoder,
interpret_geocode_result,
)
from .resolvers.registry import RegistryResolver, default_score_alias_match
from .resolvers.source_coordinates import SourceCoordinatesResolver
from .text import (
city_key,
coerce_str,
normalize_country_text,
normalize_text,
parse_float,
)
__all__ = [
"LocationCandidate",
"LocationPipeline",
"LocationQuery",
"LocationResolver",
"ResolutionDiagnostic",
"ResolutionResult",
"ResolverOutput",
"InheritFromAnotherEntityResolver",
"NominatimResolver",
"RegistryResolver",
"SourceCoordinatesResolver",
"build_default_nominatim_geocoder",
"city_key",
"coerce_str",
"default_score_alias_match",
"interpret_geocode_result",
"normalize_country_text",
"normalize_text",
"parse_float",
]

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,128 @@
"""Domain-neutral data structures for the location pipeline."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Mapping
# Renderable precision tiers, ordered from most precise to least.
RENDERABLE_PRECISIONS: tuple[str, ...] = ("precise", "site", "city")
@dataclass(frozen=True)
class LocationQuery:
"""Domain-neutral input for the resolution pipeline.
``name`` and ``aliases`` are matched against registry alias indexes;
``city`` / ``country`` / ``region`` provide geographic context for both
registry lookups and Nominatim queries; ``source_latitude`` /
``source_longitude`` short-circuit when the record already carries
coordinates; ``extra`` carries domain-specific fields (operator, site,
organization, asn, peer_ip, …) that resolvers can opt into.
"""
name: str | None = None
aliases: tuple[str, ...] = ()
city: str | None = None
country: str | None = None
region: str | None = None
source_latitude: float | None = None
source_longitude: float | None = None
extra: Mapping[str, Any] = field(default_factory=dict)
@dataclass(frozen=True)
class LocationCandidate:
"""A resolved location candidate produced by a resolver."""
latitude: float
longitude: float
display_name: str
precision: str # "precise" | "site" | "city" | (rejected: country/unknown)
confidence: float
query: str
source: str
source_note: str | None
matched_fields: tuple[str, ...]
needs_confirmation: bool
city: str | None = None
region: str | None = None
country: str | None = None
matched_location_name: str | None = None
location_verified_at: str | None = None
suggested_registry_entry: dict[str, Any] | None = None
raw_payload: dict[str, Any] | None = None
def to_dict(self) -> dict[str, Any]:
return {
"latitude": self.latitude,
"longitude": self.longitude,
"display_name": self.display_name,
"precision": self.precision,
"confidence": self.confidence,
"query": self.query,
"source": self.source,
"source_note": self.source_note,
"matched_fields": list(self.matched_fields),
"needs_confirmation": self.needs_confirmation,
"city": self.city,
"region": self.region,
"country": self.country,
"matched_location_name": self.matched_location_name,
"location_verified_at": self.location_verified_at,
"suggested_registry_entry": self.suggested_registry_entry,
"raw_payload": self.raw_payload,
}
@dataclass(frozen=True)
class ResolverOutput:
"""What a single resolver returns from one ``resolve()`` call."""
candidates: tuple[LocationCandidate, ...] = ()
attempted_queries: tuple[str, ...] = ()
@dataclass(frozen=True)
class ResolutionDiagnostic:
"""Why we could not resolve, plus what we tried."""
failure_reason: str
attempted_queries: tuple[str, ...] = ()
record_id: int | None = None
source: str | None = None
source_id: str | None = None
name: str | None = None
country: str | None = None
city: str | None = None
site: str | None = None
operator: str | None = None
extra: Mapping[str, Any] = field(default_factory=dict)
def to_dict(self) -> dict[str, Any]:
return {
"failure_reason": self.failure_reason,
"attempted_queries": list(self.attempted_queries),
"record_id": self.record_id,
"source": self.source,
"source_id": self.source_id,
"name": self.name,
"country": self.country,
"city": self.city,
"site": self.site,
"operator": self.operator,
**({"extra": dict(self.extra)} if self.extra else {}),
}
@dataclass(frozen=True)
class ResolutionResult:
"""Pipeline output: best candidate (if any) + diagnostic on miss."""
location: LocationCandidate | None
diagnostic: ResolutionDiagnostic | None
attempted_queries: tuple[str, ...] = ()
@property
def is_resolved(self) -> bool:
return bool(self.location)

View File

@@ -0,0 +1,126 @@
"""Pipeline that runs a sequence of :class:`LocationResolver` instances."""
from __future__ import annotations
from typing import Protocol, Sequence
from .models import (
LocationCandidate,
LocationQuery,
ResolutionDiagnostic,
ResolutionResult,
ResolverOutput,
)
class LocationResolver(Protocol):
"""Pluggable location resolution step.
Implementations: ``SourceCoordinatesResolver``, ``RegistryResolver``,
``NominatimResolver``, ``InheritFromAnotherEntityResolver`` — see the
``resolvers`` subpackage. New algorithms (peeringdb / IXP / user-confirmed
coordinates) plug in by implementing this protocol; the pipeline does not
care how candidates are produced.
"""
name: str
def resolve(self, query: LocationQuery) -> ResolverOutput: ...
def default_candidate_sort_key(
candidate: LocationCandidate,
) -> tuple[int, int, float]:
precision_rank = {"precise": 0, "site": 1, "city": 2}.get(
candidate.precision, 9
)
source_rank = {
"source_coordinates": 0,
"stored_compute_center_location": 1,
"stored_collector_location": 1,
"ror_organization_registry": 2,
"inherited": 3,
"nominatim_online_geocode": 4,
"local_registry": 8,
"local_registry_city": 9,
}.get(candidate.source, 9)
return (source_rank, precision_rank, -float(candidate.confidence or 0))
class LocationPipeline:
"""Orchestrate a sequence of resolvers.
``collect_candidates`` runs every resolver and returns *all* deduped
candidates plus the queries each resolver attempted (useful for
user-facing "why didn't this work?" diagnostics).
``resolve_best`` returns the top candidate per
:func:`default_candidate_sort_key` (or a custom sort).
"""
def __init__(
self,
resolvers: Sequence[LocationResolver],
*,
sort_key=default_candidate_sort_key,
failure_reason: str = (
"Could not resolve to renderable coordinates from any configured resolver."
),
) -> None:
self._resolvers = list(resolvers)
self._sort_key = sort_key
self._failure_reason = failure_reason
@property
def resolvers(self) -> tuple[LocationResolver, ...]:
return tuple(self._resolvers)
def collect_candidates(
self, query: LocationQuery
) -> tuple[list[LocationCandidate], list[str]]:
candidates: list[LocationCandidate] = []
attempted: list[str] = []
seen_keys: set[tuple[str, str, str]] = set()
for resolver in self._resolvers:
output = resolver.resolve(query)
for q in output.attempted_queries:
if q and q not in attempted:
attempted.append(q)
for candidate in output.candidates:
key = (
candidate.source,
f"{candidate.latitude:.4f}",
f"{candidate.longitude:.4f}",
)
if key in seen_keys:
continue
seen_keys.add(key)
candidates.append(candidate)
candidates.sort(key=self._sort_key)
return candidates, attempted
def resolve_best(self, query: LocationQuery) -> ResolutionResult:
candidates, attempted = self.collect_candidates(query)
if candidates:
return ResolutionResult(
location=candidates[0],
diagnostic=None,
attempted_queries=tuple(attempted),
)
return ResolutionResult(
location=None,
diagnostic=ResolutionDiagnostic(
failure_reason=self._failure_reason,
attempted_queries=tuple(attempted),
name=query.name,
country=query.country,
city=query.city,
site=str(query.extra.get("site")) if query.extra.get("site") else None,
operator=str(query.extra.get("operator"))
if query.extra.get("operator")
else None,
),
attempted_queries=tuple(attempted),
)

View File

@@ -0,0 +1,20 @@
"""Built-in resolver implementations."""
from .inherit import InheritFromAnotherEntityResolver
from .nominatim import (
NominatimResolver,
build_default_nominatim_geocoder,
interpret_geocode_result,
)
from .registry import RegistryResolver, default_score_alias_match
from .source_coordinates import SourceCoordinatesResolver
__all__ = [
"InheritFromAnotherEntityResolver",
"NominatimResolver",
"RegistryResolver",
"SourceCoordinatesResolver",
"build_default_nominatim_geocoder",
"default_score_alias_match",
"interpret_geocode_result",
]

View File

@@ -0,0 +1,31 @@
"""Resolver that inherits a candidate from another entity's resolution.
Used by BGP events to pick up the location of their owning collector. The
``source_lookup`` callable is the only domain coupling — it receives the
incoming :class:`LocationQuery` and returns either an already-resolved
:class:`LocationCandidate` (typically by querying another pipeline) or
``None`` to signal "no parent location available".
"""
from __future__ import annotations
from typing import Callable
from ..models import LocationCandidate, LocationQuery, ResolverOutput
class InheritFromAnotherEntityResolver:
def __init__(
self,
*,
source_lookup: Callable[[LocationQuery], LocationCandidate | None],
name: str = "inherited",
) -> None:
self.name = name
self._lookup = source_lookup
def resolve(self, query: LocationQuery) -> ResolverOutput:
result = self._lookup(query)
if result is None:
return ResolverOutput()
return ResolverOutput(candidates=(result,))

View File

@@ -0,0 +1,292 @@
"""Nominatim-backed online geocoder.
The actual HTTP call is encapsulated in :func:`build_default_nominatim_geocoder`
which returns an ``lru_cache``-wrapped function. Domain modules typically:
1. Build a default geocoder via :func:`build_default_nominatim_geocoder`.
2. Re-export it under a stable module-level name (e.g. ``_geocode_online``).
3. Pass a *late-binding lambda* (``lambda q: _geocode_online(q)``) to
:class:`NominatimResolver`.
This ensures tests that ``monkeypatch.setattr(module, "_geocode_online", ...)``
can swap the geocoder behavior without touching pipeline construction.
"""
from __future__ import annotations
import time
from functools import lru_cache
from typing import Any, Callable
import httpx
from ..models import LocationCandidate, LocationQuery, ResolverOutput
from ..text import (
coerce_str,
normalize_country_text,
normalize_text,
parse_float,
)
NOMINATIM_SEARCH_URL = "https://nominatim.openstreetmap.org/search"
DEFAULT_USER_AGENT = "planet-earth-location-resolver/1.0"
DEFAULT_MIN_INTERVAL_SECONDS = 1.1
DEFAULT_TIMEOUT_SECONDS = 8.0
def build_default_nominatim_geocoder(
*,
user_agent: str = DEFAULT_USER_AGENT,
min_interval_seconds: float = DEFAULT_MIN_INTERVAL_SECONDS,
timeout_seconds: float = DEFAULT_TIMEOUT_SECONDS,
cache_size: int = 512,
) -> Callable[[str], dict[str, Any] | None]:
"""Return a cached, rate-limited Nominatim geocoder."""
last_request_at = [0.0]
@lru_cache(maxsize=cache_size)
def geocode(query: str) -> dict[str, Any] | None:
if not query:
return None
elapsed = time.monotonic() - last_request_at[0]
if elapsed < min_interval_seconds:
time.sleep(min_interval_seconds - elapsed)
last_request_at[0] = time.monotonic()
response = httpx.get(
NOMINATIM_SEARCH_URL,
params={
"q": query,
"format": "jsonv2",
"limit": 1,
"addressdetails": 1,
},
headers={"User-Agent": user_agent},
timeout=timeout_seconds,
)
response.raise_for_status()
payload = response.json()
if not isinstance(payload, list) or not payload:
return None
result = payload[0]
if not isinstance(result, dict):
return None
return result
return geocode
_DEFAULT_SITE_CATEGORIES = frozenset(
{
"amenity",
"office",
"building",
"industrial",
"research",
"university",
"education",
"tourism",
"shop",
"man_made",
"campus",
"research_institute",
}
)
def interpret_geocode_result(
result: dict[str, Any],
*,
matched_fields: tuple[str, ...],
context_country: str | None,
site_categories: frozenset[str] = _DEFAULT_SITE_CATEGORIES,
site_promoting_match_fields: frozenset[str] = frozenset(
{"site", "operator", "name"}
),
) -> tuple[float, float, dict[str, Any], str] | None:
"""Validate a Nominatim raw result. Returns (lat, lon, address, classification)."""
latitude = parse_float(result.get("lat"))
longitude = parse_float(result.get("lon"))
if latitude in (None, 0.0) or longitude in (None, 0.0):
return None
address = result.get("address") if isinstance(result.get("address"), dict) else {}
if not isinstance(address, dict):
address = {}
has_city_level = bool(
address.get("city")
or address.get("town")
or address.get("village")
or address.get("municipality")
or address.get("hamlet")
or address.get("suburb")
)
osm_class = str(result.get("class") or "").lower()
osm_type = str(result.get("type") or "").lower()
is_site_like = osm_class in site_categories or osm_type in site_categories
if not has_city_level and not is_site_like:
return None
if context_country:
normalized_context = normalize_text(normalize_country_text(context_country))
normalized_result = normalize_text(
normalize_country_text(address.get("country"))
)
if (
normalized_context
and normalized_result
and normalized_context != normalized_result
):
return None
classification = (
"site"
if (
is_site_like
and has_city_level
and any(field in site_promoting_match_fields for field in matched_fields)
)
else "city"
)
return float(latitude), float(longitude), address, classification
def _candidate_from_geocode(
*,
query: LocationQuery,
geocode_query: str,
matched_fields: tuple[str, ...],
raw_result: dict[str, Any],
interpret: Callable[..., tuple[float, float, dict[str, Any], str] | None],
source: str,
site_confidence: float,
city_confidence: float,
) -> LocationCandidate | None:
interpreted = interpret(
raw_result,
matched_fields=matched_fields,
context_country=query.country,
)
if not interpreted:
return None
latitude, longitude, address, classification = interpreted
city = (
address.get("city")
or address.get("town")
or address.get("village")
or address.get("municipality")
or query.city
or None
)
region = address.get("state") or address.get("region")
country = address.get("country") or query.country or None
display_name = raw_result.get("display_name") or geocode_query
confidence = city_confidence if classification == "city" else site_confidence
extra = query.extra or {}
suggested_registry_entry = {
"canonical_name": (
(query.aliases[0] if query.aliases else None)
or query.name
or display_name
),
"aliases": list(
{
value
for value in [
query.name,
*query.aliases,
coerce_str(extra.get("operator")),
coerce_str(extra.get("site")),
]
if value
}
),
"operator": coerce_str(extra.get("operator")) or None,
"site": coerce_str(extra.get("site"))
or coerce_str(extra.get("organization"))
or None,
"country": country,
"city": city,
"region": region,
"latitude": latitude,
"longitude": longitude,
"precision": classification,
"confidence": confidence,
"source_note": (
f"Resolved via Nominatim query '{geocode_query}'{display_name}"
),
}
return LocationCandidate(
latitude=latitude,
longitude=longitude,
display_name=display_name,
precision=classification,
confidence=confidence,
query=geocode_query,
source=source,
source_note=f"Nominatim search result: {display_name}",
matched_fields=matched_fields,
needs_confirmation=True,
city=city,
region=region,
country=country,
matched_location_name=display_name,
location_verified_at=None,
suggested_registry_entry=suggested_registry_entry,
)
class NominatimResolver:
"""Run a domain-specific query plan against Nominatim."""
def __init__(
self,
*,
query_plan_builder: Callable[
[LocationQuery], list[tuple[str, tuple[str, ...]]]
],
geocoder: Callable[[str], dict[str, Any] | None],
name: str = "nominatim_online_geocode",
site_confidence: float = 0.72,
city_confidence: float = 0.62,
interpret: Callable[..., tuple[float, float, dict[str, Any], str] | None] = (
interpret_geocode_result
),
) -> None:
self.name = name
self._query_plan_builder = query_plan_builder
self._geocoder = geocoder
self._site_confidence = site_confidence
self._city_confidence = city_confidence
self._interpret = interpret
def resolve(self, query: LocationQuery) -> ResolverOutput:
plan = self._query_plan_builder(query)
candidates: list[LocationCandidate] = []
attempted: list[str] = []
for geocode_query, matched_fields in plan:
attempted.append(geocode_query)
try:
raw_result = self._geocoder(geocode_query)
except Exception:
continue
if not raw_result:
continue
candidate = _candidate_from_geocode(
query=query,
geocode_query=geocode_query,
matched_fields=matched_fields,
raw_result=raw_result,
interpret=self._interpret,
source=self.name,
site_confidence=self._site_confidence,
city_confidence=self._city_confidence,
)
if candidate is not None:
candidates.append(candidate)
return ResolverOutput(
candidates=tuple(candidates),
attempted_queries=tuple(attempted),
)

View File

@@ -0,0 +1,323 @@
"""Resolver that matches a query against a local JSON registry.
Registry schema (a single JSON file):
{
"locations": [
{
"canonical_name": "...",
"aliases": ["...", "..."],
"operator": "...",
"site": "...",
"city": "...",
"country": "...",
"region": "...",
"latitude": 0.0,
"longitude": 0.0,
"precision": "precise" | "site" | "city",
"confidence": 0.0,
"verification_status": "verified",
"source_note": "...",
"verified_at": "YYYY-MM-DD"
}
],
"city_fallbacks": [ {city, country, latitude, longitude, ...} ]
}
"""
from __future__ import annotations
import json
from functools import lru_cache
from pathlib import Path
from typing import Any, Callable, Iterable
from ..models import (
RENDERABLE_PRECISIONS,
LocationCandidate,
LocationQuery,
ResolverOutput,
)
from ..text import (
city_key,
normalize_country_text,
normalize_text,
parse_float,
)
# Field-priority weights when scoring "this query field text contains this
# alias text". Tuned to match the legacy compute-center ordering — name beats
# site beats operator beats city — which generalizes well to other domains.
_DEFAULT_FIELD_PRIORITY = {
"name": 8,
"site": 6,
"operator": 5,
"city": 3,
}
def default_score_alias_match(
alias_field: str, record_field: str, alias_text: str
) -> int:
score = max(0, len(alias_text))
score += _DEFAULT_FIELD_PRIORITY.get(alias_field, 1)
if alias_field == record_field:
score += 4
if alias_field == "name" and record_field in {"name", "name_short", "alias"}:
score += 6
if alias_field == "site" and record_field in {"site", "organization"}:
score += 4
if alias_field == "operator" and record_field in {"operator", "organization"}:
score += 4
return score
@lru_cache(maxsize=32)
def _load_registry_file(path: str) -> dict[str, Any]:
with Path(path).open("r", encoding="utf-8") as handle:
return json.load(handle)
@lru_cache(maxsize=32)
def _build_alias_index(
path: str,
) -> tuple[tuple[dict[str, Any], tuple[tuple[str, str], ...]], ...]:
index: list[tuple[dict[str, Any], tuple[tuple[str, str], ...]]] = []
for entry in _load_registry_file(path).get("locations", []):
aliases: list[tuple[str, str]] = []
seen: set[str] = set()
for alias in [entry.get("canonical_name"), *(entry.get("aliases") or [])]:
normalized = normalize_text(alias)
if normalized and normalized not in seen:
aliases.append(("name", normalized))
seen.add(normalized)
for field_name in ("operator", "site", "city"):
value = entry.get(field_name)
normalized = normalize_text(value)
if normalized and normalized not in seen:
aliases.append((field_name, normalized))
seen.add(normalized)
index.append((entry, tuple(aliases)))
return tuple(index)
def _query_corpus(query: LocationQuery) -> dict[str, str]:
"""Map a query into normalized strings keyed by source field."""
fields: dict[str, str] = {
"name": query.name or "",
"city": query.city or "",
"country": query.country or "",
}
for alias in query.aliases:
if alias and alias != query.name:
fields["name_short"] = alias
break
extra = query.extra or {}
for key in ("site", "operator", "organization"):
value = extra.get(key)
if value:
fields[key] = str(value)
return {key: normalize_text(value) for key, value in fields.items() if value}
def _country_compatible(entry: dict[str, Any], query: LocationQuery) -> bool:
record_country = normalize_country_text(query.country)
entry_country = normalize_country_text(entry.get("country"))
if not record_country or not entry_country:
return True
return normalize_text(record_country) == normalize_text(entry_country)
def _normalized_alias_matches(alias_normalized: str, record_text: str) -> bool:
alias_tokens = alias_normalized.split()
record_tokens = record_text.split()
if not alias_tokens or not record_tokens:
return False
if len(alias_tokens) == 1:
return alias_tokens[0] in record_tokens
window_size = len(alias_tokens)
return any(
record_tokens[index : index + window_size] == alias_tokens
for index in range(0, len(record_tokens) - window_size + 1)
)
def _entry_to_candidate(
entry: dict[str, Any],
*,
matched_alias: str,
matched_fields: Iterable[str],
source: str,
score_explainer: str,
confidence_floor: float,
) -> LocationCandidate:
canonical_name = entry.get("canonical_name") or matched_alias
# Registry entries are treated as candidates unless explicitly verified.
# This prevents migrated hard-coded hints from appearing as factual
# location evidence.
is_verified = entry.get("verification_status") == "verified"
precision = entry.get("precision") or "city"
if precision not in RENDERABLE_PRECISIONS:
precision = "city"
fields_summary = ", ".join(sorted(set(matched_fields))) or "name"
confidence_value = parse_float(entry.get("confidence"))
confidence = (
float(confidence_value)
if confidence_value is not None
else confidence_floor
)
return LocationCandidate(
latitude=float(parse_float(entry.get("latitude")) or 0.0),
longitude=float(parse_float(entry.get("longitude")) or 0.0),
display_name=canonical_name,
precision=precision,
confidence=confidence,
query=f"local_registry::{matched_alias or canonical_name}",
source=source,
source_note=entry.get("source_note")
or f"{score_explainer}: matched {fields_summary}",
matched_fields=tuple(sorted(set(matched_fields))) or ("name",),
needs_confirmation=bool(entry.get("needs_confirmation")) or not is_verified,
city=entry.get("city"),
region=entry.get("region"),
country=entry.get("country"),
matched_location_name=canonical_name,
location_verified_at=entry.get("verified_at") if is_verified else None,
suggested_registry_entry=None,
)
class RegistryResolver:
"""Match a query against a JSON registry (plus its city_fallbacks table)."""
def __init__(
self,
*,
registry_path: Path | str,
name: str = "local_registry",
city_fallback_source: str = "local_registry_city",
city_fallback_confidence_default: float = 0.65,
confidence_default: float = 0.85,
score_alias_match: Callable[[str, str, str], int] = default_score_alias_match,
) -> None:
self.name = name
self._registry_path = str(Path(registry_path))
self._city_fallback_source = city_fallback_source
self._city_fallback_confidence_default = city_fallback_confidence_default
self._confidence_default = confidence_default
self._score = score_alias_match
def reload(self) -> None:
"""Drop the cached registry — useful when the JSON file is edited."""
_load_registry_file.cache_clear()
_build_alias_index.cache_clear()
def resolve(self, query: LocationQuery) -> ResolverOutput:
candidates: list[LocationCandidate] = []
candidates.extend(self._registry_candidates(query))
city_candidate = self._city_fallback_candidate(query)
if city_candidate is not None:
candidates.append(city_candidate)
return ResolverOutput(candidates=tuple(candidates))
# ── internals ──────────────────────────────────────────────
def _registry_candidates(
self, query: LocationQuery
) -> list[LocationCandidate]:
corpus = _query_corpus(query)
if not corpus:
return []
# When the query carries a name (a record-specific identifier), require
# at least one alias match against a name-class field — otherwise a
# generic shared field like operator="RIPE NCC" would promote every
# registry entry that lists that operator, regardless of whether the
# name matches.
query_has_name = bool(corpus.get("name") or corpus.get("name_short"))
results: list[LocationCandidate] = []
for entry, aliases in _build_alias_index(self._registry_path):
best_alias = ""
best_score = 0
matched_fields: list[str] = []
matched_via_name_alias = False
for alias_field, alias_normalized in aliases:
for record_field, record_text in corpus.items():
if not _normalized_alias_matches(alias_normalized, record_text):
continue
score = self._score(
alias_field, record_field, alias_normalized
)
if score > best_score or (
score == best_score
and len(alias_normalized) > len(best_alias)
):
best_score = score
best_alias = alias_normalized
if record_field not in matched_fields:
matched_fields.append(record_field)
if alias_field == "name" and record_field in {"name", "name_short"}:
matched_via_name_alias = True
if not matched_fields or best_score <= 0:
continue
if query_has_name and not matched_via_name_alias:
continue
if not _country_compatible(entry, query):
continue
results.append(
_entry_to_candidate(
entry,
matched_alias=best_alias,
matched_fields=matched_fields,
source=self.name,
score_explainer="Registry alias match",
confidence_floor=self._confidence_default,
)
)
return results
def _city_fallback_candidate(
self, query: LocationQuery
) -> LocationCandidate | None:
country = normalize_country_text(query.country)
city = city_key(query.city)
if not country or not city:
return None
for fallback in _load_registry_file(self._registry_path).get(
"city_fallbacks", []
):
fallback_country = normalize_country_text(fallback.get("country"))
fallback_city = city_key(fallback.get("city"))
if fallback_country != country or fallback_city != city:
continue
confidence_value = parse_float(fallback.get("confidence"))
confidence = (
float(confidence_value)
if confidence_value is not None
else self._city_fallback_confidence_default
)
return LocationCandidate(
latitude=float(parse_float(fallback.get("latitude")) or 0.0),
longitude=float(parse_float(fallback.get("longitude")) or 0.0),
display_name=fallback.get("city") or "",
precision="city",
confidence=confidence,
query=(
f"city_fallback::{fallback.get('city')}, "
f"{fallback.get('country')}"
),
source=self._city_fallback_source,
source_note=fallback.get("source_note")
or f"City fallback for {fallback.get('city')}, {fallback.get('country')}",
matched_fields=("city", "country"),
needs_confirmation=False,
city=fallback.get("city"),
region=fallback.get("region"),
country=fallback.get("country"),
matched_location_name=fallback.get("city"),
location_verified_at=fallback.get("verified_at"),
suggested_registry_entry=None,
)
return None

View File

@@ -0,0 +1,42 @@
"""Resolver that consumes lat/lon already present on the source record."""
from __future__ import annotations
from ..models import LocationCandidate, LocationQuery, ResolverOutput
from ..text import normalize_country_text
class SourceCoordinatesResolver:
"""Pass-through for records that already carry valid coordinates."""
name = "source_coordinates"
def __init__(self, *, source: str = "source_coordinates") -> None:
self._source = source
def resolve(self, query: LocationQuery) -> ResolverOutput:
lat = query.source_latitude
lon = query.source_longitude
if lat in (None, 0.0) or lon in (None, 0.0):
return ResolverOutput()
country = normalize_country_text(query.country) or query.country
candidate = LocationCandidate(
latitude=float(lat),
longitude=float(lon),
display_name=query.name or "",
precision="precise",
confidence=1.0,
query="source_coordinates",
source=self._source,
source_note="Source record provided valid coordinates.",
matched_fields=("source_coordinates",),
needs_confirmation=False,
city=query.city,
region=query.region,
country=country,
matched_location_name=query.name,
location_verified_at=None,
suggested_registry_entry=None,
)
return ResolverOutput(candidates=(candidate,))

View File

@@ -0,0 +1,41 @@
"""Text-normalization helpers shared by every resolver."""
from __future__ import annotations
import re
from typing import Any
from app.core.countries import normalize_country
def parse_float(value: Any) -> float | None:
try:
if value in (None, ""):
return None
return float(value)
except (TypeError, ValueError):
return None
def coerce_str(value: Any) -> str:
if value in (None, ""):
return ""
return str(value).strip()
def normalize_text(value: Any) -> str:
if value in (None, ""):
return ""
normalized = str(value).casefold()
normalized = re.sub(r"[^a-z0-9一-鿿]+", " ", normalized)
return re.sub(r"\s+", " ", normalized).strip()
def normalize_country_text(value: Any) -> str:
normalized = normalize_country(value)
return normalized or coerce_str(value)
def city_key(city: Any) -> str:
text = coerce_str(city).split(",", 1)[0]
return normalize_text(text)

View File

@@ -33,6 +33,7 @@ from app.services.playground_session_store import upsert_playground_session
STREAM_CHUNK_SIZE = 24
STREAM_INTERVAL_SECONDS = 0.08
THINKING_PREVIEW_SECONDS = 2.6
ORPHANED_RUN_MESSAGE = "后台生成任务已中断,请点击上一条用户消息的重试按钮重新生成。"
class _ActiveRun:
@@ -179,6 +180,7 @@ async def _build_thread_response(
session: PlaygroundSession,
) -> PlaygroundThreadResponse:
messages = await _list_visible_messages(db, session_id=session.id)
messages = await _reconcile_orphaned_active_messages(db, messages)
id_map = {item.id: item.public_id for item in messages}
return PlaygroundThreadResponse(
session=session_to_response(session),
@@ -186,6 +188,30 @@ async def _build_thread_response(
)
async def _reconcile_orphaned_active_messages(
db: AsyncSession,
messages: list[PlaygroundMessage],
) -> list[PlaygroundMessage]:
changed = False
for item in messages:
if item.status not in {"pending", "thinking", "answering"}:
continue
if item.public_id in _ACTIVE_RUNS:
continue
item.status = "error"
item.content = item.content or ORPHANED_RUN_MESSAGE
orphan_meta = "错误: 后台任务已中断"
if orphan_meta not in (item.meta or []):
item.meta = [*(item.meta or []), orphan_meta]
changed = True
if changed:
await db.flush()
await db.commit()
for item in messages:
await db.refresh(item)
return messages
async def get_thread(
db: AsyncSession,
*,
@@ -550,6 +576,15 @@ def _build_conversation_history(messages: Sequence[PlaygroundMessage], current_u
return history[-8:]
def _format_run_exception(exc: Exception) -> str:
if isinstance(exc, HTTPException):
detail = exc.detail
if isinstance(detail, str):
return detail
return str(detail)
return str(exc) or type(exc).__name__
async def _run_assistant_message(
*,
user_id: int,
@@ -680,13 +715,18 @@ async def _run_assistant_message(
await db.commit()
raise
except Exception as exc:
error_message = _format_run_exception(exc)
async with async_session_factory() as db:
result = await db.execute(select(PlaygroundMessage).where(PlaygroundMessage.id == assistant_message_id))
message = result.scalar_one_or_none()
if message is not None:
message.status = "error"
message.content = message.content or "分析失败,请检查 AI Provider 配置或稍后再试。"
message.meta = [*(message.meta or []), f"错误: {type(exc).__name__}"]
message.content = message.content or f"分析失败{error_message}"
message.meta = [
*(message.meta or []),
f"Request ID: {request_id}",
f"错误: {error_message}",
]
await db.flush()
await db.commit()
finally:

View File

@@ -14,6 +14,11 @@ from app.core.time import to_iso8601_utc
from app.models.datasource import DataSource
from app.models.task import CollectionTask
from app.services.collectors.registry import collector_registry
from app.services.datasource_connectivity import (
build_builtin_connectivity_checksum,
get_builtin_effective_candidate,
save_connectivity_success,
)
logger = get_logger(__name__)
@@ -179,6 +184,23 @@ async def run_collector_task(collector_name: str):
task_result = await collector.run(db)
datasource.last_run_at = datetime.now(UTC)
datasource.last_status = task_result.get("status")
if datasource.last_status == "success":
effective_candidate = await get_builtin_effective_candidate(db, datasource.source)
checksum, _credential_context = await build_builtin_connectivity_checksum(
datasource.source,
effective_candidate["endpoint"],
effective_candidate["auth_type"],
effective_candidate["headers"],
effective_candidate["config"],
db,
)
await save_connectivity_success(
db,
datasource.source,
checksum,
{"status_code": None},
connected_by="collection",
)
await _update_next_run_at(datasource, db)
logger.info_event(
"Collector completed",

View File

@@ -19,6 +19,10 @@ ALLOWED_ACTIONS: dict[str, dict[str, Any]] = {
"command": ["./planet.sh", "restart", "-b"],
"recovery_mode": "backend",
},
"restart-frontend": {
"command": ["./planet.sh", "restart", "-f"],
"recovery_mode": "frontend",
},
"restart-ai-provider": {
"command": ["./planet.sh", "restart", "-a"],
"recovery_mode": "ai-provider",

View File

@@ -0,0 +1,198 @@
"""Persistence + validation for the v4 vessel_ais aggregation strategy."""
from __future__ import annotations
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.system_setting import SystemSetting
VESSEL_AGGREGATION_STRATEGY_CATEGORY = "vessel_aggregation_strategy"
DYNAMIC_FIELDS: tuple[str, ...] = ("lat", "lon", "sog", "cog", "heading", "nav_status")
STATIC_FIELDS: tuple[str, ...] = (
"name",
"callsign",
"imo",
"flag",
"vessel_type",
"vessel_type_name",
"length",
"width",
"draught",
)
ALLOWED_FIELDS: frozenset[str] = frozenset(DYNAMIC_FIELDS + STATIC_FIELDS)
ALLOWED_DYNAMIC_MODES: frozenset[str] = frozenset({"newest"})
ALLOWED_STATIC_MODES: frozenset[str] = frozenset({"source_priority", "non_empty", "newest", "locked"})
ALLOWED_LOCKED_DYNAMIC_MODES: frozenset[str] = frozenset({"newest", "source_priority", "locked"})
DEFAULT_STRATEGY: dict[str, Any] = {
"version": 1,
"vessel_ais": {
"source_priority": ["aisstream_vessels", "barentswatch_vessels"],
"field_rules": {},
"freshness": {
"realtime_stream_seconds": 900,
"polling_seconds": 3600,
},
"allow_dynamic_lock": False,
},
}
class StrategyValidationError(ValueError):
"""Raised when a saved strategy payload is malformed."""
def _coerce_str_list(value: Any, *, label: str) -> list[str]:
if value is None:
return []
if not isinstance(value, list):
raise StrategyValidationError(f"{label} must be a list of source names")
out: list[str] = []
for item in value:
if not isinstance(item, str) or not item.strip():
raise StrategyValidationError(f"{label} entries must be non-empty strings")
out.append(item.strip())
return out
def validate_strategy(payload: dict[str, Any]) -> dict[str, Any]:
"""Validate and normalize a strategy payload. Raise StrategyValidationError on issues."""
if not isinstance(payload, dict):
raise StrategyValidationError("strategy payload must be an object")
vessel_ais = payload.get("vessel_ais")
if not isinstance(vessel_ais, dict):
raise StrategyValidationError("strategy.vessel_ais is required and must be an object")
allow_dynamic_lock = bool(vessel_ais.get("allow_dynamic_lock", False))
source_priority = _coerce_str_list(
vessel_ais.get("source_priority"),
label="vessel_ais.source_priority",
)
raw_rules = vessel_ais.get("field_rules") or {}
if not isinstance(raw_rules, dict):
raise StrategyValidationError("vessel_ais.field_rules must be an object")
field_rules: dict[str, dict[str, Any]] = {}
for field, rule in raw_rules.items():
if field not in ALLOWED_FIELDS:
raise StrategyValidationError(f"unknown vessel_ais field: {field}")
if not isinstance(rule, dict):
raise StrategyValidationError(f"field_rules.{field} must be an object")
mode = str(rule.get("mode") or "").strip()
if not mode:
raise StrategyValidationError(f"field_rules.{field}.mode is required")
is_dynamic = field in DYNAMIC_FIELDS
if is_dynamic:
allowed_modes = ALLOWED_LOCKED_DYNAMIC_MODES if allow_dynamic_lock else ALLOWED_DYNAMIC_MODES
if mode not in allowed_modes:
if not allow_dynamic_lock:
raise StrategyValidationError(
f"field_rules.{field}.mode='{mode}' requires allow_dynamic_lock=true"
)
raise StrategyValidationError(
f"field_rules.{field}.mode must be one of {sorted(allowed_modes)}"
)
else:
if mode not in ALLOWED_STATIC_MODES:
raise StrategyValidationError(
f"field_rules.{field}.mode must be one of {sorted(ALLOWED_STATIC_MODES)}"
)
normalized_rule: dict[str, Any] = {"mode": mode}
rule_priority = rule.get("source_priority")
if rule_priority is not None:
normalized_rule["source_priority"] = _coerce_str_list(
rule_priority,
label=f"field_rules.{field}.source_priority",
)
if mode == "locked":
locked_source = rule.get("locked_source")
if not isinstance(locked_source, str) or not locked_source.strip():
raise StrategyValidationError(
f"field_rules.{field}.locked_source must be a non-empty string when mode=locked"
)
normalized_rule["locked_source"] = locked_source.strip()
field_rules[field] = normalized_rule
raw_freshness = vessel_ais.get("freshness") or {}
if not isinstance(raw_freshness, dict):
raise StrategyValidationError("vessel_ais.freshness must be an object")
freshness: dict[str, int] = {}
for key in ("realtime_stream_seconds", "polling_seconds"):
value = raw_freshness.get(key, DEFAULT_STRATEGY["vessel_ais"]["freshness"][key])
try:
seconds = int(value)
except (TypeError, ValueError) as exc:
raise StrategyValidationError(f"freshness.{key} must be an integer") from exc
if seconds < 0:
raise StrategyValidationError(f"freshness.{key} must be non-negative")
freshness[key] = seconds
return {
"version": int(payload.get("version") or 0) + 1,
"vessel_ais": {
"source_priority": source_priority,
"field_rules": field_rules,
"freshness": freshness,
"allow_dynamic_lock": allow_dynamic_lock,
},
}
async def _select_setting(db: AsyncSession) -> SystemSetting | None:
result = await db.execute(
select(SystemSetting).where(SystemSetting.category == VESSEL_AGGREGATION_STRATEGY_CATEGORY)
)
return result.scalar_one_or_none()
def _current_version(setting: SystemSetting | None) -> int:
if setting is None:
return 0
payload = setting.payload or {}
return int(payload.get("version") or 0)
async def load_strategy(db: AsyncSession) -> dict[str, Any]:
setting = await _select_setting(db)
if setting is None or not isinstance(setting.payload, dict):
return DEFAULT_STRATEGY
payload = setting.payload
if "vessel_ais" not in payload:
return DEFAULT_STRATEGY
return payload
async def save_strategy(db: AsyncSession, payload: dict[str, Any]) -> dict[str, Any]:
"""Validate + persist; bumps version automatically."""
existing = await _select_setting(db)
incoming = dict(payload)
incoming.setdefault("version", _current_version(existing))
validated = validate_strategy(incoming)
if existing is None:
existing = SystemSetting(category=VESSEL_AGGREGATION_STRATEGY_CATEGORY, payload=validated)
db.add(existing)
else:
existing.payload = validated
await db.commit()
return validated
async def reset_strategy(db: AsyncSession) -> dict[str, Any]:
existing = await _select_setting(db)
payload = {**DEFAULT_STRATEGY, "version": _current_version(existing) + 1}
if existing is None:
existing = SystemSetting(category=VESSEL_AGGREGATION_STRATEGY_CATEGORY, payload=payload)
db.add(existing)
else:
existing.payload = payload
await db.commit()
return payload

View File

@@ -0,0 +1,698 @@
"""AIS raw observation and aggregation support for vessel collectors."""
from datetime import UTC, datetime, timedelta
from hashlib import sha256
import json
from typing import Any, Iterable
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.vessel import AISConflictRecord, AISRawObservation, AISSourceHealth
from app.services.vessel_aggregation_strategy import (
DEFAULT_STRATEGY,
load_strategy,
)
from app.services.vessel_types import normalize_vessel_type_name
VESSEL_AIS_SCHEMA = "vessel_ais"
DEFAULT_AGGREGATION_WINDOW_HOURS = 24
BARENTSWATCH_DELIVERY_MODE = "polling"
BARENTSWATCH_TRANSPORT = "http"
AISSTREAM_DELIVERY_MODE = "realtime_stream"
AISSTREAM_TRANSPORT = "websocket"
DELIVERY_MODE_PRIORITY = {
"realtime_stream": 40,
"batch_stream": 30,
"polling": 20,
"snapshot": 10,
}
DYNAMIC_FIELDS = ("lat", "lon", "sog", "cog", "heading", "nav_status")
CONFLICT_FIELDS = (
"name",
"callsign",
"imo",
"flag",
"vessel_type",
"vessel_type_name",
"length",
"width",
"draught",
)
def _json_default(value: Any) -> Any:
if isinstance(value, datetime):
return value.astimezone(UTC).isoformat()
return str(value)
def _stable_payload(value: Any) -> str:
return json.dumps(value, sort_keys=True, separators=(",", ":"), default=_json_default)
def _jsonable(value: Any) -> Any:
if isinstance(value, datetime):
return value.astimezone(UTC).isoformat()
if isinstance(value, dict):
return {str(key): _jsonable(item) for key, item in value.items()}
if isinstance(value, list):
return [_jsonable(item) for item in value]
return value
def _coerce_datetime(value: Any) -> datetime | None:
if isinstance(value, datetime):
return value if value.tzinfo else value.replace(tzinfo=UTC)
if isinstance(value, (int, float)):
timestamp = float(value)
if timestamp > 10_000_000_000:
timestamp /= 1000
return datetime.fromtimestamp(timestamp, UTC)
if isinstance(value, str) and value:
try:
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
return parsed if parsed.tzinfo else parsed.replace(tzinfo=UTC)
except ValueError:
return None
return None
def build_observation_hash(
*,
source: str,
entity_key: str,
message_type: str | None,
observed_at: datetime,
normalized_payload: dict[str, Any],
source_message_id: str | None = None,
) -> str:
"""Build a deterministic idempotency key for one source-level AIS observation."""
if source_message_id:
basis = {
"source": source,
"entity_key": entity_key,
"source_message_id": source_message_id,
}
else:
basis = {
"source": source,
"entity_key": entity_key,
"message_type": message_type,
"observed_at": observed_at.astimezone(UTC).isoformat(),
"payload": normalized_payload,
}
return sha256(_stable_payload(basis).encode("utf-8")).hexdigest()
def build_field_conflict_candidates(
observations: Iterable[AISRawObservation],
fields: Iterable[str] = CONFLICT_FIELDS,
) -> list[dict[str, Any]]:
"""Return current field disagreements from raw observations without mutating state."""
candidates_by_field: dict[str, dict[str, Any]] = {}
for observation in observations:
payload = observation.normalized_payload or {}
for field in fields:
value = payload.get(field)
if value in (None, ""):
continue
field_candidates = candidates_by_field.setdefault(field, {})
field_candidates[observation.source] = value
conflicts = []
for field, candidates in sorted(candidates_by_field.items()):
unique_values = {_stable_payload(value) for value in candidates.values()}
if len(unique_values) <= 1:
continue
conflicts.append(
{
"field": field,
"candidates": candidates,
"status": "candidate",
}
)
return conflicts
def _payload_value(payload: dict[str, Any], field: str) -> Any:
value = payload.get(field)
return None if value in (None, "") else value
def _clean_text(value: Any) -> str | None:
if value in (None, ""):
return None
text = str(value).strip()
return text or None
def _raw_metadata_value(observation: AISRawObservation, field: str) -> Any:
raw_payload = observation.raw_payload or {}
metadata = raw_payload.get("MetaData") if isinstance(raw_payload, dict) else None
if not isinstance(metadata, dict):
return None
if field == "name":
return _clean_text(metadata.get("ShipName") or metadata.get("ship_name") or metadata.get("name"))
return None
def _delivery_priority(observation: AISRawObservation) -> int:
return DELIVERY_MODE_PRIORITY.get(str(observation.delivery_mode or ""), 0)
def _has_valid_position(payload: dict[str, Any]) -> bool:
try:
lat = float(payload.get("lat"))
lon = float(payload.get("lon"))
except (TypeError, ValueError):
return False
return -90 <= lat <= 90 and -180 <= lon <= 180
def _is_future_observation(observation: AISRawObservation, now: datetime) -> bool:
return observation.observed_at > now
def _strategy_source_rank(
source: str,
strategy: dict[str, Any],
) -> int:
priority = (strategy.get("vessel_ais") or {}).get("source_priority") or []
if source in priority:
return len(priority) - priority.index(source)
return 0
def _is_stream_stale(
observation: AISRawObservation,
*,
now: datetime,
strategy: dict[str, Any],
) -> bool:
delivery_mode = str(observation.delivery_mode or "")
freshness = (strategy.get("vessel_ais") or {}).get("freshness") or {}
if delivery_mode == "realtime_stream":
window = int(freshness.get("realtime_stream_seconds", 0) or 0)
else:
window = int(freshness.get("polling_seconds", 0) or 0)
if window <= 0:
return False
return (now - observation.observed_at).total_seconds() > window
def _select_position_observation(
observations: list[AISRawObservation],
*,
now: datetime,
strategy: dict[str, Any] | None = None,
) -> tuple[AISRawObservation | None, list[str]]:
strategy = strategy or DEFAULT_STRATEGY
rejected_flags: list[str] = []
fresh_candidates: list[AISRawObservation] = []
stale_candidates: list[AISRawObservation] = []
for observation in observations:
payload = observation.normalized_payload or {}
if not _has_valid_position(payload):
rejected_flags.append("invalid_position")
continue
if _is_future_observation(observation, now):
rejected_flags.append("future_timestamp")
continue
if _is_stream_stale(observation, now=now, strategy=strategy):
stale_candidates.append(observation)
rejected_flags.append("freshness_fallback")
continue
fresh_candidates.append(observation)
candidates = fresh_candidates or stale_candidates
if not candidates:
return None, sorted(set(rejected_flags))
candidates.sort(
key=lambda item: (
item.observed_at,
_delivery_priority(item),
_strategy_source_rank(item.source, strategy),
item.collected_at,
item.id or 0,
),
reverse=True,
)
return candidates[0], sorted(set(rejected_flags))
def _select_static_field(
observations: list[AISRawObservation],
field: str,
strategy: dict[str, Any] | None = None,
) -> tuple[Any, str | None, str | None]:
strategy = strategy or DEFAULT_STRATEGY
candidates = []
for observation in observations:
value = _payload_value(observation.normalized_payload or {}, field)
if value is None:
value = _raw_metadata_value(observation, field)
if value is None:
continue
candidates.append((observation, value))
if not candidates:
return None, None, None
field_rules = (strategy.get("vessel_ais") or {}).get("field_rules") or {}
rule = field_rules.get(field) or {"mode": "source_priority"}
mode = rule.get("mode")
if mode == "locked":
locked_source = rule.get("locked_source")
for observation, value in candidates:
if observation.source == locked_source:
return value, observation.source, "locked"
if mode in ("source_priority", "locked"):
priority = rule.get("source_priority") or (strategy.get("vessel_ais") or {}).get("source_priority") or []
ranked = sorted(
candidates,
key=lambda item: (
priority.index(item[0].source) if item[0].source in priority else len(priority) + 1,
-_delivery_priority(item[0]),
-(item[0].observed_at.timestamp() if item[0].observed_at else 0),
),
)
observation, value = ranked[0]
return value, observation.source, "source_priority"
if mode == "newest":
ranked = sorted(
candidates,
key=lambda item: (item[0].observed_at, _delivery_priority(item[0]), item[0].id or 0),
reverse=True,
)
observation, value = ranked[0]
return value, observation.source, "newest_observation"
# default / non_empty: prefer delivery mode priority, then newest
candidates.sort(
key=lambda item: (
_delivery_priority(item[0]),
item[0].observed_at,
item[0].collected_at,
item[0].id or 0,
),
reverse=True,
)
selected_observation, selected_value = candidates[0]
unique_values = {_stable_payload(value) for _, value in candidates}
reason = "delivery_mode_priority" if len(unique_values) > 1 else "non_empty_priority"
return selected_value, selected_observation.source, reason
def _build_source_summary(observations: list[AISRawObservation]) -> dict[str, dict[str, Any]]:
summary: dict[str, dict[str, Any]] = {}
for observation in observations:
source_summary = summary.setdefault(
observation.source,
{
"observation_count": 0,
"latest_observed_at": None,
"delivery_mode": observation.delivery_mode,
"transport": observation.transport,
"message_types": [],
},
)
source_summary["observation_count"] += 1
latest_observed_at = source_summary["latest_observed_at"]
if latest_observed_at is None or observation.observed_at > latest_observed_at:
source_summary["latest_observed_at"] = observation.observed_at
if observation.message_type and observation.message_type not in source_summary["message_types"]:
source_summary["message_types"].append(observation.message_type)
return summary
def _build_aggregated_vessel(
entity_key: str,
observations: list[AISRawObservation],
*,
now: datetime,
strategy: dict[str, Any] | None = None,
) -> dict[str, Any] | None:
strategy = strategy or DEFAULT_STRATEGY
position_observation, rejected_flags = _select_position_observation(
observations, now=now, strategy=strategy
)
if position_observation is None:
return None
payload = position_observation.normalized_payload or {}
mmsi = int(entity_key)
result: dict[str, Any] = {
"mmsi": mmsi,
"lat": float(payload["lat"]),
"lon": float(payload["lon"]),
"received_at": position_observation.observed_at,
"field_sources": {},
"selected_reasons": {},
"source_summary": _build_source_summary(observations),
"quality_flags": sorted(
set((position_observation.quality_flags or []) + rejected_flags)
),
"aggregation_strategy_version": int(strategy.get("version") or 0),
}
for field in DYNAMIC_FIELDS:
value = _payload_value(payload, field)
if field in ("lat", "lon") or value is not None:
result[field] = value
result["field_sources"][field] = position_observation.source
result["selected_reasons"][field] = "newest_observation"
for field in CONFLICT_FIELDS:
selected_value, selected_source, reason = _select_static_field(
observations, field, strategy=strategy
)
if selected_value is None:
continue
result[field] = selected_value
result["field_sources"][field] = selected_source
result["selected_reasons"][field] = reason
result["name"] = result.get("name") or f"MMSI {mmsi}"
result["vessel_type_name"] = result.get("vessel_type_name") or normalize_vessel_type_name(
result.get("vessel_type")
)
return result
async def _upsert_conflict_records(
db: AsyncSession,
entity_key: str,
observations: list[AISRawObservation],
aggregated: dict[str, Any],
) -> int:
conflicts = build_field_conflict_candidates(observations)
now = datetime.now(UTC)
for conflict in conflicts:
field = conflict["field"]
result = await db.execute(
select(AISConflictRecord)
.where(AISConflictRecord.target_schema == VESSEL_AIS_SCHEMA)
.where(AISConflictRecord.entity_key == entity_key)
.where(AISConflictRecord.field == field)
.limit(1)
)
record = result.scalar_one_or_none()
if record is None:
record = AISConflictRecord(
target_schema=VESSEL_AIS_SCHEMA,
entity_key=entity_key,
field=field,
)
db.add(record)
record.candidates = conflict["candidates"]
record.selected_source = (aggregated.get("field_sources") or {}).get(field)
record.selected_value = aggregated.get(field)
record.selected_reason = (aggregated.get("selected_reasons") or {}).get(field)
record.resolved_by = "system"
record.status = "open"
record.updated_at = now
return len(conflicts)
def _group_observations(observations: Iterable[AISRawObservation]) -> dict[str, list[AISRawObservation]]:
grouped: dict[str, list[AISRawObservation]] = {}
for observation in observations:
grouped.setdefault(str(observation.entity_key), []).append(observation)
return grouped
async def record_vessel_ais_observation(
db: AsyncSession,
*,
source: str,
normalized_payload: dict[str, Any],
raw_payload: dict[str, Any] | None = None,
delivery_mode: str,
transport: str,
message_type: str | None = "PositionReport",
source_message_id: str | None = None,
observed_at: datetime | None = None,
collected_at: datetime | None = None,
quality_flags: list[str] | None = None,
) -> AISRawObservation | None:
"""Insert one raw observation if the source-level fact has not already been stored."""
entity_key = str(normalized_payload["mmsi"])
collected_at = collected_at or datetime.now(UTC)
observed_at = (
_coerce_datetime(observed_at)
or _coerce_datetime(normalized_payload.get("received_at"))
or collected_at
)
normalized_json = _jsonable(normalized_payload)
raw_json = _jsonable(raw_payload or {})
observation_hash = build_observation_hash(
source=source,
entity_key=entity_key,
message_type=message_type,
observed_at=observed_at,
normalized_payload=normalized_json,
source_message_id=source_message_id,
)
existing_result = await db.execute(
select(AISRawObservation.id).where(AISRawObservation.observation_hash == observation_hash)
)
if existing_result.scalar_one_or_none() is not None:
return None
observation = AISRawObservation(
target_schema=VESSEL_AIS_SCHEMA,
source=source,
entity_key=entity_key,
delivery_mode=delivery_mode,
transport=transport,
message_type=message_type,
source_message_id=source_message_id,
observation_hash=observation_hash,
observed_at=observed_at,
collected_at=collected_at,
normalized_payload=normalized_json,
raw_payload=raw_json,
quality_flags=quality_flags or [],
)
db.add(observation)
return observation
async def aggregate_vessel_observations(
db: AsyncSession,
observations: Iterable[AISRawObservation],
*,
write_conflicts: bool = False,
strategy: dict[str, Any] | None = None,
) -> list[dict[str, Any]]:
strategy = strategy if strategy is not None else await _safe_load_strategy(db)
now = datetime.now(UTC)
vessels = []
for entity_key, entity_observations in _group_observations(observations).items():
aggregated = _build_aggregated_vessel(
entity_key, entity_observations, now=now, strategy=strategy
)
if aggregated is None:
continue
if write_conflicts:
aggregated["conflict_count"] = await _upsert_conflict_records(
db,
entity_key,
entity_observations,
aggregated,
)
else:
aggregated["conflict_count"] = len(build_field_conflict_candidates(entity_observations))
vessels.append(aggregated)
vessels.sort(key=lambda item: item.get("received_at") or datetime.min.replace(tzinfo=UTC), reverse=True)
return vessels
async def _safe_load_strategy(db: AsyncSession) -> dict[str, Any]:
"""Tolerate fake test sessions where load_strategy may misbehave."""
try:
return await load_strategy(db)
except Exception:
return DEFAULT_STRATEGY
async def get_aggregated_vessels(
db: AsyncSession,
*,
bbox: tuple[float, float, float, float] | None = None,
limit: int | None = None,
observed_since: datetime | None = None,
) -> list[dict[str, Any]]:
observed_since = observed_since or (
datetime.now(UTC) - timedelta(hours=DEFAULT_AGGREGATION_WINDOW_HOURS)
)
stmt = (
select(AISRawObservation)
.where(AISRawObservation.target_schema == VESSEL_AIS_SCHEMA)
.where(AISRawObservation.observed_at >= observed_since)
.order_by(AISRawObservation.observed_at.desc(), AISRawObservation.id.desc())
)
if limit and limit > 0:
stmt = stmt.limit(max(limit * 20, limit))
result = await db.execute(stmt)
if not hasattr(result, "scalars"):
return []
vessels = await aggregate_vessel_observations(db, result.scalars().all())
if bbox is not None:
lon_min, lat_min, lon_max, lat_max = bbox
vessels = [
vessel
for vessel in vessels
if lon_min <= float(vessel["lon"]) <= lon_max
and lat_min <= float(vessel["lat"]) <= lat_max
]
if limit and limit > 0:
return vessels[:limit]
return vessels
async def get_aggregated_vessel(db: AsyncSession, mmsi: int) -> dict[str, Any] | None:
observations = await get_vessel_raw_observations(db, mmsi, limit=1000)
vessels = await aggregate_vessel_observations(db, observations)
return vessels[0] if vessels else None
async def get_aggregated_vessel_track(
db: AsyncSession,
mmsi: int,
*,
cutoff: datetime,
) -> list[dict[str, Any]]:
result = await db.execute(
select(AISRawObservation)
.where(AISRawObservation.target_schema == VESSEL_AIS_SCHEMA)
.where(AISRawObservation.entity_key == str(mmsi))
.where(AISRawObservation.observed_at >= cutoff)
.order_by(AISRawObservation.observed_at.asc(), AISRawObservation.id.asc())
)
if not hasattr(result, "scalars"):
return []
points: list[dict[str, Any]] = []
seen: set[tuple[str, float, float, str]] = set()
for observation in result.scalars().all():
payload = observation.normalized_payload or {}
if not _has_valid_position(payload):
continue
lat = float(payload["lat"])
lon = float(payload["lon"])
key = (
observation.observed_at.isoformat(),
round(lat, 5),
round(lon, 5),
observation.source,
)
if key in seen:
continue
seen.add(key)
points.append(
{
"lat": lat,
"lon": lon,
"observed_at": observation.observed_at,
"source": observation.source,
"selected_reason": "track_timeline",
"quality_flags": observation.quality_flags or [],
}
)
return points
async def update_ais_source_health(
db: AsyncSession,
*,
source: str,
connection_state: str,
observed_count: int = 0,
last_seen_at: datetime | None = None,
last_success_at: datetime | None = None,
last_error: str | None = None,
lag_seconds: float | None = None,
) -> AISSourceHealth:
"""Upsert the health row for an AIS source."""
now = datetime.now(UTC)
health = await db.get(AISSourceHealth, source)
if health is None:
health = AISSourceHealth(source=source)
db.add(health)
health.connection_state = connection_state
health.last_seen_at = last_seen_at or health.last_seen_at
health.last_success_at = last_success_at or health.last_success_at
health.last_error = last_error
health.message_rate = float(observed_count)
health.lag_seconds = lag_seconds
health.updated_at = now
return health
async def count_unique_raw_vessel_mmsi(
db: AsyncSession,
*,
observed_since: datetime | None = None,
) -> int:
"""Count unique raw vessel MMSI values for HUD counts; never aggregates."""
from sqlalchemy import func as sa_func
unique_mmsi_stmt = (
select(AISRawObservation.entity_key)
.where(AISRawObservation.target_schema == VESSEL_AIS_SCHEMA)
.distinct()
)
if observed_since is not None:
unique_mmsi_stmt = unique_mmsi_stmt.where(
AISRawObservation.observed_at >= observed_since,
)
result = await db.execute(
select(sa_func.count()).select_from(unique_mmsi_stmt.subquery()),
)
return int(result.scalar() or 0)
async def get_vessel_raw_observations(
db: AsyncSession,
mmsi: int,
*,
limit: int = 100,
) -> list[AISRawObservation]:
result = await db.execute(
select(AISRawObservation)
.where(AISRawObservation.target_schema == VESSEL_AIS_SCHEMA)
.where(AISRawObservation.entity_key == str(mmsi))
.order_by(AISRawObservation.observed_at.desc(), AISRawObservation.id.desc())
.limit(limit)
)
return list(result.scalars().all())
async def get_vessel_conflict_records(
db: AsyncSession,
mmsi: int,
) -> list[AISConflictRecord]:
result = await db.execute(
select(AISConflictRecord)
.where(AISConflictRecord.target_schema == VESSEL_AIS_SCHEMA)
.where(AISConflictRecord.entity_key == str(mmsi))
.order_by(AISConflictRecord.updated_at.desc(), AISConflictRecord.id.desc())
)
return list(result.scalars().all())

View File

@@ -0,0 +1,109 @@
"""v5 vessel enrichment service.
Read-only side: `get_vessel_enrichment_bundle` is the only path the
aggregation/detail endpoints use. It never reaches out to third parties; it
just returns whatever the upsert side has already cached. Expired rows are
filtered out so old data never leaks back into the live UI.
"""
from __future__ import annotations
from datetime import UTC, datetime
from typing import Any
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.vessel_enrichment import VesselMediaEnrichment, VesselProfileEnrichment
def _coerce_datetime(value: Any) -> datetime | None:
if value in (None, ""):
return None
if isinstance(value, datetime):
return value if value.tzinfo else value.replace(tzinfo=UTC)
if isinstance(value, (int, float)):
ts = float(value)
if ts > 10_000_000_000:
ts /= 1000
return datetime.fromtimestamp(ts, UTC)
if isinstance(value, str):
try:
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
return parsed if parsed.tzinfo else parsed.replace(tzinfo=UTC)
except ValueError:
return None
return None
def _build_payload(record, *, now: datetime) -> dict[str, Any] | None:
if record is None:
return None
expires_at = record.expires_at
if isinstance(expires_at, datetime):
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=UTC)
if expires_at < now:
return None
return record.to_dict()
async def get_vessel_enrichment_bundle(db: AsyncSession, mmsi: int) -> dict[str, Any]:
now = datetime.now(UTC)
profile = await db.get(VesselProfileEnrichment, mmsi)
media = await db.get(VesselMediaEnrichment, mmsi)
return {
"mmsi": mmsi,
"profile": _build_payload(profile, now=now),
"media": _build_payload(media, now=now),
}
async def upsert_vessel_profile_enrichment(
db: AsyncSession,
*,
mmsi: int,
payload: dict[str, Any],
) -> dict[str, Any]:
record = await db.get(VesselProfileEnrichment, mmsi)
if record is None:
record = VesselProfileEnrichment(mmsi=mmsi)
db.add(record)
return _apply_upsert(record, payload)
async def upsert_vessel_media_enrichment(
db: AsyncSession,
*,
mmsi: int,
payload: dict[str, Any],
) -> dict[str, Any]:
record = await db.get(VesselMediaEnrichment, mmsi)
if record is None:
record = VesselMediaEnrichment(mmsi=mmsi)
db.add(record)
return _apply_upsert(record, payload)
def _apply_upsert(record, payload: dict[str, Any]) -> dict[str, Any]:
if not isinstance(payload, dict):
raise ValueError("enrichment payload must be an object")
body = payload.get("payload")
if body is not None and not isinstance(body, dict):
raise ValueError("payload.payload must be an object")
if body is not None:
record.payload = body
if "source" in payload and isinstance(payload["source"], str) and payload["source"].strip():
record.source = payload["source"].strip()
fetched_at = _coerce_datetime(payload.get("fetched_at"))
record.fetched_at = fetched_at or datetime.now(UTC)
record.expires_at = _coerce_datetime(payload.get("expires_at"))
confidence = payload.get("confidence")
if confidence is not None:
try:
record.confidence = float(confidence)
except (TypeError, ValueError):
record.confidence = None
if "reference_url" in payload:
ref = payload.get("reference_url")
record.reference_url = str(ref) if ref else None
return record.to_dict()

View File

@@ -0,0 +1,31 @@
"""Shared AIS vessel type helpers."""
from typing import Any
VESSEL_TYPE_NAMES = {
30: "Fishing",
35: "Military",
60: "Passenger",
70: "Cargo",
80: "Tanker",
}
def normalize_vessel_type_name(vessel_type: Any) -> str:
"""Map AIS numeric vessel type codes to display buckets."""
try:
type_code = int(float(vessel_type))
except (TypeError, ValueError):
return "Other"
if 70 <= type_code <= 79:
return "Cargo"
if 80 <= type_code <= 89:
return "Tanker"
if 60 <= type_code <= 69:
return "Passenger"
if type_code == 30:
return "Fishing"
if type_code == 35:
return "Military"
return VESSEL_TYPE_NAMES.get(type_code, "Other")

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
import argparse
import os
import shlex
import subprocess
import sys
import time
@@ -59,6 +60,8 @@ def wait_for_recovery(action: str) -> tuple[bool, str]:
recovery_mode = get_action_recovery_mode(action)
if recovery_mode == "backend":
return wait_for_http("http://localhost:8000/health"), "backend health recovery"
if recovery_mode == "frontend":
return wait_for_http("http://localhost:3000"), "frontend entrypoint recovery"
if recovery_mode == "ai-provider":
return wait_for_http("http://localhost:8010/health"), "ai provider health recovery"
if recovery_mode == "database":
@@ -108,8 +111,9 @@ def main() -> int:
)
append_task_log(args.task_id, "restart command started")
shell_command = shlex.join(command)
completed = subprocess.run(
command,
["zsh", "-ic", shell_command],
cwd=str(ROOT_DIR),
env=env,
capture_output=True,

View File

@@ -2,10 +2,45 @@
import pytest
import asyncio
from typing import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock, patch
import json
from unittest.mock import AsyncMock, MagicMock
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker
from sqlalchemy.ext.asyncio import AsyncSession
@pytest.fixture(autouse=True)
def bgp_collector_location_cache():
"""Mirror app startup seeding for tests that call sync BGP helpers."""
from app.services.bgp_collector_locations import (
SEED_PATH,
set_bgp_collector_location_cache,
)
payload = json.loads(SEED_PATH.read_text(encoding="utf-8"))
cache = {}
for entry in payload.get("locations", []):
collector_id = next(
alias for alias in entry.get("aliases", []) if str(alias).startswith("rrc")
)
cache[collector_id] = {
"city": entry.get("city"),
"country": entry.get("country"),
"latitude": entry.get("latitude"),
"longitude": entry.get("longitude"),
"precision": entry.get("precision") or "city",
"source": "legacy_seed",
"needs_confirmation": True,
"matched_location_name": entry.get("site") or collector_id,
"verified_at": None,
"confidence": entry.get("confidence"),
"operator": entry.get("operator"),
"site": entry.get("site"),
"verification_status": "unverified",
"source_note": entry.get("source_note"),
}
set_bgp_collector_location_cache(cache)
yield
set_bgp_collector_location_cache({})
@pytest.fixture(scope="session")

View File

@@ -1,6 +1,7 @@
"""Tests for BGP observability helpers."""
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
import pytest
from httpx import ASGITransport, AsyncClient
@@ -54,11 +55,34 @@ class _FakeResult:
def scalars(self):
return _FakeScalarResult(self._rows)
def all(self):
if self._rows and all(isinstance(row, BGPObservation) for row in self._rows):
return [
(row.prefix, row.origin_asn, row.collector, row.collector_geo)
for row in self._rows
]
return self._rows
def scalar(self):
if not self._rows:
return 0
first = self._rows[0]
if isinstance(first, (int, float, str)):
return first
if isinstance(first, tuple) and len(first) == 1:
return first[0]
return len(self._rows)
def fetchall(self):
return self._rows
def fetchone(self):
return self._rows[0] if self._rows else None
if not self._rows:
return None
first = self._rows[0]
if isinstance(first, CollectedData):
return {"extra_data": first.extra_data}
return first
class _FakeAsyncSession:
@@ -988,7 +1012,7 @@ async def test_infer_related_infrastructure_links_nearby_cables():
data_type="cable",
extra_data={"cable_id": 20},
)
db = _FakeAsyncSession([[landing], [relation], [cable]])
db = _FakeAsyncSession([[landing, relation, cable]])
result = await infer_related_infrastructure(
db,
@@ -1012,27 +1036,37 @@ async def test_infer_related_infrastructure_links_nearby_cables():
@pytest.mark.asyncio
async def test_build_bgp_collector_coverage_summarizes_observations():
now = datetime.now(UTC)
obs_one = BGPObservation(
source="ris_live_bgp",
aggregate = SimpleNamespace(
collector="rrc00",
observation_count=2,
prefix_count=2,
origin_asn_count=2,
peer_asn_count=2,
recent_15m_observation_count=2,
recent_24h_observation_count=2,
recent_7d_observation_count=2,
recent_15m_prefix_count=2,
recent_24h_prefix_count=2,
recent_7d_prefix_count=2,
latest_observed_at=now + timedelta(minutes=5),
)
latest = SimpleNamespace(
collector="rrc00",
latest_event_type="withdrawal",
country="Netherlands",
city="Amsterdam",
)
top_event = SimpleNamespace(
collector="rrc00",
prefix="203.0.113.0/24",
origin_asn=64496,
peer_asn=3333,
event_type="announcement",
observed_at=now,
collector_geo={"city": "Amsterdam", "country": "Netherlands"},
count=1,
)
obs_two = BGPObservation(
source="ris_live_bgp",
scope = SimpleNamespace(
collector="rrc00",
prefix="198.51.100.0/24",
origin_asn=64497,
peer_asn=3334,
event_type="withdrawal",
observed_at=now + timedelta(minutes=5),
collector_geo={"city": "Amsterdam", "country": "Netherlands"},
country="Netherlands",
city="Amsterdam",
)
db = _FakeAsyncSession([[obs_one, obs_two]])
db = _FakeAsyncSession([[aggregate], [latest], [top_event], [scope]])
coverage = await build_bgp_collector_coverage(db, source_filter=BGP_SOURCES)
@@ -1363,18 +1397,39 @@ async def test_bgp_event_summary_api_returns_aggregates():
@pytest.mark.asyncio
async def test_bgp_collectors_api_returns_coverage():
now = datetime.now(UTC)
observation = BGPObservation(
id=1,
source="ris_live_bgp",
aggregate = SimpleNamespace(
collector="rrc00",
peer_asn=3333,
prefix="203.0.113.0/24",
event_type="announcement",
origin_asn=64496,
observed_at=now,
collector_geo={"city": "Amsterdam", "country": "Netherlands"},
observation_count=1,
prefix_count=1,
origin_asn_count=1,
peer_asn_count=1,
recent_15m_observation_count=1,
recent_24h_observation_count=1,
recent_7d_observation_count=1,
recent_15m_prefix_count=1,
recent_24h_prefix_count=1,
recent_7d_prefix_count=1,
latest_observed_at=now,
)
latest = SimpleNamespace(
collector="rrc00",
latest_event_type="announcement",
country="Netherlands",
city="Amsterdam",
)
top_event = SimpleNamespace(
collector="rrc00",
event_type="announcement",
count=1,
)
scope = SimpleNamespace(
collector="rrc00",
country="Netherlands",
city="Amsterdam",
)
db = _FakeAsyncSession(
[[aggregate], [latest], [top_event], [scope], [aggregate], [latest], [top_event], [scope]]
)
db = _FakeAsyncSession([[observation], [observation]])
client = await _bgp_test_client(db)
try:

View File

@@ -0,0 +1,222 @@
"""Tests for the BGP collector + event location services."""
from __future__ import annotations
from unittest.mock import AsyncMock
import pytest
from app.api.v1 import bgp as bgp_api
from app.services import bgp_collector_locations
from app.services.location.llm_fallback import LocationLLMFallbackResult
from app.services.bgp_collector_locations import (
RIPE_RIS_COLLECTOR_COORDS,
collect_bgp_collector_location_candidates,
iter_known_collector_names,
resolve_bgp_collector_location,
)
from app.services.bgp_event_locations import (
resolve_bgp_event_geo_dict,
resolve_bgp_event_location,
)
def test_legacy_dict_view_preserves_backward_compatible_keys():
rrc00 = RIPE_RIS_COLLECTOR_COORDS["rrc00"]
assert rrc00["city"] == "Amsterdam"
assert rrc00["country"] == "Netherlands"
assert rrc00["latitude"] == pytest.approx(52.3676)
assert rrc00["longitude"] == pytest.approx(4.9041)
# New richer fields layered on top.
assert rrc00["precision"] == "city"
assert rrc00["source"] == "legacy_seed"
assert rrc00["needs_confirmation"] is True
def test_every_legacy_collector_present():
expected = {
"rrc00", "rrc01", "rrc03", "rrc04", "rrc05", "rrc06", "rrc07",
"rrc10", "rrc11", "rrc12", "rrc13", "rrc14", "rrc15", "rrc16",
"rrc18", "rrc19", "rrc20", "rrc21", "rrc22", "rrc23", "rrc24",
"rrc25", "rrc26",
}
assert set(iter_known_collector_names()) == expected
def test_resolve_bgp_collector_returns_stored_location():
result = resolve_bgp_collector_location("rrc12")
assert result.location is not None
assert result.location.city == "Frankfurt"
assert result.location.country == "Germany"
assert result.location.precision == "city"
assert result.location.source == "legacy_seed"
assert result.location.needs_confirmation is True
def test_resolve_unknown_bgp_collector_returns_diagnostic(monkeypatch):
monkeypatch.setattr(bgp_collector_locations, "_geocode_online", lambda q: None)
result = resolve_bgp_collector_location("rrc-doesnotexist")
assert result.location is None
assert result.diagnostic is not None
assert result.diagnostic.failure_reason
def test_collect_bgp_collector_candidates_uses_stored_context_without_registry(monkeypatch):
bgp_collector_locations._geocode_online.cache_clear()
def _fake_geocode(query):
assert "CIXP" in query or "Geneva" in query
return {
"lat": "46.2044",
"lon": "6.1432",
"display_name": "Geneva, Switzerland",
"address": {"city": "Geneva", "country": "Switzerland"},
}
monkeypatch.setattr(bgp_collector_locations, "_geocode_online", _fake_geocode)
candidates, attempted = collect_bgp_collector_location_candidates(
collector="rrc04",
)
assert attempted, "stored context should feed online query attempts"
assert candidates, "online geocoding should produce at least one candidate"
best = candidates[0]
assert best.source == "nominatim_online_geocode"
assert best.needs_confirmation is True
assert all(candidate.source != "local_registry" for candidate in candidates)
def test_collect_bgp_collector_candidates_uses_nominatim_when_registry_misses(monkeypatch):
bgp_collector_locations._geocode_online.cache_clear()
def _fake_geocode(query):
if "Lyon" not in query and "France-IX" not in query and "FR-IX" not in query:
return None
return {
"lat": "45.764",
"lon": "4.8357",
"display_name": "Lyon, Auvergne-Rhône-Alpes, France",
"address": {"city": "Lyon", "country": "France"},
}
monkeypatch.setattr(bgp_collector_locations, "_geocode_online", _fake_geocode)
candidates, attempted = collect_bgp_collector_location_candidates(
collector="rrc-mystery",
city="Lyon",
country="France",
)
assert attempted, "Nominatim plan should run"
online = [c for c in candidates if c.source == "nominatim_online_geocode"]
assert online, "online resolver must produce a candidate when registry misses"
assert online[0].needs_confirmation is True
@pytest.mark.asyncio
async def test_collect_bgp_collector_location_uses_llm_when_candidates_empty(monkeypatch):
llm_candidate = bgp_collector_locations.LocationCandidate(
latitude=45.764,
longitude=4.8357,
display_name="Lyon, France",
precision="city",
confidence=0.74,
query="llm_factcheck:bgp_collector:rrc-mystery",
source="llm_location_factcheck",
source_note="LLM location factcheck fallback",
matched_fields=("collector",),
needs_confirmation=True,
city="Lyon",
country="France",
)
monkeypatch.setattr(
bgp_api,
"get_bgp_collector_location_dict",
lambda _collector: {},
)
monkeypatch.setattr(
bgp_api,
"collect_bgp_collector_location_candidates",
lambda **_kwargs: ([], ["Lyon, France"]),
)
from app.services.location.llm_fallback import LocationSearchEvidenceResult
async def _search_evidence(**_kwargs):
return LocationSearchEvidenceResult(
evidence=[
{
"title": "RRC source",
"url": "https://example.test/rrc",
"snippet": "rrc-mystery is in Lyon.",
}
],
attempted_queries=["web_search:bgp_collector:rrc-mystery Lyon France physical location route collector city"],
)
async def _fallback(**_kwargs):
return LocationLLMFallbackResult(
candidates=[llm_candidate],
attempted_queries=["llm_factcheck:bgp_collector:rrc-mystery"],
)
monkeypatch.setattr(bgp_api, "get_ai_provider_client", AsyncMock(return_value=object()))
monkeypatch.setattr(bgp_api, "get_web_search_client", AsyncMock(return_value=object()))
monkeypatch.setattr(bgp_api, "collect_location_search_evidence", _search_evidence)
monkeypatch.setattr(bgp_api, "collect_llm_location_fallback_candidate", _fallback)
response = await bgp_api.collect_bgp_collector_location(
"rrc-mystery",
bgp_api.CollectBGPCollectorLocationRequest(city="Lyon", country="France"),
current_user=object(),
db=AsyncMock(),
)
assert response["success"] is True
assert response["best_candidate"]["source"] == "llm_location_factcheck"
assert response["best_candidate"]["needs_confirmation"] is True
assert response["attempted_queries"] == [
"Lyon, France",
"web_search:bgp_collector:rrc-mystery Lyon France physical location route collector city",
"llm_factcheck:bgp_collector:rrc-mystery",
]
# ── BGP event resolver ─────────────────────────────────────────────
def test_event_resolver_inherits_from_owning_collector():
geo = resolve_bgp_event_geo_dict("rrc25")
assert geo["city"] == "Amsterdam"
assert geo["country"] == "Netherlands"
assert geo["source"] == "inherited_from_collector"
assert geo["precision"] == "city"
def test_event_resolver_does_not_match_unrelated_collectors():
"""Regression: passing operator=RIPE NCC must NOT make every collector match."""
rrc12 = resolve_bgp_event_geo_dict("rrc12")
rrc25 = resolve_bgp_event_geo_dict("rrc25")
assert rrc12["city"] == "Frankfurt"
assert rrc25["city"] == "Amsterdam"
assert rrc12["latitude"] != rrc25["latitude"]
def test_event_resolver_uses_source_coordinates_when_present():
geo = resolve_bgp_event_geo_dict(
"rrc12",
source_latitude=12.34,
source_longitude=56.78,
)
assert geo["latitude"] == pytest.approx(12.34)
assert geo["longitude"] == pytest.approx(56.78)
assert geo["precision"] == "precise"
assert geo["source"] == "source_coordinates"
def test_event_resolver_returns_empty_for_unknown_collector_without_source_coords():
geo = resolve_bgp_event_geo_dict("rrc-doesnotexist")
assert geo == {}
def test_event_resolver_full_result_carries_diagnostic_on_miss():
result = resolve_bgp_event_location(collector="rrc-doesnotexist")
assert result.location is None
assert result.diagnostic is not None

View File

@@ -1,11 +1,14 @@
"""Unit tests for data collectors"""
import pytest
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, patch
from app.core.datasource_defaults import DEFAULT_DATASOURCES
from app.services.credential_guides import DEFAULT_CREDENTIAL_GUIDES
from app.services.collectors.top500 import TOP500Collector
from app.services.collectors.base import BaseCollector, HTTPCollector
from app.services.collectors.registry import collector_registry
from app.services.datasource_connectivity import SUPPORTED_CREDENTIAL_PROVIDERS
from app.models.task import CollectionTask
class TestBaseCollector:
@@ -19,6 +22,31 @@ class TestBaseCollector:
assert collector.module == "L1"
assert collector.frequency_hours == 4
@pytest.mark.asyncio
async def test_update_phase_progress_tracks_phase_fields(self, mock_db_session):
"""Test phase-level progress updates independently from record totals"""
collector = TOP500Collector()
task = CollectionTask(datasource_id=1, status="running", phase="fetching")
collector._current_task = task
collector._db_session = mock_db_session
with patch.object(collector, "_publish_task_update", new=AsyncMock()) as publish:
await collector.update_phase_progress(
current=512,
total=1024,
unit="bytes",
message="Downloading dataset",
commit=True,
)
assert task.phase_progress == 50.0
assert task.phase_current == 512
assert task.phase_total == 1024
assert task.phase_unit == "bytes"
assert task.phase_message == "Downloading dataset"
mock_db_session.commit.assert_awaited_once()
publish.assert_awaited_once()
class TestTOP500Collector:
"""Tests for TOP500Collector"""
@@ -119,3 +147,30 @@ class TestHTTPCollector:
assert hasattr(collector, "parse_response")
assert callable(collector.fetch)
assert callable(collector.parse_response)
def test_aisstream_collector_is_registered():
collector = collector_registry.get("aisstream_vessels")
assert collector is not None
assert collector.data_type == "vessel_ais"
def test_supported_credential_collectors_have_guides_and_connectivity_provider():
missing: list[str] = []
for source, info in DEFAULT_DATASOURCES.items():
if not info.get("requires_credentials"):
continue
if info.get("credential_status") != "supported":
continue
provider = info.get("credential_provider")
if not provider:
missing.append(f"{source}: missing credential_provider")
continue
if provider not in DEFAULT_CREDENTIAL_GUIDES:
missing.append(f"{source}: missing credential guide for {provider}")
if provider not in SUPPORTED_CREDENTIAL_PROVIDERS:
missing.append(f"{source}: missing connectivity provider for {provider}")
assert missing == []

View File

@@ -0,0 +1,149 @@
"""End-to-end integration test for the custom WebSocket datasource runner.
Boots an in-process WebSocket server that mimics the bun mock AIS server
(`scripts/mock-ais-ws-server.ts`) and runs the real
`run_mapped_websocket_config` against it. Catches regressions where the
runner stops connecting, fails to extract the configured message path,
or quietly drops mapped records before broadcasting.
"""
from __future__ import annotations
import asyncio
import json
from contextlib import asynccontextmanager
from datetime import UTC, datetime
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import websockets
from app.models.datasource_config import DataSourceConfig
from app.services import custom_datasource_runtime
from app.services.custom_datasource_runtime import run_mapped_websocket_config
def _make_payload(seq: int) -> str:
return json.dumps(
{
"type": "vessel",
"sequence": seq,
"data": {
"mmsi": str(999_000_000 + seq),
"name": f"MOCK VESSEL {seq:03d}",
"lat": 36.20 + seq * 0.001,
"lon": 14.20 + seq * 0.001,
"sog": 12.0,
"cog": 90.0,
"heading": 90,
"vessel_type": 70,
"vessel_type_name": "Cargo",
"received_at": datetime.now(UTC).isoformat(),
},
}
)
@asynccontextmanager
async def _mock_ais_server(emit_count: int):
received_subscribe: list[str] = []
async def handler(ws):
try:
try:
msg = await asyncio.wait_for(ws.recv(), timeout=0.5)
received_subscribe.append(msg)
except (asyncio.TimeoutError, websockets.ConnectionClosed):
pass
for seq in range(1, emit_count + 1):
await ws.send(_make_payload(seq))
await asyncio.sleep(0.01)
# keep the socket open briefly so the runner observes the messages
await asyncio.sleep(0.05)
except websockets.ConnectionClosed:
return
async with websockets.serve(handler, "127.0.0.1", 0) as server:
port = next(iter(server.sockets)).getsockname()[1]
yield port, received_subscribe
@pytest.mark.asyncio
async def test_websocket_runner_streams_from_live_mock(monkeypatch):
mapping = SimpleNamespace(
id=11,
version=3,
target_schema="vessel_ais",
mapping_json={
"source": {"items_path": "$"},
"fields": {
"mmsi": {"path": "$.mmsi", "type": "integer"},
"lat": {"path": "$.lat", "type": "float"},
"lon": {"path": "$.lon", "type": "float"},
"name": {"path": "$.name", "type": "string"},
"vessel_type": {"path": "$.vessel_type", "type": "integer", "default": None},
"vessel_type_name": {"path": "$.vessel_type_name", "type": "string", "default": None},
"sog": {"path": "$.sog", "type": "float", "default": None},
"cog": {"path": "$.cog", "type": "float", "default": None},
"heading": {"path": "$.heading", "type": "integer", "default": None},
"received_at": {"path": "$.received_at", "type": "datetime"},
},
},
)
class FakeResult:
def scalar_one_or_none(self):
return mapping
class FakeDB:
async def execute(self, _stmt):
return FakeResult()
persist = AsyncMock(return_value=1)
monkeypatch.setattr(custom_datasource_runtime, "persist_mapped_records", persist)
async with _mock_ais_server(emit_count=3) as (port, received_subscribe):
result = await run_mapped_websocket_config(
FakeDB(),
DataSourceConfig(
id=99,
name="mock_ais_ws",
source_type="websocket",
endpoint=f"ws://127.0.0.1:{port}",
auth_type="none",
headers={},
config={
"ws_message_path": "$.data",
"ws_subscribe_message": {
"type": "subscribe",
"anchor": {"lat": 36.2, "lon": 14.2},
"spread_km": 50,
"rate_hz": 1,
},
"debug_max_messages": 2,
"delivery_mode": "realtime_stream",
"ws_reconnect": False,
},
),
use_config_debug_max_messages=True,
)
assert result["status"] == "success"
assert result["messages_seen"] == 2
assert result["written_count"] == 2
assert result["mapped_count"] == 2
assert result["target_schema"] == "vessel_ais"
# subscribe message must reach the server unchanged
assert received_subscribe, "runner did not forward ws_subscribe_message"
parsed = json.loads(received_subscribe[0])
assert parsed["type"] == "subscribe"
assert parsed["anchor"] == {"lat": 36.2, "lon": 14.2}
assert parsed["rate_hz"] == 1
# mapped records carry the real MMSIs from the mock stream
persisted_records = []
for call in persist.await_args_list:
persisted_records.extend(call.kwargs["records"])
assert {record["mmsi"] for record in persisted_records} == {999_000_001, 999_000_002}
assert all(record["vessel_type"] == 70 for record in persisted_records)
assert all(record["vessel_type_name"] == "Cargo" for record in persisted_records)

View File

@@ -0,0 +1,328 @@
from types import SimpleNamespace
import pytest
from unittest.mock import AsyncMock
from httpx import ASGITransport, AsyncClient
from app.api.v1.datasource_config import get_ai_provider_client
from app.core.websocket import broadcaster as broadcaster_module
from app.core.security import get_current_user
from app.core.target_schema_registry import get_target_schema, list_target_schemas
from app.main import app
from app.models.user import User
from app.models.datasource_config import DataSourceConfig
from app.services import custom_datasource_runtime
from app.services.custom_datasource_runtime import run_mapped_websocket_config
from app.services.datasource_mapping import execute_mapping, persist_mapped_records, redact_for_llm
SAMPLE_AIS = {
"data": [
{
"mmsi": "257123000",
"latitude": "59.91",
"longitude": "10.75",
"speedOverGround": "12.4",
"timestamp": "2026-04-28T00:00:00Z",
"api_token": "secret-value",
}
]
}
def test_registry_exposes_v1_target_schemas():
keys = {schema["key"] for schema in list_target_schemas()}
assert {"vessel_ais", "geo_points", "generic_records"}.issubset(keys)
assert get_target_schema("vessel_ais").destination == "vessel_position"
def test_mapping_engine_maps_and_validates_vessel_ais():
mapping = {
"source": {"items_path": "$.data[*]"},
"fields": {
"mmsi": {"path": "$.mmsi", "type": "integer"},
"lat": {"path": "$.latitude", "type": "float"},
"lon": {"path": "$.longitude", "type": "float"},
"sog": {"path": "$.speedOverGround", "type": "float"},
"received_at": {"path": "$.timestamp", "type": "datetime"},
},
}
result = execute_mapping(SAMPLE_AIS, mapping, "vessel_ais")
assert result["mapped_count"] == 1
assert result["failed_count"] == 0
assert result["records"][0]["mmsi"] == 257123000
assert result["records"][0]["lat"] == 59.91
def test_mapping_engine_reports_schema_errors():
mapping = {
"source": {"items_path": "$.data[*]"},
"fields": {
"mmsi": {"path": "$.mmsi", "type": "integer"},
"lat": {"path": "$.missing_lat", "type": "float"},
"lon": {"path": "$.longitude", "type": "float"},
},
}
result = execute_mapping(SAMPLE_AIS, mapping, "vessel_ais")
assert result["mapped_count"] == 0
assert result["failed_count"] == 1
assert any("lat" in error for error in result["errors"][0]["errors"])
def test_redact_for_llm_masks_secret_like_fields():
redacted = redact_for_llm(SAMPLE_AIS)
assert redacted["data"][0]["api_token"] == "[REDACTED]"
@pytest.mark.asyncio
async def test_persist_mapped_records_writes_generic_records():
class FakeDB:
def __init__(self):
self.added = []
self.committed = False
def add(self, value):
self.added.append(value)
async def commit(self):
self.committed = True
db = FakeDB()
count = await persist_mapped_records(
db,
datasource_name="custom_weather",
datasource_config_id=42,
target_schema="generic_records",
records=[{"source_id": "row-1", "data": {"temp": 25}}],
mapping_version=3,
)
assert count == 1
assert db.committed is True
assert db.added[0].source == "custom_weather"
assert db.added[0].data_type == "generic_records"
assert db.added[0].extra_data["mapping_version"] == 3
@pytest.mark.asyncio
async def test_persist_mapped_vessel_records_writes_raw_and_broadcasts(monkeypatch):
record_observation = AsyncMock(return_value=object())
update_health = AsyncMock()
broadcast_custom = AsyncMock()
monkeypatch.setattr(
"app.services.vessel_ais_aggregation.record_vessel_ais_observation",
record_observation,
)
monkeypatch.setattr(
"app.services.vessel_ais_aggregation.update_ais_source_health",
update_health,
)
monkeypatch.setattr(broadcaster_module, "broadcast_custom", broadcast_custom)
class FakeDB:
def __init__(self):
self.committed = False
async def commit(self):
self.committed = True
db = FakeDB()
count = await persist_mapped_records(
db,
datasource_name="mock_ais_ws",
datasource_config_id=42,
target_schema="vessel_ais",
records=[
{
"mmsi": 999000001,
"lat": 31.2,
"lon": 121.4,
"name": "MOCK VESSEL 001",
"received_at": "2026-05-01T00:00:00Z",
}
],
mapping_version=1,
delivery_mode="realtime_stream",
transport="websocket",
)
assert count == 1
assert db.committed is True
record_observation.assert_awaited_once()
assert record_observation.await_args.kwargs["source"] == "mock_ais_ws"
assert record_observation.await_args.kwargs["delivery_mode"] == "realtime_stream"
assert record_observation.await_args.kwargs["transport"] == "websocket"
update_health.assert_awaited_once()
broadcast_custom.assert_awaited_once()
assert broadcast_custom.await_args.args[0] == "vessels"
assert broadcast_custom.await_args.args[1]["vessels"][0]["mmsi_display"] == "999000001"
@pytest.mark.asyncio
async def test_custom_websocket_runner_maps_and_persists_vessel_records(monkeypatch):
mapping = SimpleNamespace(
id=7,
version=2,
target_schema="vessel_ais",
mapping_json={
"source": {"items_path": "$"},
"fields": {
"mmsi": {"path": "$.mmsi", "type": "integer"},
"lat": {"path": "$.lat", "type": "float"},
"lon": {"path": "$.lon", "type": "float"},
"name": {"path": "$.name", "type": "string"},
"received_at": {"path": "$.received_at", "type": "datetime"},
},
},
)
class FakeResult:
def scalar_one_or_none(self):
return mapping
class FakeDB:
async def execute(self, _stmt):
return FakeResult()
class FakeWebSocket:
async def __aenter__(self):
return self
async def __aexit__(self, *_args):
return None
async def send(self, _message):
return None
async def recv(self):
return (
'{"type":"vessel","data":{"mmsi":"999000001","name":"MOCK VESSEL 001",'
'"lat":31.2,"lon":121.4,"received_at":"2026-05-01T00:00:00Z"}}'
)
persist = AsyncMock(return_value=1)
monkeypatch.setattr(custom_datasource_runtime, "_connect_websocket", AsyncMock(return_value=FakeWebSocket()))
monkeypatch.setattr(custom_datasource_runtime, "persist_mapped_records", persist)
result = await run_mapped_websocket_config(
FakeDB(),
DataSourceConfig(
id=42,
name="mock_ais_ws",
source_type="websocket",
endpoint="ws://localhost:8787/ais",
auth_type="none",
headers={},
config={"ws_message_path": "$.data", "debug_max_messages": 1},
),
)
assert result["status"] == "success"
assert result["messages_seen"] == 1
assert result["written_count"] == 1
persist.assert_awaited_once()
assert persist.await_args.kwargs["datasource_name"] == "mock_ais_ws"
assert persist.await_args.kwargs["records"][0]["mmsi"] == 999000001
assert persist.await_args.kwargs["delivery_mode"] == "realtime_stream"
assert persist.await_args.kwargs["transport"] == "websocket"
@pytest.mark.asyncio
async def test_mapping_preview_api_uses_deterministic_engine():
def override_get_current_user():
return User(
id=1,
username="testuser",
email="test@example.com",
password_hash="hashed",
role="admin",
is_active=True,
)
app.dependency_overrides = {get_current_user: override_get_current_user}
transport = ASGITransport(app=app)
try:
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/api/v1/datasources/mappings/preview",
json={
"sample_payload": SAMPLE_AIS,
"target_schema": "vessel_ais",
"mapping_json": {
"source": {"items_path": "$.data[*]"},
"fields": {
"mmsi": {"path": "$.mmsi", "type": "integer"},
"lat": {"path": "$.latitude", "type": "float"},
"lon": {"path": "$.longitude", "type": "float"},
},
},
},
)
finally:
app.dependency_overrides.clear()
assert response.status_code == 200
payload = response.json()
assert payload["success"] is True
assert payload["preview"]["records"][0]["mmsi"] == 257123000
@pytest.mark.asyncio
async def test_mapping_propose_api_redacts_sample_before_ai():
seen_context = {}
class FakeAIClient:
async def analyze(self, request, request_id=None):
seen_context.update(request.context)
return SimpleNamespace(
content=(
'{"source":{"items_path":"$.data[*]"},"fields":{'
'"mmsi":{"path":"$.mmsi","type":"integer"},'
'"lat":{"path":"$.latitude","type":"float"},'
'"lon":{"path":"$.longitude","type":"float"}}}'
)
)
def override_get_current_user():
return User(
id=1,
username="testuser",
email="test@example.com",
password_hash="hashed",
role="admin",
is_active=True,
)
def override_ai_client():
return FakeAIClient()
app.dependency_overrides = {
get_current_user: override_get_current_user,
get_ai_provider_client: override_ai_client,
}
transport = ASGITransport(app=app)
try:
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/api/v1/datasources/mappings/propose",
json={
"sample_payload": SAMPLE_AIS,
"target_schema": "vessel_ais",
"use_ai": True,
},
)
finally:
app.dependency_overrides.clear()
assert response.status_code == 200
payload = response.json()
assert payload["mapping_json"]["meta"]["generated_by"] == "ai_provider"
assert seen_context["sample_payload"]["data"][0]["api_token"] == "[REDACTED]"

View File

@@ -0,0 +1,116 @@
"""Docs Gatekeeper API tests."""
import pytest
from httpx import ASGITransport, AsyncClient
from app.api.v1 import docs as docs_api
from app.main import app
from app.models.user import User
def make_user(role: str = "viewer", groups: list[str] | None = None) -> User:
user = User(
id=1,
username="docs-user",
email="docs@example.com",
password_hash="x",
role=role,
is_active=True,
)
user.gatekeeper_groups = groups or []
return user
async def get_json(path: str, user: User | None = None):
if user is not None:
async def override_user():
return user
app.dependency_overrides[docs_api.get_optional_current_user] = override_user
transport = ASGITransport(app=app)
try:
async with AsyncClient(transport=transport, base_url="http://test") as client:
return await client.get(path)
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_public_catalog_only_for_anonymous_user():
response = await get_json("/api/v1/docs/catalog")
assert response.status_code == 200
items = response.json()["items"]
assert {item["access"] for item in items} == {"public"}
assert {item["slug"] for item in items if item["lang"] == "zh"} == {
"overview",
"quickstart",
"manual",
"faq",
"location-pipeline-user",
}
@pytest.mark.asyncio
async def test_anonymous_can_read_public_doc():
response = await get_json("/api/v1/docs/zh/quickstart")
assert response.status_code == 200
assert response.json()["access"] == "public"
assert "快速开始" in response.json()["markdown"]
@pytest.mark.asyncio
async def test_anonymous_protected_doc_requires_authentication():
response = await get_json("/api/v1/docs/zh/backend-collectors")
assert response.status_code == 401
@pytest.mark.asyncio
async def test_viewer_without_group_cannot_read_developer_doc():
response = await get_json(
"/api/v1/docs/zh/backend-collectors",
make_user(role="viewer"),
)
assert response.status_code == 403
@pytest.mark.asyncio
async def test_developer_group_can_read_developer_but_not_admin_doc():
user = make_user(role="viewer", groups=["docs_developer"])
developer_response = await get_json("/api/v1/docs/zh/backend-collectors", user)
admin_response = await get_json("/api/v1/docs/zh/backend-system-service-control", user)
assert developer_response.status_code == 200
assert developer_response.json()["access"] == "docs_developer"
assert admin_response.status_code == 403
@pytest.mark.asyncio
async def test_admin_and_super_admin_can_read_admin_docs():
admin_response = await get_json(
"/api/v1/docs/zh/backend-system-service-control",
make_user(role="admin"),
)
super_admin_response = await get_json(
"/api/v1/docs/zh/backend-system-service-control",
make_user(role="super_admin"),
)
assert admin_response.status_code == 200
assert super_admin_response.status_code == 200
@pytest.mark.asyncio
async def test_unknown_language_slug_and_path_traversal_do_not_read_files():
bad_lang = await get_json("/api/v1/docs/fr/quickstart")
bad_slug = await get_json("/api/v1/docs/zh/not-a-doc")
traversal = await get_json("/api/v1/docs/zh/..%2Fmanual")
assert bad_lang.status_code == 404
assert bad_slug.status_code == 404
assert traversal.status_code == 404

View File

@@ -0,0 +1,957 @@
"""Tests for the shared location resolution pipeline.
Validates the abstraction itself: the protocol contract, the orchestrator,
each built-in resolver, and the pluggability promise (a custom resolver can
be slotted in without touching consumers).
"""
from __future__ import annotations
import json
from pathlib import Path
import pytest
from app.services.location import (
InheritFromAnotherEntityResolver,
LocationCandidate,
LocationPipeline,
LocationQuery,
NominatimResolver,
RegistryResolver,
ResolverOutput,
SourceCoordinatesResolver,
)
from app.schemas.ai import SituationalAnalysisResponse
import app.services.location.llm_fallback as llm_fallback
from app.services.location.llm_fallback import collect_llm_location_fallback_candidate
# ── Test fixtures ────────────────────────────────────────────────────
@pytest.fixture
def tmp_registry(tmp_path: Path) -> Path:
payload = {
"locations": [
{
"canonical_name": "Test Site Alpha",
"aliases": ["alpha", "alpha-one", "Acme HQ"],
"operator": "Acme Networks",
"site": "Acme HQ",
"city": "Lyon",
"country": "France",
"latitude": 45.764,
"longitude": 4.8357,
"precision": "site",
"confidence": 0.92,
"source_note": "Test fixture",
"verified_at": "2026-05-08",
},
{
"canonical_name": "Test Site Bravo",
"aliases": ["bravo"],
"operator": "Acme Networks",
"site": "Bravo POP",
"city": "Berlin",
"country": "Germany",
"latitude": 52.52,
"longitude": 13.405,
"precision": "city",
"confidence": 0.85,
},
],
"city_fallbacks": [
{
"city": "Bhutan-Capital",
"country": "Bhutan",
"latitude": 27.4728,
"longitude": 89.639,
"precision": "city",
"confidence": 0.5,
}
],
}
path = tmp_path / "registry.json"
path.write_text(json.dumps(payload), encoding="utf-8")
return path
# ── SourceCoordinatesResolver ────────────────────────────────────────
def test_source_coordinates_resolver_passes_through_valid_coordinates():
resolver = SourceCoordinatesResolver()
query = LocationQuery(
name="Acme HQ",
source_latitude=45.0,
source_longitude=4.0,
country="France",
)
output = resolver.resolve(query)
assert len(output.candidates) == 1
candidate = output.candidates[0]
assert candidate.latitude == 45.0
assert candidate.longitude == 4.0
assert candidate.precision == "precise"
assert candidate.source == "source_coordinates"
assert candidate.needs_confirmation is False
def test_source_coordinates_resolver_skips_zero_coordinates():
resolver = SourceCoordinatesResolver()
output = resolver.resolve(
LocationQuery(name="X", source_latitude=0.0, source_longitude=0.0)
)
assert output.candidates == ()
def test_source_coordinates_resolver_skips_when_missing():
resolver = SourceCoordinatesResolver()
output = resolver.resolve(LocationQuery(name="X"))
assert output.candidates == ()
# ── RegistryResolver ─────────────────────────────────────────────────
def test_registry_resolver_matches_alias(tmp_registry):
resolver = RegistryResolver(registry_path=tmp_registry)
resolver.reload()
output = resolver.resolve(
LocationQuery(name="alpha", country="France")
)
candidates = list(output.candidates)
assert candidates, "should match registry entry"
assert any(c.matched_location_name == "Test Site Alpha" for c in candidates)
alpha = next(c for c in candidates if c.matched_location_name == "Test Site Alpha")
assert alpha.precision == "site"
assert alpha.confidence == pytest.approx(0.92)
assert alpha.needs_confirmation is True
assert alpha.location_verified_at is None
def test_registry_resolver_filters_country_mismatch(tmp_registry):
resolver = RegistryResolver(registry_path=tmp_registry)
resolver.reload()
# alpha is in France; query says Spain → should reject
output = resolver.resolve(
LocationQuery(name="alpha", country="Spain")
)
assert all(
c.matched_location_name != "Test Site Alpha" for c in output.candidates
)
def test_registry_resolver_emits_city_fallback_candidate(tmp_registry):
resolver = RegistryResolver(registry_path=tmp_registry)
resolver.reload()
output = resolver.resolve(
LocationQuery(city="Bhutan-Capital", country="Bhutan")
)
candidates = list(output.candidates)
assert candidates, "city fallback should fire"
assert any(c.source == "local_registry_city" for c in candidates)
# ── NominatimResolver ───────────────────────────────────────────────
def test_nominatim_resolver_calls_geocoder_with_plan_queries():
calls = []
def fake_geocoder(query: str):
calls.append(query)
return {
"lat": "12.34",
"lon": "56.78",
"display_name": "Test City, Country",
"address": {"city": "Test City", "country": "Country"},
}
def plan(query: LocationQuery):
return [
("primary query", ("name",)),
("secondary query", ("city",)),
]
resolver = NominatimResolver(
query_plan_builder=plan,
geocoder=fake_geocoder,
)
output = resolver.resolve(LocationQuery(name="X", country="Country"))
assert calls == ["primary query", "secondary query"]
assert output.attempted_queries == ("primary query", "secondary query")
assert len(output.candidates) == 2
assert all(c.precision == "city" for c in output.candidates)
assert all(c.needs_confirmation for c in output.candidates)
def test_nominatim_resolver_skips_when_geocoder_returns_none():
resolver = NominatimResolver(
query_plan_builder=lambda q: [("only", ("name",))],
geocoder=lambda q: None,
)
output = resolver.resolve(LocationQuery(name="X"))
assert output.candidates == ()
assert output.attempted_queries == ("only",)
def test_nominatim_resolver_swallows_exceptions_per_query():
def boom(query):
raise RuntimeError("network down")
resolver = NominatimResolver(
query_plan_builder=lambda q: [("a", ()), ("b", ())],
geocoder=boom,
)
output = resolver.resolve(LocationQuery(name="X"))
assert output.candidates == ()
assert output.attempted_queries == ("a", "b")
# ── InheritFromAnotherEntityResolver ────────────────────────────────
def test_inherit_resolver_returns_provided_candidate():
sentinel = LocationCandidate(
latitude=10.0,
longitude=20.0,
display_name="Inherited",
precision="city",
confidence=0.7,
query="inherit::test",
source="inherited",
source_note=None,
matched_fields=("collector",),
needs_confirmation=False,
)
resolver = InheritFromAnotherEntityResolver(
source_lookup=lambda q: sentinel
)
output = resolver.resolve(LocationQuery(name="X"))
assert output.candidates == (sentinel,)
def test_inherit_resolver_skips_when_lookup_returns_none():
resolver = InheritFromAnotherEntityResolver(source_lookup=lambda q: None)
assert resolver.resolve(LocationQuery(name="X")).candidates == ()
# ── LocationPipeline orchestration ──────────────────────────────────
def test_pipeline_aggregates_candidates_across_resolvers(tmp_registry):
pipeline = LocationPipeline(
[
SourceCoordinatesResolver(),
RegistryResolver(registry_path=tmp_registry),
NominatimResolver(
query_plan_builder=lambda q: [("nominatim attempt", ("name",))],
geocoder=lambda q: {
"lat": "1.0",
"lon": "2.0",
"display_name": "Online City",
"address": {"city": "Online City", "country": "France"},
},
),
]
)
pipeline.resolvers[1].reload()
candidates, attempted = pipeline.collect_candidates(
LocationQuery(
name="alpha",
country="France",
source_latitude=44.0,
source_longitude=5.0,
)
)
sources = {c.source for c in candidates}
assert "source_coordinates" in sources
assert "local_registry" in sources
assert "nominatim_online_geocode" in sources
assert "nominatim attempt" in attempted
def test_pipeline_dedupes_by_source_and_coordinates():
same = LocationCandidate(
latitude=1.0,
longitude=2.0,
display_name="dup",
precision="city",
confidence=0.5,
query="x",
source="dup_source",
source_note=None,
matched_fields=(),
needs_confirmation=False,
)
class _DupResolver:
name = "dup_source"
def resolve(self, query):
return ResolverOutput(candidates=(same, same))
pipeline = LocationPipeline([_DupResolver()])
candidates, _ = pipeline.collect_candidates(LocationQuery(name="X"))
assert len(candidates) == 1
def test_registry_short_aliases_do_not_match_inside_larger_tokens(tmp_path: Path):
registry_path = tmp_path / "registry.json"
registry_path.write_text(
json.dumps(
{
"locations": [
{
"canonical_name": "Aurora",
"aliases": ["Aurora", "ANL"],
"site": "DOE/SC/Argonne National Laboratory",
"country": "United States",
"city": "Lemont",
"latitude": 41.713,
"longitude": -87.982,
"precision": "site",
},
{
"canonical_name": "Venado",
"aliases": ["Venado"],
"site": "DOE/NNSA/LANL",
"country": "United States",
"city": "Los Alamos",
"latitude": 35.8443,
"longitude": -106.2872,
"precision": "site",
},
],
"city_fallbacks": [],
}
),
encoding="utf-8",
)
resolver = RegistryResolver(registry_path=registry_path)
resolver.reload()
output = resolver.resolve(
LocationQuery(
name="Venado",
country="United States",
extra={"site": "DOE/NNSA/LANL"},
)
)
assert len(output.candidates) == 1
assert output.candidates[0].matched_location_name == "Venado"
def test_pipeline_resolve_best_returns_highest_priority():
online = LocationCandidate(
latitude=10.0,
longitude=20.0,
display_name="online",
precision="city",
confidence=0.9,
query="x",
source="nominatim_online_geocode",
source_note=None,
matched_fields=(),
needs_confirmation=True,
)
source = LocationCandidate(
latitude=11.0,
longitude=21.0,
display_name="src",
precision="precise",
confidence=1.0,
query="x",
source="source_coordinates",
source_note=None,
matched_fields=(),
needs_confirmation=False,
)
class _StubResolver:
def __init__(self, c, name):
self._c = c
self.name = name
def resolve(self, query):
return ResolverOutput(candidates=(self._c,))
pipeline = LocationPipeline(
[
_StubResolver(online, "online"),
_StubResolver(source, "src"),
]
)
result = pipeline.resolve_best(LocationQuery(name="X"))
assert result.location is source, "source_coordinates should beat nominatim"
def test_pipeline_returns_diagnostic_when_nothing_resolves():
pipeline = LocationPipeline([SourceCoordinatesResolver()])
result = pipeline.resolve_best(LocationQuery(name="X", country="Bhutan"))
assert result.location is None
assert result.diagnostic is not None
assert result.diagnostic.country == "Bhutan"
def test_pluggability_custom_resolver_works_without_changing_pipeline():
"""Validates the abstraction promise: a new algorithm = a new class."""
class _PeeringDBStubResolver:
name = "fake_peeringdb"
def resolve(self, query):
asn = (query.extra or {}).get("asn")
if asn != 174:
return ResolverOutput()
return ResolverOutput(
candidates=(
LocationCandidate(
latitude=1.0,
longitude=2.0,
display_name="Cogent HQ",
precision="site",
confidence=0.8,
query=f"peeringdb::{asn}",
source="peeringdb_stub",
source_note="Stub for testing",
matched_fields=("asn",),
needs_confirmation=False,
),
)
)
pipeline = LocationPipeline([_PeeringDBStubResolver()])
candidates, _ = pipeline.collect_candidates(
LocationQuery(name="X", extra={"asn": 174})
)
assert len(candidates) == 1
assert candidates[0].source == "peeringdb_stub"
# ── LLM fallback helper ─────────────────────────────────────────────
class _FakeAIProviderClient:
def __init__(self, content: str | list[str]):
self.contents = content if isinstance(content, list) else [content]
self.calls = 0
async def analyze(self, payload, request_id=None):
self.calls += 1
content = self.contents[min(self.calls - 1, len(self.contents) - 1)]
return SituationalAnalysisResponse(
provider="test",
model="test-model",
content=content,
raw_response={},
)
@pytest.mark.asyncio
async def test_llm_location_fallback_returns_candidate_from_strict_json():
client = _FakeAIProviderClient(
json.dumps(
{
"latitude": 45.764,
"longitude": 4.8357,
"precision": "city",
"confidence": 0.74,
"city": "Lyon",
"region": "Auvergne-Rhone-Alpes",
"country": "France",
"matched_location_name": "Lyon, France",
"evidence": ["operator and city point to Lyon"],
"reasoning_summary": "Best supported city-level match.",
}
)
)
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(
name="Mystery GPU Cluster",
city="Lyon",
country="France",
extra={"operator": "Mystery Operator"},
),
entity_type="compute_center",
attempted_queries=("Mystery Operator, Lyon, France",),
)
assert client.calls == 1
assert result.failure_reason is None
assert result.attempted_queries == ["llm_factcheck:compute_center:Mystery GPU Cluster"]
candidate = result.candidates[0]
assert candidate.source == "llm_location_factcheck"
assert candidate.needs_confirmation is True
assert candidate.precision == "city"
assert candidate.city == "Lyon"
@pytest.mark.asyncio
async def test_llm_location_fallback_accepts_common_precision_aliases():
client = _FakeAIProviderClient(
json.dumps(
{
"candidate": {
"latitude": 43.2389,
"longitude": 76.8897,
"precision": "city-level",
"confidence": "0.68",
"city": "Almaty",
"country": "Kazakhstan",
"matched_location_name": "Almaty, Kazakhstan",
"evidence": ["NITEC context points to Almaty"],
"reasoning_summary": "City-level fallback.",
}
}
)
)
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(name="Alem.Cloud", country="Kazakhstan"),
entity_type="compute_center",
)
assert result.failure_reason is None
assert result.candidates[0].precision == "city"
assert result.candidates[0].confidence >= 0.55
@pytest.mark.asyncio
async def test_llm_location_fallback_accepts_lat_lng_aliases():
client = _FakeAIProviderClient(
json.dumps(
{
"lat": 51.1694,
"lng": 71.4491,
"precision": "city",
"confidence": 0.62,
"city": "Astana",
"country": "Kazakhstan",
"matched_location_name": "Astana, Kazakhstan",
"evidence": [
{
"source": "Official source",
"source_type": "official",
"entity_match": True,
"text": "Alem.Cloud is in Astana.",
}
],
}
)
)
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(name="Alem.Cloud", country="Kazakhstan"),
entity_type="compute_center",
)
assert result.failure_reason is None
assert result.candidates[0].latitude == pytest.approx(51.1694)
assert result.candidates[0].longitude == pytest.approx(71.4491)
@pytest.mark.asyncio
async def test_llm_location_fallback_geocodes_city_when_coordinates_missing(monkeypatch):
monkeypatch.setattr(
llm_fallback,
"_geocode_llm_city",
lambda query: {
"lat": "51.1694",
"lon": "71.4491",
"display_name": "Astana, Kazakhstan",
"address": {"city": "Astana", "country": "Kazakhstan"},
},
)
client = _FakeAIProviderClient(
json.dumps(
{
"precision": "city",
"confidence": 0.62,
"city": "Astana",
"country": "Kazakhstan",
"matched_location_name": "Astana, Kazakhstan",
"evidence": [
{
"source": "Official source",
"source_type": "official",
"entity_match": True,
"text": "Alem.Cloud is in Astana.",
}
],
}
)
)
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(name="Alem.Cloud", country="Kazakhstan"),
entity_type="compute_center",
)
assert result.failure_reason is None
candidate = result.candidates[0]
assert candidate.latitude == pytest.approx(51.1694)
assert candidate.longitude == pytest.approx(71.4491)
assert "Nominatim city fallback" in candidate.source_note
@pytest.mark.asyncio
async def test_llm_location_fallback_geocodes_matched_location_without_city(monkeypatch):
def _fake_geocode(query):
if "Falun" not in query:
return None
return {
"lat": "60.6065",
"lon": "15.6355",
"display_name": "Falun, Dalarna County, Sweden",
"address": {"city": "Falun", "state": "Dalarna County", "country": "Sweden"},
}
monkeypatch.setattr(llm_fallback, "_geocode_llm_city", _fake_geocode)
client = _FakeAIProviderClient(
json.dumps(
{
"precision": "city",
"confidence": 0.64,
"country": "Sweden",
"matched_location_name": "Falun, Sweden",
"evidence": [
{
"source": "Credible public source",
"source_type": "news",
"entity_match": True,
"text": "DeepL Mercury supercomputer is located in Falun.",
}
],
}
)
)
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(name="DeepL Mercury", country="Sweden"),
entity_type="compute_center",
)
assert result.failure_reason is None
candidate = result.candidates[0]
assert candidate.city == "Falun"
assert candidate.country == "瑞典"
assert candidate.latitude == pytest.approx(60.6065)
assert candidate.longitude == pytest.approx(15.6355)
@pytest.mark.asyncio
async def test_llm_location_fallback_repairs_non_json_answer(monkeypatch):
monkeypatch.setattr(
llm_fallback,
"_geocode_llm_city",
lambda query: {
"lat": "25.033",
"lon": "121.5654",
"display_name": "Taipei, Taiwan",
"address": {"city": "Taipei", "country": "Taiwan"},
},
)
client = _FakeAIProviderClient(
[
"TAIPEI-1 appears to be located in Taipei, Taiwan, based on NVIDIA context.",
json.dumps(
{
"latitude": None,
"longitude": None,
"precision": "city",
"confidence": 0.62,
"city": "Taipei",
"country": "Taiwan",
"matched_location_name": "Taipei, Taiwan",
"evidence": [
{
"source": "NVIDIA context",
"source_type": "generic",
"entity_match": True,
"text": "TAIPEI-1 appears to be located in Taipei.",
}
],
"reasoning_summary": "City-level location extracted from prose.",
}
),
]
)
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(name="TAIPEI-1", country="Taiwan"),
entity_type="compute_center",
)
assert client.calls == 2
assert result.failure_reason is None
assert result.candidates[0].city == "Taipei"
assert result.candidates[0].source == "llm_location_factcheck"
@pytest.mark.asyncio
async def test_llm_location_fallback_accepts_taipei_name_hint_with_weak_wording(monkeypatch):
monkeypatch.setattr(
llm_fallback,
"_geocode_llm_city",
lambda query: {
"lat": "25.033",
"lon": "121.5654",
"display_name": "Taipei, Taiwan",
"address": {"city": "Taipei", "country": "Taiwan"},
},
)
client = _FakeAIProviderClient(
json.dumps(
{
"latitude": None,
"longitude": None,
"precision": "city",
"confidence": 0.43,
"city": "Taipei",
"country": "Taiwan",
"matched_location_name": "Taipei, Taiwan",
"evidence": [
{
"source": "NVIDIA context",
"source_type": "generic",
"entity_match": True,
"text": "TAIPEI-1 points to Taipei city-level placement.",
}
],
"reasoning_summary": "Weak city-level evidence, but the entity name and geography align.",
}
)
)
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(name="TAIPEI-1", country="Taiwan"),
entity_type="compute_center",
)
assert result.failure_reason is None
candidate = result.candidates[0]
assert candidate.city == "Taipei"
assert candidate.confidence >= 0.55
breakdown = candidate.suggested_registry_entry["llm_score_breakdown"]
assert breakdown["weak_evidence_penalty"] <= 0.15
assert breakdown["conflict_penalty"] == 0
assert breakdown["name_location_hint"] > 0
@pytest.mark.asyncio
async def test_llm_location_fallback_geocodes_city_from_entity_name_when_llm_unparseable(monkeypatch):
def _fake_geocode(query):
if query != "Taipei, 中国(台湾)":
return None
return {
"lat": "25.033",
"lon": "121.5654",
"display_name": "Taipei, Taiwan",
"address": {"city": "Taipei", "country": "Taiwan"},
}
monkeypatch.setattr(llm_fallback, "_geocode_llm_city", _fake_geocode)
client = _FakeAIProviderClient(["not a location answer", "still not json"])
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(name="TAIPEI-1", country="中国(台湾)"),
entity_type="compute_center",
)
assert client.calls == 2
assert result.failure_reason is None
candidate = result.candidates[0]
assert candidate.city == "Taipei"
assert candidate.latitude == pytest.approx(25.033)
assert candidate.longitude == pytest.approx(121.5654)
assert "Entity name city hint" in candidate.source_note
@pytest.mark.asyncio
async def test_llm_location_fallback_extracts_city_from_non_json_when_repair_fails(monkeypatch):
monkeypatch.setattr(
llm_fallback,
"_geocode_llm_city",
lambda query: {
"lat": "60.6065",
"lon": "15.6355",
"display_name": "Falun, Sweden",
"address": {"city": "Falun", "country": "Sweden"},
},
)
client = _FakeAIProviderClient(
[
"DeepL Mercury 超級電腦位於瑞典的 法倫 (Falun)。",
"still not json",
]
)
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(name="DeepL Mercury", country="Sweden"),
entity_type="compute_center",
)
assert client.calls == 2
assert result.failure_reason is None
assert result.candidates[0].city == "Falun"
assert result.candidates[0].needs_confirmation is True
@pytest.mark.asyncio
async def test_llm_location_fallback_combines_model_score_with_evidence_score():
client = _FakeAIProviderClient(
json.dumps(
{
"latitude": 51.1694,
"longitude": 71.4491,
"precision": "city",
"confidence": 0.38,
"city": "Astana",
"country": "Kazakhstan",
"matched_location_name": "Astana, Kazakhstan",
"evidence": [
{
"source": "Kazakhstan National Supercomputing Center",
"url": "https://example.test/alem-cloud",
"source_type": "official",
"entity_match": True,
"text": "Alem.Cloud is located in Astana.",
}
],
"reasoning_summary": "Evidence supports city-level location but not exact facility coordinates.",
}
)
)
result = await collect_llm_location_fallback_candidate(
provider_client=client,
query=LocationQuery(name="Alem.Cloud", country="Kazakhstan"),
entity_type="compute_center",
)
assert result.failure_reason is None
candidate = result.candidates[0]
assert candidate.city == "Astana"
assert candidate.confidence >= 0.55
assert candidate.suggested_registry_entry["llm_model_confidence"] == pytest.approx(0.38)
assert candidate.suggested_registry_entry["llm_combined_confidence"] == pytest.approx(
candidate.confidence
)
@pytest.mark.asyncio
async def test_llm_location_fallback_rejects_low_combined_score():
result = await collect_llm_location_fallback_candidate(
provider_client=_FakeAIProviderClient(
json.dumps(
{
"latitude": 51.1694,
"longitude": 71.4491,
"precision": "city",
"confidence": 0.38,
"city": "Astana",
"country": "Kazakhstan",
"matched_location_name": "Astana, Kazakhstan",
"evidence": ["some page mentions Kazakhstan"],
"reasoning_summary": "Weak and ambiguous city evidence.",
"ambiguity": "weak city evidence",
}
)
),
query=LocationQuery(name="Alem.Cloud", country="Kazakhstan"),
entity_type="compute_center",
)
assert result.candidates == []
assert "combined evidence score" in result.failure_reason
assert "below minimum 0.55" in result.failure_reason
@pytest.mark.asyncio
async def test_llm_location_fallback_rejects_explicit_conflicts():
result = await collect_llm_location_fallback_candidate(
provider_client=_FakeAIProviderClient(
json.dumps(
{
"latitude": 25.033,
"longitude": 121.5654,
"precision": "city",
"confidence": 0.70,
"city": "Taipei",
"country": "Taiwan",
"matched_location_name": "Taipei, Taiwan",
"evidence": [
{
"source": "Conflicting source",
"source_type": "generic",
"entity_match": True,
"has_conflict": True,
"text": "One source says Taipei, another contradicts it.",
}
],
"reasoning_summary": "Conflicting evidence prevents confirmation.",
}
)
),
query=LocationQuery(name="TAIPEI-1", country="Taiwan"),
entity_type="compute_center",
)
assert result.candidates == []
assert "conflict=" in result.failure_reason
@pytest.mark.asyncio
@pytest.mark.parametrize(
"content",
[
"not json",
json.dumps({"latitude": 0, "longitude": 0, "precision": "city", "confidence": 0.9}),
json.dumps({"latitude": 45, "longitude": 4, "precision": "country", "confidence": 0.9}),
json.dumps({"latitude": 45, "longitude": 4, "precision": "city", "confidence": 0.2}),
],
)
async def test_llm_location_fallback_rejects_unsafe_outputs(content):
result = await collect_llm_location_fallback_candidate(
provider_client=_FakeAIProviderClient(content),
query=LocationQuery(name="Unsafe", country="France"),
entity_type="compute_center",
)
assert result.candidates == []
assert result.failure_reason
assert result.attempted_queries == ["llm_factcheck:compute_center:Unsafe"]
@pytest.mark.asyncio
async def test_llm_location_fallback_failure_explains_rejection_reason():
result = await collect_llm_location_fallback_candidate(
provider_client=_FakeAIProviderClient(
json.dumps({
"latitude": 45,
"longitude": 4,
"precision": "region",
"confidence": 0.9,
})
),
query=LocationQuery(name="Unsafe", country="France"),
entity_type="compute_center",
)
assert result.candidates == []
assert "precision" in result.failure_reason
assert "region" in result.failure_reason

View File

@@ -121,6 +121,25 @@ class TestCollectionTaskModel:
)
assert task.records_processed == 100
def test_task_with_phase_progress(self):
"""Test collection task phase-level progress fields"""
task = CollectionTask(
datasource_id=1,
status="running",
phase="fetching",
phase_progress=42.5,
phase_message="Downloading dataset",
phase_current=1024,
phase_total=4096,
phase_unit="bytes",
)
assert task.phase == "fetching"
assert task.phase_progress == 42.5
assert task.phase_message == "Downloading dataset"
assert task.phase_current == 1024
assert task.phase_total == 4096
assert task.phase_unit == "bytes"
def test_task_error_message(self):
"""Test collection task with error message"""
task = CollectionTask(

View File

@@ -0,0 +1,242 @@
import json
import pytest
from motion_agent.cameras import (
MotionAgentCameraError,
MotionAgentDependencyError,
UrlCameraInput,
UrlCameraSpec,
UsbCameraInput,
UsbCameraSpec,
)
import motion_agent.cameras as motion_cameras
from motion_agent.config import MotionAgentConfig
from motion_agent.events import GestureEvent, HeartbeatEvent, SkeletonEvent, SkeletonJoint
from motion_agent.recognizer import GestureObservation
from motion_agent.server import MotionAgentServer
from motion_agent.state import GestureStateMachine
from motion_agent import cli as motion_cli
def test_gesture_event_serializes_stable_protocol_fields():
event = GestureEvent(
gesture="rotate_left",
confidence=0.91,
intensity=0.75,
timestamp_ms=1000,
seq=7,
mode="single",
)
payload = json.loads(event.to_json())
assert payload["type"] == "gesture"
assert payload["gesture"] == "rotate_left"
assert payload["phase"] == "discrete"
assert payload["confidence"] == 0.91
assert payload["intensity"] == 0.75
assert payload["timestamp_ms"] == 1000
assert payload["seq"] == 7
assert payload["source"] == "motion-agent"
assert payload["mode"] == "single"
assert payload["payload"] == {}
def test_state_machine_ignores_low_confidence_observations():
state = GestureStateMachine(confidence_threshold=0.8, cooldown_ms=400)
event = state.accept(
GestureObservation(
gesture="confirm",
confidence=0.79,
intensity=1,
timestamp_ms=1000,
)
)
assert event is None
def test_state_machine_applies_per_gesture_cooldown():
state = GestureStateMachine(confidence_threshold=0.7, cooldown_ms=400)
first = state.accept(
GestureObservation("rotate_right", confidence=0.9, intensity=0.8, timestamp_ms=1000)
)
repeated = state.accept(
GestureObservation("rotate_right", confidence=0.95, intensity=0.9, timestamp_ms=1200)
)
later = state.accept(
GestureObservation("rotate_right", confidence=0.95, intensity=0.9, timestamp_ms=1500)
)
assert first is not None
assert first.seq == 1
assert repeated is None
assert later is not None
assert later.seq == 2
def test_motion_server_status_includes_dry_run_camera_and_heartbeat():
server = MotionAgentServer(MotionAgentConfig(dry_run=True))
status = json.loads(server.status_event().to_json())
heartbeat = json.loads(HeartbeatEvent(timestamp_ms=123).to_json())
assert status["type"] == "status"
assert status["camera_count"] == 1
assert status["active_camera_ids"] == ["dry-run:null-camera"]
assert status["recognizer"] == "dry-run"
assert heartbeat == {
"timestamp_ms": 123,
"source": "motion-agent",
"type": "heartbeat",
}
def test_skeleton_event_serializes_without_raw_image_fields():
event = SkeletonEvent(
joints=[SkeletonJoint("left_wrist", 0.42, 0.61, 0.98)],
bones=[("left_shoulder", "left_elbow"), ("left_elbow", "left_wrist")],
matched_gesture="rotate_left",
confidence=0.91,
camera_id="usb:0",
timestamp_ms=1000,
mode="single",
)
payload = json.loads(event.to_json())
assert payload["type"] == "skeleton"
assert payload["matched_gesture"] == "rotate_left"
assert payload["confidence"] == 0.91
assert payload["camera_id"] == "usb:0"
assert payload["joints"] == [
{"id": "left_wrist", "x": 0.42, "y": 0.61, "confidence": 0.98}
]
assert payload["bones"] == [["left_shoulder", "left_elbow"], ["left_elbow", "left_wrist"]]
assert "image" not in payload
assert "frame" not in payload
def test_dry_run_recognizer_produces_debug_skeleton():
server = MotionAgentServer(MotionAgentConfig(dry_run=True))
skeleton = server.recognizer.debug_skeleton(
None,
camera_id="dry-run:null-camera",
mode="single",
)
assert skeleton is not None
assert skeleton.type == "skeleton"
assert skeleton.camera_id == "dry-run:null-camera"
assert skeleton.joints
assert skeleton.bones
class ServerRecognizerStub:
name = "stub"
def recognize(self, frame):
_ = frame
return None
def debug_skeleton(self, frame, **kwargs):
_ = frame, kwargs
return None
def test_motion_server_prefers_camera_urls_over_usb_indexes():
server = MotionAgentServer(
MotionAgentConfig(
dry_run=False,
camera_indexes=(0,),
camera_urls=("rtsp://camera.example/live", "http://camera.example/video"),
),
recognizer=ServerRecognizerStub(),
)
assert [camera.camera_id for camera in server.cameras] == ["url:0", "url:1"]
assert all(isinstance(camera, UrlCameraInput) for camera in server.cameras)
def test_usb_camera_reports_missing_opencv_as_readable_dependency_error(monkeypatch):
import builtins
original_import = builtins.__import__
original_exists = motion_cameras.Path.exists
def fake_import(name, *args, **kwargs):
if name == "cv2":
raise ImportError("cv2 missing")
return original_import(name, *args, **kwargs)
monkeypatch.setattr(builtins, "__import__", fake_import)
monkeypatch.setattr(
motion_cameras.Path,
"exists",
lambda self: True if str(self) in {"/dev", "/dev/video0"} else original_exists(self),
)
camera = UsbCameraInput(UsbCameraSpec(index=0))
with pytest.raises(MotionAgentDependencyError, match="Add opencv-python with uv"):
camera.open()
def test_usb_camera_reports_missing_device_before_opencv_noise(monkeypatch):
original_exists = motion_cameras.Path.exists
monkeypatch.setattr(
motion_cameras.Path,
"exists",
lambda self: True if str(self) == "/dev" else False if str(self) == "/dev/video0" else original_exists(self),
)
camera = UsbCameraInput(UsbCameraSpec(index=0))
with pytest.raises(MotionAgentCameraError, match="/dev/video0"):
camera.open()
def test_url_camera_reports_unreachable_stream(monkeypatch):
class BrokenCapture:
def __init__(self, _url):
pass
def isOpened(self):
return False
class Cv2Stub:
VideoCapture = BrokenCapture
import builtins
original_import = builtins.__import__
def fake_import(name, *args, **kwargs):
if name == "cv2":
return Cv2Stub()
return original_import(name, *args, **kwargs)
monkeypatch.setattr(builtins, "__import__", fake_import)
camera = UrlCameraInput(UrlCameraSpec(url="rtsp://camera.example/live"))
with pytest.raises(MotionAgentCameraError, match="Unable to open camera URL"):
camera.open()
@pytest.mark.asyncio
async def test_motion_agent_cli_reports_dependency_error_without_traceback(monkeypatch, capsys):
class BrokenServer:
def __init__(self, _config):
raise MotionAgentDependencyError("missing cv stack")
monkeypatch.setattr(motion_cli, "MotionAgentServer", BrokenServer)
exit_code = await motion_cli.async_main([])
captured = capsys.readouterr()
assert exit_code == 2
assert "Motion agent failed: missing cv stack" in captured.err
assert "Traceback" not in captured.err

Some files were not shown because too many files have changed in this diff Show More