Compare commits

...

15 Commits

Author SHA1 Message Date
rayd1o
8955c58d19 release: bump version to 0.51.1 2026-05-11 10:59:41 +08:00
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
237 changed files with 39150 additions and 4321 deletions

View File

@@ -1,151 +1,93 @@
---
description: 分析本次 git 变更,在 docs/technical/zh/ 中新建或更新对应的技术文档
argument-hint: 可选:指定要记录的主题,或留空自动从 git diff 推断
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 — 技术文档写入工作流
# /docs — Documentation Workflow
## 目标
## Goal
根据当前 git 变更(或用户指定主题)在 `docs/technical/zh/` 中写入或更新技术文档,记录**为什么**这样做,而不只是记录做了什么。
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
### Step 1 — 理解变更范围
Before deciding scope, check whether the repository has a documentation rules file:
```bash
git diff HEAD --stat # 变更文件一览
git diff HEAD --name-only # 变更文件列表
git log --oneline -10 # 近期 commit 上下文
test -f docs/documentation-coverage-rules.md && sed -n '1,240p' docs/documentation-coverage-rules.md
```
`$ARGUMENTS` 指定了主题,优先聚焦该主题;否则从文件列表和 diff stat 推断变更主题。不要默认读取完整仓库 diff只对决定文档主题所需的文件读取 focused diff
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 — 确认文档范围
### 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.
1. **应写几篇文档**:单一主题写一篇,跨领域变更可拆分(如后端性能优化 + 运维启动脚本分开写)
2. **是新建还是更新**:检查 `docs/technical/zh/` 中是否已有相关文档
3. **文档命名**:按 `领域-主题-副题.md` 格式,全小写,用连字符,如:
- `backend-datasources-api-performance.md`
- `ops-planet-sh-startup.md`
- `earth-bgp-context.md`
For ambiguous or large documentation changes, briefly state the intended doc plan before editing. For clear small changes, proceed directly.
```bash
ls docs/technical/zh/ # 查看现有文档
### 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
```
文档计划:
新建docs/technical/zh/ops-planet-sh-startup.md — planet.sh 启动性能优化
更新docs/technical/zh/backend-datasources-api-performance.md — 补充并行化细节
```
### Step 3 — 写文档
遵循以下原则:
**记录 WHY不只记录 WHAT**
- 好:`将戳文件从 /tmp 移到 ~/.cache/planet/,因为 WSL 重启后 /tmp 被清空`
- 差:`修改了 AI_PROVIDER_BUILD_STAMP_FILE 的值`
**必须包含的内容**
- 背景/问题:改动之前存在什么问题,为什么要改
- 核心设计决策及其理由
- 关键代码片段(用 diff 或 before/after 展示)
- 相关文件列表
**格式要求**
- 使用 `##``###` 分级,不要超过三级
- 代码块注明语言python / bash / typescript / sql
- 表格用于对比多个选项或列出参数
- 中文写作,技术术语保留英文原文
- `docs/technical/zh/` 中的文档不得用英文原文占位;如果存在 `docs/technical/en/` 对应文件,禁止逐字复制成中文文件
- 中文文档内部链接应指向 `docs/technical/zh/...`,除非明确引用英文专属文档
**文档结构模板**
```markdown
# 标题(说明做了什么)
## 背景
为什么要做这个改动,改动前存在什么问题。
## 核心变更
### 子主题一
before/after 或决策说明 + 关键代码
### 子主题二
...
## 相关文件
- `path/to/file.py` — 简短说明
```
### Step 4 — 验证
- 读一遍写好的文档,确认逻辑清晰、代码片段无明显错误
-`rg --files``test -e` 确认文档中的文件路径在项目中真实存在,避免凭记忆判断:
- 检查中文文档没有误复制英文版:
```bash
python - <<'PY'
from pathlib import Path
same = []
for en in sorted(Path("docs/technical/en").glob("*.md")):
zh = Path("docs/technical/zh") / en.name
if zh.exists() and en.read_text() == zh.read_text():
same.append(en.name)
if same:
raise SystemExit("identical en/zh docs: " + ", ".join(same))
print("no identical en/zh docs")
PY
```
- 检查中文文档内部链接没有继续指向无语言目录:
```bash
rg -n "/home/ray/dev/linkong/planet/docs/technical/(?!zh|en)" docs/technical/zh --pcre2
```
```bash
# 对文档中提到的关键路径做快速验证
ls <mentioned_paths>
```
如需检查大量链接,优先用确定性提取:
```bash
rg -n "\]\(([^)]+)\)" docs/technical/zh/<doc>.md
```
### Step 5 — 完成确认
输出摘要:
```
✓ 新建docs/technical/zh/ops-planet-sh-startup.md约 xxx 字)
✓ 更新docs/technical/zh/backend-datasources-api-performance.md
```
## 注意事项
- 不要写流水账式的"改了 A、改了 B、改了 C",要写改动背后的约束和权衡
- 不要在文档中引用 PR 号、issue 号、或当前对话——这些会随时间失效
- 代码片段保持简洁,只保留说明问题的关键部分,省略无关样板代码
- 如果某个变更已有文档记录,优先在原文档中追加,而不是新建
- 文档是给未来的开发者看的,假设读者熟悉项目但不了解这次改动的背景
- 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

@@ -1,107 +1,72 @@
---
name: docs
description: Analyze current Planet repo changes and create or update technical documentation under docs/technical/zh. Use when the user asks to write docs, update technical docs, summarize implementation changes into documentation, or port the Claude docs-codex workflow into Codex.
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 user asks to create or update Planet technical documentation, especially under `docs/technical/zh/`.
Use this skill when the task is documentation work: creating, updating, checking, or summarizing docs for code or behavior changes.
## Goal
Write or update technical docs that explain why a change exists, not only what files changed.
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.
Default target directory:
## Repository Rules
- `docs/technical/zh/`
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 change context:
1. Gather focused context:
```bash
git diff HEAD --stat
git diff HEAD --name-only
git log --oneline -10
ls docs/technical/zh/
rg --files docs
```
If the user gives a specific topic, focus on that topic. Otherwise infer the documentation topic from the file list and diff stat. Do **not** read the full repository diff by default; inspect focused diffs only for the files that define the doc topic:
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 document scope:
2. Decide scope:
- Use one document for one coherent topic.
- Split documents when the changes cross meaningful domains, such as backend performance and ops startup behavior.
- Prefer updating an existing relevant doc over creating a duplicate.
- Name new files as lowercase hyphenated `domain-topic-detail.md`, for example:
- `backend-datasources-api-performance.md`
- `ops-planet-sh-startup.md`
- `earth-bgp-context.md`
- Use one document for one coherent topic.
- Split documents only when changes cross meaningful domains.
- Keep filenames lowercase and hyphenated.
3. Write the doc in Chinese:
3. Write the doc:
- Write Chinese prose for `docs/technical/zh/`.
- Keep technical identifiers, API paths, config keys, code symbols, and standard product names in English where appropriate.
- Use `##` and `###` headings; avoid going deeper than three levels.
- Use fenced code blocks with language tags.
- Use tables when comparing options or listing parameters.
- 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. Required content:
4. Verify:
- Background/problem: what was wrong before and why the change was needed.
- Core design decisions and rationale.
- Key code snippets, preferably before/after or focused excerpts.
- Related files and what each file contributes.
5. Verification:
- Read the completed doc and check that the reasoning is clear.
- Verify important referenced paths exist.
- Use `rg --files` or `test -e` for path existence instead of relying on memory.
- Run a quick duplicate-language check when editing bilingual docs:
```bash
python - <<'PY'
from pathlib import Path
same = []
for en in sorted(Path("docs/technical/en").glob("*.md")):
zh = Path("docs/technical/zh") / en.name
if zh.exists() and en.read_text() == zh.read_text():
same.append(en.name)
if same:
raise SystemExit("identical en/zh docs: " + ", ".join(same))
print("no identical en/zh docs")
PY
```
Also check that Chinese docs do not link to the old language-less technical docs path:
```bash
rg -n "/home/ray/dev/linkong/planet/docs/technical/(?!zh|en)" docs/technical/zh --pcre2
```
This command should return no matches.
If checking many links, prefer deterministic extraction:
```bash
rg -n "\]\(([^)]+)\)" docs/technical/zh/<doc>.md
```
- 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
- A file under `docs/technical/zh/` must not be an English source file copied as a placeholder.
- Do not leave a Chinese doc with only an English title and English first-screen content.
- When an English counterpart exists in `docs/technical/en/`, never duplicate it byte-for-byte into `docs/technical/zh/`.
- Internal links inside `docs/technical/zh/` should point to `docs/technical/zh/...` for Chinese docs, unless intentionally linking to an English-only file.
- Do not reference PR numbers, issue numbers, or the current conversation.
- Do not write changelog-style lists like "changed A, changed B, changed C" without the constraints and tradeoffs behind those changes.
- Keep code snippets concise and relevant.
- 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
@@ -109,9 +74,9 @@ After editing, summarize:
```md
Updated:
- docs/technical/zh/example.md — what changed
- path/to/doc.md — what changed
Verified:
- no identical en/zh docs
- no language-less docs/technical links in zh docs
- checks that passed
- checks that could not be run, if any
```

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.43.0
0.51.1

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

@@ -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

@@ -5,17 +5,20 @@ from datetime import datetime
import base64
import json
import re
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy import select, func
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
@@ -25,10 +28,25 @@ from app.services.datasource_mapping import (
MappingError,
build_heuristic_mapping,
execute_mapping,
persist_mapped_records,
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()
@@ -36,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={})
@@ -73,6 +91,32 @@ 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
@@ -186,6 +230,8 @@ def _build_query_params(auth_type: str, auth_config: dict, config: dict) -> dict
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"}:
@@ -285,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)
@@ -312,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,
@@ -356,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)
@@ -388,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()
@@ -405,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),
):
@@ -415,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")
@@ -437,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,
@@ -467,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,
@@ -490,41 +649,62 @@ 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(
{
"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}",
}
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 {"total": len(result), "data": result}
return {
**result,
"connected": False,
}
@router.post("/custom/sample")
@@ -771,6 +951,8 @@ async def update_datasource_mapping(
@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),
):
@@ -779,20 +961,24 @@ async def run_mapped_datasource(
if not datasource:
raise HTTPException(status_code=404, detail="Configuration not found")
result = await db.execute(
select(DataSourceMappingTemplate)
.where(DataSourceMappingTemplate.datasource_config_id == config_id)
.where(DataSourceMappingTemplate.is_active.is_(True))
.order_by(DataSourceMappingTemplate.version.desc())
.limit(1)
)
mapping = result.scalar_one_or_none()
if not mapping:
raise HTTPException(status_code=404, detail="No active mapping template found")
try:
sample = await fetch_custom_sample_from_config(datasource, 5_000_000)
mapped = execute_mapping(sample, mapping.mapping_json, mapping.target_schema)
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,
@@ -800,36 +986,26 @@ async def run_mapped_datasource(
) from exc
except httpx.HTTPError as exc:
raise HTTPException(status_code=502, detail=f"Datasource request failed: {exc}") from exc
except (MappingError, ValueError) as exc:
except (CustomDatasourceRuntimeError, MappingError, ValueError) as exc:
raise HTTPException(status_code=400, detail=f"Mapping failed: {exc}") from exc
if mapped["failed_count"] > 0:
return {
"status": "failed",
"datasource_config_id": config_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],
}
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,
)
@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": "success",
"status": "stopped" if stopped else "not_running",
"datasource_config_id": config_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,
"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

@@ -398,6 +398,11 @@ async def list_datasources(
"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,
}
@@ -626,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,
@@ -709,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 = {
@@ -32,6 +31,7 @@ COLLECTOR_URL_KEYS = {
"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",
}
@@ -74,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

@@ -98,3 +98,7 @@ news_live_streams:
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

@@ -245,6 +245,18 @@ DEFAULT_DATASOURCES = {
"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",
},
}
ID_TO_COLLECTOR = {info["id"]: name for name, info in DEFAULT_DATASOURCES.items()}

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

@@ -16,8 +16,11 @@ class VesselAISRecord(BaseModel):
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
@@ -104,8 +107,11 @@ TARGET_SCHEMAS: dict[str, TargetSchema] = {
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("vessel_type", "string", False, "船型", "cargo"),
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"),
),
),

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,14 +103,17 @@ 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(
@@ -127,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(
"""
@@ -146,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)
"""
)
)
@@ -158,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(
"""
@@ -178,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,13 +6,15 @@ 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 VesselPosition, VesselStatic
from app.models.vessel import AISConflictRecord, AISRawObservation, AISSourceHealth, VesselPosition, VesselStatic
from app.models.datasource_mapping import DataSourceMappingTemplate
__all__ = [
@@ -27,11 +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

@@ -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

@@ -1,6 +1,6 @@
"""Vessel AIS models for live maritime tracking."""
from sqlalchemy import BigInteger, Column, DateTime, Float, Index, Integer, SmallInteger, String
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
@@ -73,3 +73,114 @@ class VesselPosition(Base):
"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

@@ -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,7 @@ 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())
@@ -65,3 +66,40 @@ 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

@@ -1,28 +1,26 @@
"""BarentsWatch AIS collector for vessel tracking."""
from datetime import UTC, datetime, timedelta
import os
from datetime import UTC, datetime
from typing import Any
import httpx
from sqlalchemy import delete, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.vessel import VesselPosition, VesselStatic
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
BARENTSWATCH_LATEST_URL = "https://live.ais.barentswatch.no/v1/latest/combined"
BARENTSWATCH_TOKEN_URL = "https://id.barentswatch.no/connect/token"
VESSEL_TYPE_NAMES = {
30: "Fishing",
35: "Military",
60: "Passenger",
70: "Cargo",
80: "Tanker",
}
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):
@@ -38,62 +36,9 @@ class VesselAISCollector(BaseCollector):
def base_url(self) -> str:
return self._resolved_url or BARENTSWATCH_LATEST_URL
async def _load_datasource_config(self) -> dict[str, Any]:
if not self._db_session:
return {}
try:
from sqlalchemy import select
from app.models.datasource_config import DataSourceConfig
result = await self._db_session.execute(
select(DataSourceConfig)
.where(DataSourceConfig.name == self.name)
.where(DataSourceConfig.is_active.is_(True))
)
datasource_config = result.scalar_one_or_none()
except Exception:
return {}
if not datasource_config:
return {}
return {
"auth_config": datasource_config.auth_config or {},
"config": datasource_config.config or {},
}
async def _get_access_token(self, client: httpx.AsyncClient) -> str | None:
datasource_config = await self._load_datasource_config()
auth_config = datasource_config.get("auth_config") or {}
config = datasource_config.get("config") or {}
client_id = (
auth_config.get("client_id")
or config.get("client_id")
or os.getenv("BARENTSWATCH_CLIENT_ID")
or os.getenv("BARRENTSWATCH_CLIENT_ID")
)
client_secret = (
auth_config.get("client_secret")
or config.get("client_secret")
or os.getenv("BARENTSWATCH_CLIENT_SECRET")
or os.getenv("BARRENTSWATCH_CLIENT_SECRET")
)
if not client_id or not client_secret:
return None
response = await client.post(
BARENTSWATCH_TOKEN_URL,
data={
"client_id": client_id,
"client_secret": 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
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:
@@ -145,51 +90,75 @@ class VesselAISCollector(BaseCollector):
records_added = 0
for index, item in enumerate(data):
static = await db.get(VesselStatic, item["mmsi"])
if static is None:
static = VesselStatic(mmsi=item["mmsi"])
db.add(static)
for field in (
"name",
"callsign",
"vessel_type",
"vessel_type_name",
"flag",
"length",
"width",
"draught",
"imo",
):
value = item.get(field)
if value not in (None, ""):
setattr(static, field, value)
static.updated_at = now
db.add(
VesselPosition(
mmsi=item["mmsi"],
lat=item["lat"],
lon=item["lon"],
sog=item.get("sog"),
cog=item.get("cog"),
heading=item.get("heading"),
nav_status=item.get("nav_status"),
received_at=item.get("received_at") or now,
)
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)
await db.execute(
delete(VesselPosition).where(VesselPosition.received_at < now - timedelta(hours=24))
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"))
@@ -209,7 +178,7 @@ class VesselAISCollector(BaseCollector):
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 _vessel_type_name(vessel_type)
or normalize_vessel_type_name(vessel_type)
)
received_at = _parse_datetime(_pick(item, "received_at", "timestamp", "time", "msgtime"))
@@ -308,19 +277,3 @@ def _parse_datetime(value: Any) -> datetime | None:
except ValueError:
return None
return None
def _vessel_type_name(vessel_type: int | None) -> str:
if vessel_type is None:
return "Other"
if 70 <= vessel_type <= 79:
return "Cargo"
if 80 <= vessel_type <= 89:
return "Tanker"
if 60 <= vessel_type <= 69:
return "Passenger"
if vessel_type == 30:
return "Fishing"
if vessel_type == 35:
return "Military"
return VESSEL_TYPE_NAMES.get(vessel_type, "Other")

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

@@ -290,25 +290,77 @@ async def persist_mapped_records(
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.models.vessel import VesselPosition
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:
db.add(
VesselPosition(
mmsi=record["mmsi"],
lat=record["lat"],
lon=record["lon"],
sog=record.get("sog"),
cog=record.get("cog"),
heading=record.get("heading"),
received_at=_parse_datetime(record.get("received_at")) or datetime.now(UTC),
)
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()
return len(records)
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

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,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

@@ -1,13 +1,18 @@
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
@@ -106,6 +111,130 @@ async def test_persist_mapped_records_writes_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():

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

View File

@@ -0,0 +1,169 @@
from types import SimpleNamespace
import pytest
from app.api.v1 import settings as settings_api
from app.api.v1.settings import (
AIProviderIntegrationUpdate,
_build_ai_provider_payload,
_mask_secret,
_normalize_ai_provider_payload,
_resolve_provider_api_key,
get_runtime_ai_provider_config,
)
@pytest.fixture(autouse=True)
def isolated_ai_provider_env_file(monkeypatch, tmp_path):
env_file = tmp_path / ".env"
monkeypatch.setattr(settings_api, "AI_PROVIDER_ENV_FILE", env_file)
return env_file
def test_legacy_ai_provider_payload_maps_to_provider_config():
payload = _normalize_ai_provider_payload(
{
"provider": "openai",
"provider_api": "openai-completions",
"base_url": "https://api.openai.example/v1",
"model": "gpt-test",
"api_key": "old-openai-key",
"max_tokens": 2048,
"anthropic_version": "2023-06-01",
}
)
assert payload["default_provider"] == "openai"
assert payload["providers"]["openai"]["api_key"] == "old-openai-key"
assert payload["providers"]["openai"]["model"] == "gpt-test"
assert payload["providers"]["openai"]["base_url"] == "https://api.openai.example/v1"
def test_provider_key_prefers_specific_env_file_key(isolated_ai_provider_env_file):
isolated_ai_provider_env_file.write_text(
"OPENAI_API_KEY=openai-env-file-key\nAI_API_KEY=generic-env-file-key\n",
encoding="utf-8",
)
value, source = _resolve_provider_api_key("openai", {"api_key": ""})
assert value == "openai-env-file-key"
assert source == "env_file"
def test_provider_key_falls_back_to_generic_ai_api_key(isolated_ai_provider_env_file):
isolated_ai_provider_env_file.write_text(
"AI_API_KEY=generic-env-file-key\n",
encoding="utf-8",
)
value, source = _resolve_provider_api_key("openai", {"api_key": ""})
assert value == "generic-env-file-key"
assert source == "env_file"
def test_mask_secret_without_prefix_is_fully_masked():
assert _mask_secret("plainsecret")["preview"] == "***********"
assert _mask_secret("sk-prefixed")["preview"] == "sk-********"
def test_build_payload_updates_only_selected_provider_key():
current = {
"ai_provider": {
"default_provider": "openai",
"providers": {
"openai": {
"provider": "openai",
"provider_api": "openai-completions",
"base_url": "https://api.openai.com/v1",
"model": "gpt-old",
"api_key": "openai-old-key",
"max_tokens": 4096,
"anthropic_version": "2023-06-01",
},
"minimax": {
"provider": "minimax",
"api_key": "minimax-old-key",
},
},
}
}
update = AIProviderIntegrationUpdate(
provider="openai",
provider_api="openai-completions",
base_url="https://api.openai.com/v1",
model="gpt-new",
api_key="openai-new-key",
max_tokens=8192,
)
payload = _build_ai_provider_payload(current, update)
assert payload["default_provider"] == "openai"
assert payload["providers"]["openai"]["api_key"] == "openai-new-key"
assert payload["providers"]["openai"]["model"] == "gpt-new"
assert payload["providers"]["minimax"]["api_key"] == "minimax-old-key"
def test_build_payload_keeps_saved_key_when_preview_submitted():
current = {
"ai_provider": {
"providers": {
"openai": {
"provider": "openai",
"api_key": "sk-old-secret",
},
},
}
}
update = AIProviderIntegrationUpdate(
provider="openai",
provider_api="openai-completions",
base_url="https://api.openai.com/v1",
model="gpt-test",
api_key="sk-*********",
)
payload = _build_ai_provider_payload(current, update)
assert payload["providers"]["openai"]["api_key"] == "sk-old-secret"
@pytest.mark.asyncio
async def test_runtime_config_uses_default_provider_specific_key(monkeypatch):
record = SimpleNamespace(
payload={
"ai_provider": {
"default_provider": "minimax",
"providers": {
"openai": {
"provider": "openai",
"api_key": "openai-key",
"provider_api": "openai-completions",
"base_url": "https://api.openai.com/v1",
"model": "gpt-test",
},
"minimax": {
"provider": "minimax",
"api_key": "minimax-key",
"provider_api": "anthropic-messages",
"base_url": "https://api.minimaxi.com/anthropic",
"model": "MiniMax-test",
},
},
}
}
)
async def fake_get_setting_record(_db, category):
assert category == "external_integrations"
return record
monkeypatch.setattr(settings_api, "get_setting_record", fake_get_setting_record)
runtime_config = await get_runtime_ai_provider_config(object())
assert runtime_config["llm_config"]["provider"] == "minimax"
assert runtime_config["llm_config"]["api_key"] == "minimax-key"
assert runtime_config["llm_config"]["model"] == "MiniMax-test"

View File

@@ -0,0 +1,161 @@
"""Tests for the v4 vessel_ais aggregation strategy."""
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock
import pytest
from app.models.vessel import AISRawObservation
from app.services.vessel_aggregation_strategy import (
DEFAULT_STRATEGY,
StrategyValidationError,
validate_strategy,
)
from app.services.vessel_ais_aggregation import aggregate_vessel_observations
def _obs(*, source: str, mmsi: int, observed_at: datetime, **payload) -> AISRawObservation:
payload = {"mmsi": mmsi, "lat": 50.0, "lon": 10.0, **payload}
delivery_mode = "realtime_stream" if source == "aisstream_vessels" else "polling"
transport = "websocket" if source == "aisstream_vessels" else "http"
return AISRawObservation(
target_schema="vessel_ais",
source=source,
entity_key=str(mmsi),
delivery_mode=delivery_mode,
transport=transport,
message_type="PositionReport",
observation_hash=f"{source}:{mmsi}:{observed_at.isoformat()}",
observed_at=observed_at,
collected_at=observed_at,
normalized_payload=payload,
raw_payload=payload,
quality_flags=[],
)
def test_validate_rejects_unknown_field():
with pytest.raises(StrategyValidationError, match="unknown vessel_ais field"):
validate_strategy({"vessel_ais": {"field_rules": {"definitely_not_a_field": {"mode": "newest"}}}})
def test_validate_rejects_dynamic_lock_without_flag():
with pytest.raises(StrategyValidationError, match="allow_dynamic_lock"):
validate_strategy(
{
"vessel_ais": {
"field_rules": {"lat": {"mode": "source_priority"}},
"allow_dynamic_lock": False,
}
}
)
def test_validate_allows_dynamic_lock_with_flag():
normalized = validate_strategy(
{
"version": 0,
"vessel_ais": {
"field_rules": {"lat": {"mode": "source_priority", "source_priority": ["barentswatch_vessels"]}},
"allow_dynamic_lock": True,
},
}
)
assert normalized["vessel_ais"]["field_rules"]["lat"]["mode"] == "source_priority"
assert normalized["version"] == 1
def test_validate_increments_version():
first = validate_strategy({"version": 5, "vessel_ais": {}})
assert first["version"] == 6
@pytest.mark.asyncio
async def test_strategy_field_rule_promotes_specific_source(monkeypatch):
now = datetime(2026, 5, 4, 12, 0, tzinfo=timezone.utc)
obs_a = _obs(
source="aisstream_vessels",
mmsi=257123000,
observed_at=now,
name="AISSTREAM ONE",
vessel_type_name="Cargo",
)
obs_b = _obs(
source="barentswatch_vessels",
mmsi=257123000,
observed_at=now - timedelta(seconds=1),
name="BARENTSWATCH ONE",
vessel_type_name="Cargo",
)
strategy = {
"version": 7,
"vessel_ais": {
"source_priority": [],
"field_rules": {
"name": {"mode": "source_priority", "source_priority": ["barentswatch_vessels", "aisstream_vessels"]},
},
"freshness": {"realtime_stream_seconds": 0, "polling_seconds": 0},
"allow_dynamic_lock": False,
},
}
db = AsyncMock()
vessels = await aggregate_vessel_observations(
db,
[obs_a, obs_b],
write_conflicts=False,
strategy=strategy,
)
assert len(vessels) == 1
vessel = vessels[0]
assert vessel["name"] == "BARENTSWATCH ONE"
assert vessel["field_sources"]["name"] == "barentswatch_vessels"
assert vessel["selected_reasons"]["name"] == "source_priority"
assert vessel["aggregation_strategy_version"] == 7
@pytest.mark.asyncio
async def test_strategy_freshness_falls_back_to_polling_when_realtime_stale():
now = datetime(2026, 5, 4, 12, 0, tzinfo=timezone.utc)
stale_realtime = _obs(
source="aisstream_vessels",
mmsi=257123000,
observed_at=now - timedelta(hours=1),
lat=58.0,
lon=10.0,
)
fresh_polling = _obs(
source="barentswatch_vessels",
mmsi=257123000,
observed_at=now - timedelta(seconds=30),
lat=60.0,
lon=11.0,
)
strategy = {
"version": 1,
"vessel_ais": {
"source_priority": ["aisstream_vessels", "barentswatch_vessels"],
"field_rules": {},
"freshness": {"realtime_stream_seconds": 900, "polling_seconds": 7200},
"allow_dynamic_lock": False,
},
}
db = AsyncMock()
vessels = await aggregate_vessel_observations(
db,
[stale_realtime, fresh_polling],
write_conflicts=False,
strategy=strategy,
)
assert vessels[0]["field_sources"]["lat"] == "barentswatch_vessels"
assert vessels[0]["lat"] == 60.0
def test_default_strategy_is_stable():
assert DEFAULT_STRATEGY["vessel_ais"]["allow_dynamic_lock"] is False
assert "freshness" in DEFAULT_STRATEGY["vessel_ais"]

View File

@@ -0,0 +1,155 @@
"""Tests for v5 enrichment + conflict promote-to-rule."""
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock
import pytest
from app.models.vessel import AISConflictRecord, AISRawObservation
from app.models.vessel_enrichment import VesselMediaEnrichment, VesselProfileEnrichment
from app.services.vessel_ais_aggregation import aggregate_vessel_observations
from app.services.vessel_enrichment import (
_apply_upsert,
get_vessel_enrichment_bundle,
)
class _StoreSession:
"""Minimal AsyncSession stand-in that tracks mmsi-keyed enrichment + a strategy."""
def __init__(self, *, profile=None, media=None, conflicts=None):
self.profile = profile
self.media = media
self.conflicts = list(conflicts or [])
self.added: list = []
self.committed = False
async def get(self, model, key):
if model is VesselProfileEnrichment:
return self.profile if self.profile and self.profile.mmsi == key else None
if model is VesselMediaEnrichment:
return self.media if self.media and self.media.mmsi == key else None
return None
@pytest.mark.asyncio
async def test_enrichment_bundle_filters_expired_records():
now = datetime.now(timezone.utc)
fresh = VesselProfileEnrichment(
mmsi=257123000,
source="local_cache",
payload={"vessel_subtype": "Container"},
fetched_at=now - timedelta(hours=1),
expires_at=now + timedelta(days=7),
confidence=0.9,
)
expired_media = VesselMediaEnrichment(
mmsi=257123000,
source="vesselfinder",
payload={"images": ["https://example.com/a.jpg"]},
fetched_at=now - timedelta(days=30),
expires_at=now - timedelta(days=1),
)
db = _StoreSession(profile=fresh, media=expired_media)
bundle = await get_vessel_enrichment_bundle(db, 257123000)
assert bundle["profile"]["payload"]["vessel_subtype"] == "Container"
assert bundle["media"] is None
def test_apply_upsert_preserves_payload_and_metadata():
record = VesselProfileEnrichment(mmsi=257123000)
out = _apply_upsert(
record,
{
"source": "vesselfinder",
"payload": {"vessel_subtype": "Container", "operator": "Maersk"},
"expires_at": "2026-12-31T00:00:00Z",
"confidence": 0.85,
"reference_url": "https://www.vesselfinder.com/vessels/257123000",
},
)
assert out["payload"]["operator"] == "Maersk"
assert out["confidence"] == 0.85
assert record.reference_url == "https://www.vesselfinder.com/vessels/257123000"
assert record.expires_at is not None
assert record.expires_at.year == 2026
def _obs(*, source: str, mmsi: int, observed_at, **payload) -> AISRawObservation:
payload = {"mmsi": mmsi, "lat": 60.0, "lon": 5.0, **payload}
delivery_mode = "realtime_stream" if source == "aisstream_vessels" else "polling"
transport = "websocket" if source == "aisstream_vessels" else "http"
return AISRawObservation(
target_schema="vessel_ais",
source=source,
entity_key=str(mmsi),
delivery_mode=delivery_mode,
transport=transport,
message_type="PositionReport",
observation_hash=f"{source}:{mmsi}:{observed_at.isoformat()}",
observed_at=observed_at,
collected_at=observed_at,
normalized_payload=payload,
raw_payload=payload,
quality_flags=[],
)
@pytest.mark.asyncio
async def test_promoted_rule_wins_during_aggregation():
"""Simulate the strategy that conflict-promote-to-rule writes."""
now = datetime.now(timezone.utc)
obs_a = _obs(
source="aisstream_vessels",
mmsi=257111000,
observed_at=now,
name="STREAM NAME",
vessel_type_name="Cargo",
)
obs_b = _obs(
source="barentswatch_vessels",
mmsi=257111000,
observed_at=now - timedelta(seconds=1),
name="REST NAME",
vessel_type_name="Cargo",
)
promoted_strategy = {
"version": 99,
"vessel_ais": {
"source_priority": [],
"field_rules": {
"name": {"mode": "source_priority", "source_priority": ["barentswatch_vessels"]}
},
"freshness": {"realtime_stream_seconds": 0, "polling_seconds": 0},
"allow_dynamic_lock": False,
},
}
db = AsyncMock()
vessels = await aggregate_vessel_observations(
db,
[obs_a, obs_b],
write_conflicts=False,
strategy=promoted_strategy,
)
assert vessels[0]["name"] == "REST NAME"
assert vessels[0]["selected_reasons"]["name"] == "source_priority"
assert vessels[0]["aggregation_strategy_version"] == 99
def test_conflict_record_holds_selected_source():
"""Sanity: the promote-to-rule API reads selected_source from this column."""
record = AISConflictRecord(
target_schema="vessel_ais",
entity_key="257111000",
field="name",
candidates={"a": "X", "b": "Y"},
selected_source="barentswatch_vessels",
selected_value="Y",
selected_reason="delivery_mode_priority",
)
serialized = record.to_dict()
assert serialized["selected_source"] == "barentswatch_vessels"
assert serialized["field"] == "name"

View File

@@ -1,13 +1,23 @@
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock
import pytest
from httpx import ASGITransport, AsyncClient
from app.api.v1 import visualization
from app.api.v1.visualization import convert_vessels_to_geojson
from app.db.session import get_db
from app.main import app
from app.models.vessel import VesselPosition, VesselStatic
from app.models.vessel import AISRawObservation, VesselPosition, VesselStatic
from app.services import barentswatch
from app.services.collectors.aisstream import AISStreamCollector
from app.services.collectors.vessel_ais import VesselAISCollector
from app.services.vessel_ais_aggregation import (
aggregate_vessel_observations,
build_field_conflict_candidates,
build_observation_hash,
record_vessel_ais_observation,
)
def test_vessel_collector_transforms_barentswatch_like_records():
@@ -34,6 +44,423 @@ def test_vessel_collector_transforms_barentswatch_like_records():
assert records[0]["lat"] == pytest.approx(59.91)
def test_vessel_observation_hash_is_stable_for_same_payload():
observed_at = datetime(2026, 4, 30, 12, 0, tzinfo=timezone.utc)
payload = {
"mmsi": 257123000,
"lat": 59.91,
"lon": 10.73,
"received_at": observed_at,
}
first = build_observation_hash(
source="barentswatch_vessels",
entity_key="257123000",
message_type="PositionReport",
observed_at=observed_at,
normalized_payload=payload,
)
second = build_observation_hash(
source="barentswatch_vessels",
entity_key="257123000",
message_type="PositionReport",
observed_at=observed_at,
normalized_payload=dict(reversed(payload.items())),
)
assert first == second
assert len(first) == 64
@pytest.mark.asyncio
async def test_record_vessel_ais_observation_skips_existing_hash():
observed_at = datetime(2026, 4, 30, 12, 0, tzinfo=timezone.utc)
class _Result:
def scalar_one_or_none(self):
return 123
class _Session:
def __init__(self):
self.added = []
async def execute(self, _stmt):
return _Result()
def add(self, item):
self.added.append(item)
db = _Session()
observation = await record_vessel_ais_observation(
db,
source="barentswatch_vessels",
normalized_payload={
"mmsi": 257123000,
"lat": 59.91,
"lon": 10.73,
"received_at": observed_at,
},
delivery_mode="polling",
transport="http",
observed_at=observed_at.isoformat(),
)
assert observation is None
assert db.added == []
def test_build_field_conflict_candidates_from_raw_observations():
observations = [
AISRawObservation(
source="barentswatch_vessels",
normalized_payload={"name": "OSLO TRADER", "flag": "NO"},
),
AISRawObservation(
source="aisstream_vessels",
normalized_payload={"name": "OSLO TRADER II", "flag": "NO"},
),
]
conflicts = build_field_conflict_candidates(observations)
assert conflicts == [
{
"field": "name",
"candidates": {
"aisstream_vessels": "OSLO TRADER II",
"barentswatch_vessels": "OSLO TRADER",
},
"status": "candidate",
}
]
@pytest.mark.asyncio
async def test_aggregate_vessel_observations_prefers_realtime_and_records_conflict():
observed_at = datetime.now(timezone.utc) - timedelta(minutes=5)
class _Result:
def scalar_one_or_none(self):
return None
class _Session:
def __init__(self):
self.added = []
async def execute(self, _stmt):
return _Result()
def add(self, item):
self.added.append(item)
db = _Session()
observations = [
AISRawObservation(
id=1,
source="barentswatch_vessels",
entity_key="257123000",
delivery_mode="polling",
transport="http",
observed_at=observed_at,
collected_at=observed_at,
normalized_payload={
"mmsi": 257123000,
"name": "OSLO TRADER",
"lat": 59.91,
"lon": 10.73,
},
),
AISRawObservation(
id=2,
source="aisstream_vessels",
entity_key="257123000",
delivery_mode="realtime_stream",
transport="websocket",
observed_at=observed_at + timedelta(seconds=10),
collected_at=observed_at + timedelta(seconds=10),
normalized_payload={
"mmsi": 257123000,
"vessel_type": 79,
"lat": 59.92,
"lon": 10.74,
},
raw_payload={"MetaData": {"ShipName": "OSLO TRADER II "}},
),
]
vessels = await aggregate_vessel_observations(db, observations)
assert vessels[0]["lat"] == pytest.approx(59.92)
assert vessels[0]["field_sources"]["lat"] == "aisstream_vessels"
assert vessels[0]["name"] == "OSLO TRADER II"
assert vessels[0]["vessel_type_name"] == "Cargo"
assert vessels[0]["source_summary"]["aisstream_vessels"]["observation_count"] == 1
assert vessels[0]["source_summary"]["barentswatch_vessels"]["delivery_mode"] == "polling"
assert vessels[0]["conflict_count"] == 0
assert db.added == []
@pytest.mark.asyncio
async def test_vessel_collector_writes_raw_observations_only(monkeypatch):
collector = VesselAISCollector()
collector.update_progress = AsyncMock()
record_observation = AsyncMock()
update_health = AsyncMock()
broadcast_custom = AsyncMock()
monkeypatch.setattr(
"app.services.collectors.vessel_ais.record_vessel_ais_observation",
record_observation,
)
monkeypatch.setattr(
"app.services.collectors.vessel_ais.update_ais_source_health",
update_health,
)
monkeypatch.setattr(
"app.services.collectors.vessel_ais.broadcaster.broadcast_custom",
broadcast_custom,
)
class _Session:
def __init__(self):
self.added = []
self.committed = False
async def get(self, *_args):
return None
def add(self, item):
self.added.append(item)
async def execute(self, _stmt):
return None
async def commit(self):
self.committed = True
db = _Session()
observed_at = datetime(2026, 4, 30, 12, 0, tzinfo=timezone.utc)
saved = await collector._save_data(
db,
[
{
"mmsi": 257123000,
"name": "OSLO TRADER",
"lat": 59.91,
"lon": 10.73,
"received_at": observed_at,
}
],
)
assert saved == 1
assert db.committed is True
# BarentsWatch must funnel through the unified AIS pipeline only — no legacy writes.
assert not any(isinstance(item, VesselStatic) for item in db.added)
assert not any(isinstance(item, VesselPosition) for item in db.added)
record_observation.assert_awaited_once()
assert record_observation.await_args.kwargs["source"] == "barentswatch_vessels"
assert record_observation.await_args.kwargs["normalized_payload"]["mmsi"] == 257123000
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]["action"] == "upsert"
assert broadcast_custom.await_args.args[1]["vessels"][0]["mmsi_display"] == "257123000"
def test_aisstream_collector_normalizes_position_report():
collector = AISStreamCollector()
records = collector.transform(
[
{
"MessageType": "PositionReport",
"MetaData": {
"MMSI": 257123000,
"ShipName": "OSLO TRADER ",
"time_utc": "2026-04-30T12:00:00Z",
},
"Message": {
"PositionReport": {
"Latitude": 59.91,
"Longitude": 10.73,
"Sog": 12.4,
"Cog": 214,
"TrueHeading": 215,
"NavigationalStatus": 0,
}
},
}
]
)
assert len(records) == 1
assert records[0]["mmsi"] == 257123000
assert records[0]["lat"] == pytest.approx(59.91)
assert records[0]["name"] == "OSLO TRADER"
assert records[0]["_message_type"] == "PositionReport"
def test_aisstream_collector_maps_ship_static_type_name():
collector = AISStreamCollector()
records = collector.transform(
[
{
"MessageType": "ShipStaticData",
"MetaData": {
"MMSI": 257123000,
"time_utc": "2026-04-30T12:00:00Z",
},
"Message": {
"ShipStaticData": {
"Name": "OSLO TRADER",
"Type": 79,
"CallSign": "LAAB",
}
},
}
]
)
assert len(records) == 1
assert records[0]["vessel_type"] == 79
assert records[0]["vessel_type_name"] == "Cargo"
@pytest.mark.asyncio
async def test_aisstream_collector_writes_only_raw_observations(monkeypatch):
collector = AISStreamCollector()
collector.update_progress = AsyncMock()
record_observation = AsyncMock(return_value=object())
update_health = AsyncMock()
monkeypatch.setattr(
"app.services.collectors.aisstream.record_vessel_ais_observation",
record_observation,
)
monkeypatch.setattr(
"app.services.collectors.aisstream.update_ais_source_health",
update_health,
)
class _Session:
def __init__(self):
self.added = []
self.committed = False
def add(self, item):
self.added.append(item)
async def commit(self):
self.committed = True
db = _Session()
saved = await collector._save_data(
db,
[
{
"mmsi": 257123000,
"lat": 59.91,
"lon": 10.73,
"received_at": datetime(2026, 4, 30, 12, 0, tzinfo=timezone.utc),
"_message_type": "PositionReport",
}
],
)
assert saved == 1
assert db.added == []
assert db.committed is True
record_observation.assert_awaited_once()
assert record_observation.await_args.kwargs["source"] == "aisstream_vessels"
update_health.assert_awaited_once()
@pytest.mark.asyncio
async def test_aisstream_stream_record_broadcasts_vessel_delta(monkeypatch):
collector = AISStreamCollector()
record_observation = AsyncMock(return_value=object())
update_health = AsyncMock()
broadcast_custom = AsyncMock()
monkeypatch.setattr(
"app.services.collectors.aisstream.record_vessel_ais_observation",
record_observation,
)
monkeypatch.setattr(
"app.services.collectors.aisstream.update_ais_source_health",
update_health,
)
monkeypatch.setattr(
"app.services.collectors.aisstream.broadcaster.broadcast_custom",
broadcast_custom,
)
class _Session:
async def commit(self):
pass
created = await collector._save_stream_record(
_Session(),
{
"mmsi": 257123000,
"lat": 59.91,
"lon": 10.73,
"cog": 214,
"received_at": datetime(2026, 4, 30, 12, 0, tzinfo=timezone.utc),
},
)
assert created is True
record_observation.assert_awaited_once()
broadcast_custom.assert_awaited_once()
assert broadcast_custom.await_args.args[0] == "vessels"
assert broadcast_custom.await_args.args[1]["action"] == "upsert"
assert broadcast_custom.await_args.args[1]["vessels"][0]["mmsi_display"] == "257123000"
def test_barentswatch_reads_credentials_from_zshrc(tmp_path):
zshrc = tmp_path / ".zshrc"
zshrc.write_text(
"\n".join(
[
"export BARENTSWATCH_CLIENT_ID='client-from-zshrc'",
'export BARENTSWATCH_CLIENT_SECRET="secret-from-zshrc" # local dev credential',
]
),
encoding="utf-8",
)
values = barentswatch._read_zshrc_env(zshrc)
assert values["BARENTSWATCH_CLIENT_ID"] == "client-from-zshrc"
assert values["BARENTSWATCH_CLIENT_SECRET"] == "secret-from-zshrc"
@pytest.mark.asyncio
async def test_barentswatch_resolves_config_from_zshrc(tmp_path, monkeypatch):
zshrc = tmp_path / ".zshrc"
zshrc.write_text(
"\n".join(
[
"export BARENTSWATCH_CLIENT_ID=client-from-zshrc",
"export BARENTSWATCH_CLIENT_SECRET=secret-from-zshrc",
]
),
encoding="utf-8",
)
monkeypatch.delenv("BARENTSWATCH_CLIENT_ID", raising=False)
monkeypatch.delenv("BARENTSWATCH_CLIENT_SECRET", raising=False)
monkeypatch.delenv("BARRENTSWATCH_CLIENT_ID", raising=False)
monkeypatch.delenv("BARRENTSWATCH_CLIENT_SECRET", raising=False)
monkeypatch.setattr(barentswatch.Path, "home", lambda: tmp_path)
config = await barentswatch.resolve_barentswatch_config(None)
assert config.client_id == "client-from-zshrc"
assert config.client_secret == "secret-from-zshrc"
assert config.credential_source == "~/.zshrc"
def test_convert_vessels_to_geojson():
position = VesselPosition(
mmsi=257123000,
@@ -62,6 +489,39 @@ def test_convert_vessels_to_geojson():
assert payload["features"][0]["properties"]["vessel_type_name"] == "Cargo"
def test_convert_vessels_to_geojson_dedupes_mmsi_rows():
first = VesselPosition(
mmsi=257123000,
lat=59.91,
lon=10.73,
received_at=datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc),
)
duplicate = VesselPosition(
mmsi=257123000,
lat=60.01,
lon=10.83,
received_at=datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc),
)
other = VesselPosition(
mmsi=257456000,
lat=60.3,
lon=5.3,
received_at=datetime(2026, 4, 28, 0, 59, tzinfo=timezone.utc),
)
payload = convert_vessels_to_geojson(
[
(first, VesselStatic(mmsi=257123000, name="OSLO TRADER")),
(duplicate, VesselStatic(mmsi=257123000, name="OSLO TRADER DUP")),
(other, VesselStatic(mmsi=257456000, name="BERGEN FERRY")),
]
)
mmsis = [feature["properties"]["mmsi"] for feature in payload["features"]]
assert mmsis == [257123000, 257456000]
assert payload["features"][0]["geometry"]["coordinates"] == [10.73, 59.91]
@pytest.mark.asyncio
async def test_vessels_geojson_endpoint_filters_type_and_bbox():
now = datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc)
@@ -93,7 +553,7 @@ async def test_vessels_geojson_endpoint_filters_type_and_bbox():
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.get(
"/api/v1/visualization/geo/vessels",
params={"bbox": "0,50,20,70", "type": "cargo"},
params={"bbox": "0,50,20,70", "type": "cargo", "limit": 0},
)
assert response.status_code == 200
@@ -103,3 +563,113 @@ async def test_vessels_geojson_endpoint_filters_type_and_bbox():
assert data["stats"]["by_type"]["Cargo"] == 1
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_vessels_geojson_merges_raw_and_legacy_sources(monkeypatch):
now = datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc)
monkeypatch.setattr(
visualization,
"get_aggregated_vessels",
AsyncMock(
return_value=[
{
"mmsi": 1,
"lat": 59.9,
"lon": 10.7,
"received_at": now,
"name": "AISSTREAM SHIP",
"vessel_type_name": "Cargo",
"source_summary": {"aisstream_vessels": {"message_types": ["PositionReport"]}},
}
]
),
)
rows = [
(
VesselPosition(mmsi=1, lat=60.0, lon=10.8, received_at=now),
VesselStatic(mmsi=1, name="LEGACY DUP", vessel_type_name="Cargo"),
),
(
VesselPosition(mmsi=2, lat=60.3, lon=5.3, received_at=now),
VesselStatic(mmsi=2, name="BARENTSWATCH ONLY", vessel_type_name="Passenger"),
),
]
class _Result:
def all(self):
return rows
class _FakeSession:
async def execute(self, _query):
return _Result()
async def override_get_db():
yield _FakeSession()
app.dependency_overrides[get_db] = override_get_db
transport = ASGITransport(app=app)
try:
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.get("/api/v1/visualization/geo/vessels")
assert response.status_code == 200
data = response.json()
names = {feature["properties"]["mmsi"]: feature["properties"]["name"] for feature in data["features"]}
assert data["count"] == 2
assert names == {1: "AISSTREAM SHIP", 2: "BARENTSWATCH ONLY"}
assert data["diagnostics"]["legacy_backfilled_mmsi"] == 1
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_vessel_name_fallbacks_reports_mmsi_display_names(monkeypatch):
now = datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc)
monkeypatch.setattr(
visualization,
"get_aggregated_vessels",
AsyncMock(
return_value=[
{
"mmsi": 257123000,
"lat": 59.9,
"lon": 10.7,
"received_at": now,
"name": "MMSI 257123000",
"vessel_type_name": "Other",
"source_summary": {
"aisstream_vessels": {
"latest_observed_at": now,
"message_types": ["PositionReport"],
}
},
}
]
),
)
class _Result:
def all(self):
return []
class _FakeSession:
async def execute(self, _query):
return _Result()
async def override_get_db():
yield _FakeSession()
app.dependency_overrides[get_db] = override_get_db
transport = ASGITransport(app=app)
try:
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.get("/api/v1/visualization/vessels/name-fallbacks")
assert response.status_code == 200
data = response.json()
assert data["count"] == 1
assert data["items"][0]["mmsi"] == "257123000"
assert data["items"][0]["message_types"] == ["PositionReport"]
finally:
app.dependency_overrides.clear()

View File

@@ -1,9 +1,15 @@
from datetime import datetime, timezone
from unittest.mock import AsyncMock
import pytest
from httpx import ASGITransport, AsyncClient
from app.api.v1.visualization import convert_compute_centers_to_geojson
from app.api.v1 import visualization as visualization_api
from app.api.v1.visualization import (
CollectComputeCenterLocationRequest,
convert_compute_centers_to_geojson,
)
import app.services.compute_center_locations as compute_center_locations
from app.db.session import get_db
from app.main import app
from app.models.collected_data import CollectedData
@@ -89,6 +95,8 @@ def test_convert_compute_centers_to_geojson_unifies_sources():
assert supercomputer_feature["properties"]["operator"] == "ORNL"
assert supercomputer_feature["properties"]["location_precision"] == "precise"
assert supercomputer_feature["properties"]["is_estimated"] is False
assert supercomputer_feature["properties"]["location_source"] == "source_coordinates"
assert supercomputer_feature["properties"]["location_confidence"] == 1.0
gpu_feature = payload["features"][1]
assert gpu_feature["properties"]["site_type"] == "gpu_cluster"
@@ -98,8 +106,110 @@ def test_convert_compute_centers_to_geojson_unifies_sources():
assert gpu_feature["properties"]["location_precision"] == "precise"
def test_convert_compute_centers_to_geojson_uses_coordinate_hints():
hinted_record = _build_record(
def test_convert_compute_centers_to_geojson_accepts_source_coordinate_aliases():
record = _build_record(
record_id=3,
source="epoch_ai_gpu",
data_type="gpu_cluster",
name="Alias Coordinates",
country="United States",
city="New York",
latitude=0.0,
longitude=0.0,
metadata={
"latitude": "",
"longitude": "",
"location": {
"lat": 40.7128,
"lng": -74.0060,
},
"value": "1200",
"unit": "TFlop/s",
},
)
payload = convert_compute_centers_to_geojson([record])
assert len(payload["features"]) == 1
feature = payload["features"][0]
assert feature["geometry"]["coordinates"] == [-74.006, 40.7128]
assert feature["properties"]["location_source"] == "source_coordinates"
def test_compute_center_source_coordinates_win_over_stored_location():
compute_center_locations.set_compute_center_location_cache({
"top500:top500-31": {
"source": "top500",
"source_id": "top500-31",
"name": "Stored Wrong",
"latitude": 1.0,
"longitude": 2.0,
"precision": "city",
"confidence": 0.5,
"needs_confirmation": True,
}
})
record = _build_record(
record_id=31,
source="top500",
data_type="supercomputer",
name="Source Wins",
country="United States",
city="Oak Ridge",
latitude=35.93,
longitude=-84.31,
metadata={"organization": "ORNL"},
)
payload = convert_compute_centers_to_geojson([record])
assert payload["features"][0]["geometry"]["coordinates"] == [-84.31, 35.93]
assert payload["features"][0]["properties"]["location_source"] == "source_coordinates"
compute_center_locations.set_compute_center_location_cache({})
def test_compute_center_geojson_uses_stored_location_when_source_coords_missing():
compute_center_locations.set_compute_center_location_cache({
"epoch_ai_gpu:epoch_ai_gpu-32": {
"source": "epoch_ai_gpu",
"source_id": "epoch_ai_gpu-32",
"name": "Stored Cluster",
"city": "Memphis",
"country": "United States",
"latitude": 35.1495,
"longitude": -90.049,
"precision": "city",
"confidence": 0.72,
"location_source": "manual_selection",
"source_note": "Saved by user",
"needs_confirmation": False,
"verified_at": "2026-05-08T00:00:00Z",
}
})
record = _build_record(
record_id=32,
source="epoch_ai_gpu",
data_type="gpu_cluster",
name="Stored Cluster",
country="United States",
city="",
latitude=0.0,
longitude=0.0,
metadata={"value": "1200", "unit": "TFlop/s"},
)
payload = convert_compute_centers_to_geojson([record])
assert len(payload["features"]) == 1
feature = payload["features"][0]
assert feature["geometry"]["coordinates"] == [-90.049, 35.1495]
assert feature["properties"]["location_source"] == "stored_compute_center_location"
assert feature["properties"]["needs_confirmation"] is False
compute_center_locations.set_compute_center_location_cache({})
def test_convert_compute_centers_to_geojson_does_not_use_registry_aliases():
registry_record = _build_record(
record_id=3,
source="top500",
data_type="supercomputer",
@@ -114,23 +224,51 @@ def test_convert_compute_centers_to_geojson_uses_coordinate_hints():
},
)
payload = convert_compute_centers_to_geojson([hinted_record])
payload = convert_compute_centers_to_geojson([registry_record])
assert len(payload["features"]) == 1
coords = payload["features"][0]["geometry"]["coordinates"]
assert coords[0] == pytest.approx(-84.3107)
assert coords[1] == pytest.approx(35.9319)
assert payload["features"][0]["properties"]["is_estimated"] is True
assert payload["features"][0]["properties"]["location_precision"] == "estimated_site"
assert payload["features"] == []
assert len(payload["unresolved"]) == 1
assert payload["unresolved"][0]["name"] == "Frontier"
assert "source coords" in payload["unresolved"][0]["failure_reason"]
def test_convert_compute_centers_to_geojson_falls_back_to_country_centroid():
centroid_record = _build_record(
def test_convert_compute_centers_to_geojson_does_not_use_city_fallback():
city_record = _build_record(
record_id=4,
source="epoch_ai_gpu",
data_type="gpu_cluster",
name="Sample GPU Cluster",
country="United States",
city="San Francisco, CA",
latitude=0.0,
longitude=0.0,
metadata={
"organization": "Sample Operator",
"value": "10000",
"unit": "TFlop/s",
},
)
payload = convert_compute_centers_to_geojson([city_record])
assert payload["features"] == []
assert len(payload["unresolved"]) == 1
assert payload["unresolved"][0]["city"] == "San Francisco, CA"
def test_convert_compute_centers_to_geojson_does_not_online_geocode_on_startup(monkeypatch):
compute_center_locations._geocode_online.cache_clear()
def _explode(_query):
raise AssertionError("startup GeoJSON must not call online geocoding")
monkeypatch.setattr(compute_center_locations, "_geocode_online", _explode)
country_record = _build_record(
record_id=4,
source="epoch_ai_gpu",
data_type="gpu_cluster",
name="Unknown Cluster",
country="United States",
country="France",
city="",
latitude=0.0,
longitude=0.0,
@@ -141,16 +279,339 @@ def test_convert_compute_centers_to_geojson_falls_back_to_country_centroid():
},
)
payload = convert_compute_centers_to_geojson([centroid_record])
payload = convert_compute_centers_to_geojson([country_record])
assert len(payload["features"]) == 1
props = payload["features"][0]["properties"]
coords = payload["features"][0]["geometry"]["coordinates"]
assert coords[0] == pytest.approx(-98.5795)
assert coords[1] == pytest.approx(39.8283)
assert props["is_estimated"] is True
assert props["location_precision"] == "estimated_country"
assert props["geography_mode"] == "country_centroid"
assert payload["features"] == []
assert len(payload["unresolved"]) == 1
assert payload["unresolved"][0]["operator"] == "Unknown Operator"
def test_convert_compute_centers_to_geojson_records_diagnostics_when_online_geocode_fails(monkeypatch):
compute_center_locations._geocode_online.cache_clear()
monkeypatch.setattr(compute_center_locations, "_geocode_online", lambda _query: None)
country_record = _build_record(
record_id=5,
source="epoch_ai_gpu",
data_type="gpu_cluster",
name="Unknown French Cluster",
country="France",
city="",
latitude=0.0,
longitude=0.0,
metadata={
"organization": "Unknown Operator",
"value": "10000",
"unit": "TFlop/s",
},
)
payload = convert_compute_centers_to_geojson([country_record])
assert payload["features"] == []
assert len(payload["unresolved"]) == 1
diagnostic = payload["unresolved"][0]
assert diagnostic["record_id"] == 5
assert diagnostic["source_id"] == "epoch_ai_gpu-5"
assert diagnostic["country"] == "France"
assert diagnostic["operator"] == "Unknown Operator"
assert diagnostic["failure_reason"]
assert diagnostic["attempted_queries"] == []
def test_convert_compute_centers_to_geojson_records_diagnostics_when_no_country(monkeypatch):
compute_center_locations._geocode_online.cache_clear()
monkeypatch.setattr(compute_center_locations, "_geocode_online", lambda _query: None)
unknown_record = _build_record(
record_id=6,
source="epoch_ai_gpu",
data_type="gpu_cluster",
name="Unknown Offshore Cluster",
country="",
city="",
latitude=0.0,
longitude=0.0,
metadata={
"organization": "Unknown Operator",
"value": "10000",
"unit": "TFlop/s",
},
)
payload = convert_compute_centers_to_geojson([unknown_record])
assert payload["features"] == []
assert len(payload["unresolved"]) == 1
assert payload["unresolved"][0]["failure_reason"]
def test_convert_compute_centers_to_geojson_never_emits_zero_coordinates(monkeypatch):
compute_center_locations._geocode_online.cache_clear()
def _zero_geocode(query):
return {
"lat": "0",
"lon": "0",
"display_name": "Null Island",
"address": {"city": "", "country": ""},
}
monkeypatch.setattr(compute_center_locations, "_geocode_online", _zero_geocode)
record = _build_record(
record_id=7,
source="epoch_ai_gpu",
data_type="gpu_cluster",
name="Null Island Cluster",
country="",
city="",
latitude=0.0,
longitude=0.0,
metadata={"organization": "Null Inc"},
)
payload = convert_compute_centers_to_geojson([record])
for feature in payload["features"]:
coords = feature["geometry"]["coordinates"]
assert coords[0] not in (0, 0.0)
assert coords[1] not in (0, 0.0)
def test_convert_compute_centers_to_geojson_rejects_country_or_unknown_precision(monkeypatch):
compute_center_locations._geocode_online.cache_clear()
monkeypatch.setattr(compute_center_locations, "_geocode_online", lambda _query: None)
record = _build_record(
record_id=8,
source="top500",
data_type="supercomputer",
name="Phantom System",
country="Liechtenstein",
city="",
latitude=0.0,
longitude=0.0,
metadata={"organization": "Phantom Operator", "rmax": 100.0},
)
payload = convert_compute_centers_to_geojson([record])
for feature in payload["features"]:
assert feature["properties"]["location_precision"] in {"precise", "site", "city"}
assert payload["features"] == []
assert payload["unresolved"], "phantom record must surface as diagnostic"
def test_resolve_full_returns_diagnostic_for_unresolved(monkeypatch):
compute_center_locations._geocode_online.cache_clear()
monkeypatch.setattr(compute_center_locations, "_geocode_online", lambda _query: None)
record = _build_record(
record_id=11,
source="epoch_ai_gpu",
data_type="gpu_cluster",
name="Phantom Cluster",
country="Bhutan",
city="",
latitude=0.0,
longitude=0.0,
metadata={"organization": "Mystery Operator"},
)
result = compute_center_locations.resolve_compute_center_location_full(record, record.extra_data)
assert result.location is None
assert result.diagnostic is not None
assert result.diagnostic.failure_reason
assert result.diagnostic.country == "Bhutan"
def test_collect_location_candidates_ignores_registry_and_uses_online(monkeypatch):
compute_center_locations._geocode_online.cache_clear()
def _fake_ror(query):
assert query == "Oak Ridge National Laboratory"
return {
"id": "https://ror.org/01qz5mb56",
"names": [
{"types": ["ror_display"], "value": "Oak Ridge National Laboratory"}
],
"locations": [
{
"geonames_id": 4646571,
"geonames_details": {
"name": "Oak Ridge",
"country_subdivision_name": "Tennessee",
"country_name": "United States",
"lat": 36.01036,
"lng": -84.26964,
},
}
],
}
monkeypatch.setattr(compute_center_locations, "_lookup_ror_organization", _fake_ror)
monkeypatch.setattr(compute_center_locations, "_geocode_online", lambda _query: None)
candidates, attempted = compute_center_locations.collect_location_candidates(
name="Frontier",
operator="Oak Ridge National Laboratory",
country="United States",
)
assert candidates, "online source-traced query must produce a candidate"
best = candidates[0]
assert best.source == "ror_organization_registry"
assert best.precision == "city"
assert best.needs_confirmation is True
assert attempted[0] == "ror:Oak Ridge National Laboratory"
def test_collect_location_candidates_returns_online_when_registry_misses(monkeypatch):
compute_center_locations._geocode_online.cache_clear()
def _fake_geocode(query):
if "Lyon" not in query and "Mystery Operator" not in query and "Lyon, France" not in query:
return None
return {
"lat": "45.7640",
"lon": "4.8357",
"display_name": "Lyon, Auvergne-Rhône-Alpes, France",
"address": {"city": "Lyon", "state": "Auvergne-Rhône-Alpes", "country": "France"},
}
monkeypatch.setattr(compute_center_locations, "_geocode_online", _fake_geocode)
monkeypatch.setattr(compute_center_locations, "_lookup_ror_organization", lambda _query: None)
candidates, attempted = compute_center_locations.collect_location_candidates(
name="Mystery System",
operator="Mystery Operator",
city="Lyon",
country="France",
)
assert candidates, "online geocoding must produce a candidate"
online_candidates = [c for c in candidates if c.source == "nominatim_online_geocode"]
assert online_candidates, "must include at least one online candidate"
online = online_candidates[0]
assert online.precision == "city"
assert online.needs_confirmation is True
assert online.suggested_registry_entry is not None
assert attempted, "must record attempted query strings"
def test_collect_location_candidates_failure_returns_attempted_queries(monkeypatch):
compute_center_locations._geocode_online.cache_clear()
monkeypatch.setattr(compute_center_locations, "_lookup_ror_organization", lambda _query: None)
monkeypatch.setattr(compute_center_locations, "_geocode_online", lambda _query: None)
candidates, attempted = compute_center_locations.collect_location_candidates(
name="Mystery Offshore Cluster",
operator="Mystery Operator",
country="Bhutan",
)
assert candidates == []
assert attempted, "even on failure we record attempted queries for diagnostics"
@pytest.mark.asyncio
async def test_collect_compute_center_location_skips_llm_when_candidates_exist(monkeypatch):
candidate = compute_center_locations.LocationCandidate(
latitude=45.764,
longitude=4.8357,
display_name="Lyon",
precision="city",
confidence=0.62,
query="Lyon, France",
source="nominatim_online_geocode",
source_note="fixture",
matched_fields=("city", "country"),
needs_confirmation=True,
city="Lyon",
country="France",
)
monkeypatch.setattr(visualization_api, "_load_compute_center_record", AsyncMock(return_value=None))
monkeypatch.setattr(
visualization_api,
"collect_location_candidates",
lambda **_kwargs: ([candidate], ["Lyon, France"]),
)
async def _explode(**_kwargs):
raise AssertionError("LLM fallback should not run when a normal candidate exists")
monkeypatch.setattr(visualization_api, "collect_llm_location_fallback_candidate", _explode)
response = await visualization_api.collect_compute_center_location(
"epoch_ai_gpu-test",
CollectComputeCenterLocationRequest(
name="Mystery Cluster",
source="epoch_ai_gpu",
city="Lyon",
country="France",
),
db=AsyncMock(),
)
assert response["success"] is True
assert response["best_candidate"]["source"] == "nominatim_online_geocode"
@pytest.mark.asyncio
async def test_collect_compute_center_location_uses_llm_when_candidates_empty(monkeypatch):
llm_candidate = compute_center_locations.LocationCandidate(
latitude=45.764,
longitude=4.8357,
display_name="Lyon, France",
precision="city",
confidence=0.74,
query="llm_factcheck:compute_center:Mystery Cluster",
source="llm_location_factcheck",
source_note="LLM location factcheck fallback",
matched_fields=("name",),
needs_confirmation=True,
city="Lyon",
country="France",
)
monkeypatch.setattr(visualization_api, "_load_compute_center_record", AsyncMock(return_value=None))
monkeypatch.setattr(
visualization_api,
"collect_location_candidates",
lambda **_kwargs: ([], ["Mystery Cluster, France"]),
)
from app.services.location.llm_fallback import LocationLLMFallbackResult, LocationSearchEvidenceResult
async def _search_evidence(**_kwargs):
return LocationSearchEvidenceResult(
evidence=[
{
"title": "Mystery Cluster source",
"url": "https://example.test/mystery",
"snippet": "Mystery Cluster is in Lyon.",
}
],
attempted_queries=["web_search:compute_center:Mystery Cluster France physical location"],
)
async def _fallback(**_kwargs):
return LocationLLMFallbackResult(
candidates=[llm_candidate],
attempted_queries=["llm_factcheck:compute_center:Mystery Cluster"],
)
monkeypatch.setattr(visualization_api, "get_ai_provider_client", AsyncMock(return_value=object()))
monkeypatch.setattr(visualization_api, "get_web_search_client", AsyncMock(return_value=object()))
monkeypatch.setattr(visualization_api, "collect_location_search_evidence", _search_evidence)
monkeypatch.setattr(visualization_api, "collect_llm_location_fallback_candidate", _fallback)
response = await visualization_api.collect_compute_center_location(
"epoch_ai_gpu-test",
CollectComputeCenterLocationRequest(
name="Mystery Cluster",
source="epoch_ai_gpu",
country="France",
),
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"] == [
"Mystery Cluster, France",
"web_search:compute_center:Mystery Cluster France physical location",
"llm_factcheck:compute_center:Mystery Cluster",
]
@pytest.mark.asyncio
@@ -292,6 +753,9 @@ async def test_visualization_geo_summary_returns_counts(monkeypatch):
def scalar(self):
return self._scalar_value
def all(self):
return list(self._rows)
def scalars(self):
class _Scalars:
def __init__(self, rows):
@@ -304,13 +768,18 @@ async def test_visualization_geo_summary_returns_counts(monkeypatch):
class _FakeSession:
async def execute(self, query):
query_text = str(query)
query_text = str(query).lower()
if "bgp_incidents" in query_text:
return _ScalarResult(scalar_value=2)
if "bgp_anomalies" in query_text:
return _ScalarResult(scalar_value=3)
if "ais_raw_observations" in query_text or "vessel_position" in query_text:
return _ScalarResult(rows=[])
return _ScalarResult(rows=records)
async def get(self, *_args, **_kwargs):
return None
async def override_get_db():
yield _FakeSession()
@@ -345,3 +814,304 @@ async def test_visualization_geo_summary_returns_counts(monkeypatch):
assert stats["bgp_collector_count"] == 2
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_collect_location_endpoint_returns_candidates_for_known_record(monkeypatch):
def _fake_ror(query):
assert query == "Oak Ridge National Laboratory"
return {
"id": "https://ror.org/01qz5mb56",
"names": [
{"types": ["ror_display"], "value": "Oak Ridge National Laboratory"}
],
"locations": [
{
"geonames_id": 4646571,
"geonames_details": {
"name": "Oak Ridge",
"country_subdivision_name": "Tennessee",
"country_name": "United States",
"lat": 36.01036,
"lng": -84.26964,
},
}
],
}
monkeypatch.setattr(compute_center_locations, "_lookup_ror_organization", _fake_ror)
monkeypatch.setattr(compute_center_locations, "_geocode_online", lambda _query: None)
target_record = _build_record(
record_id=42,
source="top500",
data_type="supercomputer",
name="Frontier",
country="United States",
city="",
latitude=0.0,
longitude=0.0,
metadata={"organization": "Oak Ridge National Laboratory", "rmax": 1102000.0},
)
class _ScalarResult:
def __init__(self, rows):
self._rows = rows
def scalars(self):
class _Scalars:
def __init__(self, rows):
self._rows = rows
def first(self):
return self._rows[0] if self._rows else None
def all(self):
return self._rows
return _Scalars(self._rows)
class _FakeSession:
async def execute(self, _query):
return _ScalarResult([target_record])
async def override_get_db():
yield _FakeSession()
app.dependency_overrides[get_db] = override_get_db
transport = ASGITransport(app=app)
try:
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/api/v1/visualization/compute-centers/top500-42/collect-location",
json={
"name": "Frontier",
"operator": "Oak Ridge National Laboratory",
"country": "United States",
},
)
assert response.status_code == 200
body = response.json()
assert body["success"] is True
assert body["candidates"], "must include candidates"
best = body["best_candidate"]
assert best["precision"] in {"precise", "site", "city"}
assert best["source"] == "ror_organization_registry"
assert best["needs_confirmation"] is True
assert best["matched_fields"], "matched_fields must be populated"
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_collect_location_endpoint_returns_failure_reason(monkeypatch):
monkeypatch.setattr(compute_center_locations, "_lookup_ror_organization", lambda _query: None)
monkeypatch.setattr(compute_center_locations, "_geocode_online", lambda _query: None)
class _ScalarResult:
def __init__(self, rows):
self._rows = rows
def scalars(self):
class _Scalars:
def __init__(self, rows):
self._rows = rows
def first(self):
return self._rows[0] if self._rows else None
def all(self):
return self._rows
return _Scalars(self._rows)
class _FakeSession:
async def execute(self, _query):
return _ScalarResult([])
async def override_get_db():
yield _FakeSession()
app.dependency_overrides[get_db] = override_get_db
transport = ASGITransport(app=app)
try:
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/api/v1/visualization/compute-centers/epoch-mystery-99/collect-location",
json={
"name": "Mystery Cluster",
"operator": "Mystery Operator",
"country": "Bhutan",
},
)
assert response.status_code == 200
body = response.json()
assert body["success"] is False
assert body["failure_reason"]
assert body["candidates"] == []
assert body["attempted_queries"], "must include attempted queries"
finally:
app.dependency_overrides.clear()
@pytest.mark.asyncio
async def test_save_location_endpoint_upserts_and_geojson_can_render():
target_record = _build_record(
record_id=52,
source="epoch_ai_gpu",
data_type="gpu_cluster",
name="Saved Cluster",
country="United States",
city="",
latitude=0.0,
longitude=0.0,
metadata={"value": "1200", "unit": "TFlop/s"},
)
class _ScalarResult:
def __init__(self, rows):
self._rows = rows
def scalars(self):
class _Scalars:
def __init__(self, rows):
self._rows = rows
def first(self):
return self._rows[0] if self._rows else None
def all(self):
return self._rows
return _Scalars(self._rows)
class _FakeSession:
def __init__(self):
self.saved = []
async def execute(self, _query):
if self.saved:
return _ScalarResult(self.saved)
return _ScalarResult([target_record])
async def scalar(self, _query):
return None
def add(self, record):
self.saved.append(record)
async def commit(self):
return None
async def refresh(self, _record):
return None
fake_session = _FakeSession()
async def override_get_db():
yield fake_session
app.dependency_overrides[get_db] = override_get_db
transport = ASGITransport(app=app)
try:
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/api/v1/visualization/compute-centers/epoch_ai_gpu-52/location",
json={
"source": "epoch_ai_gpu",
"name": "Saved Cluster",
"latitude": 35.1495,
"longitude": -90.049,
"precision": "city",
"confidence": 0.72,
"location_source": "ror_organization_registry",
"source_note": "Selected by user",
"raw_payload": {"source": "ror_organization_registry"},
},
)
assert response.status_code == 200
body = response.json()
assert body["success"] is True
assert fake_session.saved
payload = convert_compute_centers_to_geojson([target_record])
assert len(payload["features"]) == 1
feature = payload["features"][0]
assert feature["geometry"]["coordinates"] == [-90.049, 35.1495]
assert feature["properties"]["location_source"] == "stored_compute_center_location"
finally:
app.dependency_overrides.clear()
compute_center_locations.set_compute_center_location_cache({})
def test_resolution_chain_orders_source_coords_first(monkeypatch):
def _explode(_query):
raise AssertionError("source coords must short-circuit before online geocoding")
monkeypatch.setattr(compute_center_locations, "_geocode_online", _explode)
record = _build_record(
record_id=20,
source="top500",
data_type="supercomputer",
name="Frontier",
country="United States",
city="Oak Ridge",
latitude=35.93,
longitude=-84.31,
metadata={"organization": "ORNL"},
)
result = compute_center_locations.resolve_compute_center_location_full(record, record.extra_data)
assert result.is_resolved
assert result.location.location_precision == "precise"
assert result.location.location_source == "source_coordinates"
def test_no_country_centroid_or_major_compute_city_fallback(monkeypatch):
monkeypatch.setattr(compute_center_locations, "_geocode_online", lambda _query: None)
record = _build_record(
record_id=21,
source="top500",
data_type="supercomputer",
name="Phantom System",
country="France",
city="",
latitude=0.0,
longitude=0.0,
metadata={"organization": "Phantom Operator"},
)
result = compute_center_locations.resolve_compute_center_location_full(record, record.extra_data)
assert result.location is None, "must NOT fall back to country centroid or hashed major city"
assert result.diagnostic is not None
assert result.diagnostic.failure_reason
def test_repository_has_no_forbidden_precision_tokens():
"""Static guard: forbidden fallback strategies must not regress into the codebase.
Each forbidden token may appear at most once per target file, and only inside
the FORBIDDEN_PRECISIONS guard list (so we still reject them at runtime).
"""
from pathlib import Path
backend_root = Path(__file__).resolve().parents[1]
forbidden_tokens = (
"country_centroid",
"country_major_compute_city",
"estimated_country",
)
targets = [
backend_root / "app" / "services" / "compute_center_locations.py",
backend_root / "app" / "api" / "v1" / "visualization.py",
]
for target in targets:
text = target.read_text(encoding="utf-8")
for token in forbidden_tokens:
occurrences = text.count(token)
assert occurrences <= 1, (
f"{token} appears {occurrences} times in {target}; "
"should only appear in FORBIDDEN_PRECISIONS guard list."
)
if occurrences == 1:
assert "FORBIDDEN_PRECISIONS" in text, (
f"{token} appears in {target} outside the FORBIDDEN_PRECISIONS guard"
)

View File

@@ -0,0 +1,261 @@
from types import SimpleNamespace
import pytest
from app.api.v1 import settings as settings_api
from app.api.v1.settings import (
WebSearchIntegrationUpdate,
_build_web_search_payload,
_mask_secret,
_normalize_web_search_payload,
_resolve_web_search_api_key,
)
from app.services.ai_tools.schemas import WebSearchConfig, WebSearchProviderConfig
from app.services.ai_tools.web_search import WebSearchClient
from app.services.credential_guides import generate_credential_guide
from app.services.location.llm_fallback import (
collect_llm_location_fallback_candidate,
collect_location_search_evidence,
)
from app.services.location.models import LocationQuery
@pytest.fixture(autouse=True)
def isolated_web_search_env_files(monkeypatch, tmp_path):
env_file = tmp_path / ".env"
monkeypatch.setattr(settings_api, "WEB_SEARCH_ENV_FILES", (env_file,))
return env_file
def test_normalize_web_search_payload_adds_default_provider():
payload = _normalize_web_search_payload({})
assert payload["default_provider"] == "tavily"
assert payload["providers"]["tavily"]["base_url"] == "https://api.tavily.com"
def test_web_search_key_prefers_provider_env(isolated_web_search_env_files):
isolated_web_search_env_files.write_text(
"TAVILY_API_KEY=tavily-env-key\nWEB_SEARCH_API_KEY=generic-search-key\n",
encoding="utf-8",
)
value, source = _resolve_web_search_api_key("tavily", {"api_key": ""})
assert value == "tavily-env-key"
assert source == "env_file"
def test_build_web_search_payload_keeps_saved_key_when_preview_submitted():
current = {
"web_search": {
"default_provider": "tavily",
"providers": {
"tavily": {
"provider": "tavily",
"api_key": "tvly-old-secret",
},
},
}
}
update = WebSearchIntegrationUpdate(
enabled=True,
provider="tavily",
base_url="https://api.tavily.com",
api_key=_mask_secret("tvly-old-secret")["preview"],
)
payload = _build_web_search_payload(current, update)
assert payload["enabled"] is True
assert payload["providers"]["tavily"]["api_key"] == "tvly-old-secret"
@pytest.mark.asyncio
async def test_tavily_adapter_normalizes_results(monkeypatch):
config = WebSearchConfig(
enabled=True,
default_provider="tavily",
provider="tavily",
providers={
"tavily": WebSearchProviderConfig(
provider="tavily",
base_url="https://api.tavily.com",
api_key="key",
)
},
)
client = WebSearchClient(config)
async def fake_request_json(*args, **kwargs):
return {
"query": "Alem.Cloud",
"results": [
{
"title": "Alem.Cloud official",
"url": "https://example.test/alem",
"content": "Alem.Cloud is in Astana.",
"score": 0.9,
}
],
}
monkeypatch.setattr(client, "_request_json", fake_request_json)
results = await client.search("Alem.Cloud")
assert results[0].source_provider == "tavily"
assert results[0].url == "https://example.test/alem"
assert "Astana" in results[0].snippet
@pytest.mark.asyncio
async def test_searxng_adapter_allows_empty_api_key(monkeypatch):
config = WebSearchConfig(
enabled=True,
default_provider="searxng",
provider="searxng",
providers={
"searxng": WebSearchProviderConfig(
provider="searxng",
base_url="http://localhost:8080",
api_key="",
)
},
)
client = WebSearchClient(config)
async def fake_request_json(*args, **kwargs):
return {
"results": [
{
"title": "TAIPEI-1",
"url": "https://example.test/taipei",
"content": "TAIPEI-1 is in Taipei.",
"score": 2,
"engine": "duckduckgo",
}
],
}
monkeypatch.setattr(client, "_request_json", fake_request_json)
results = await client.search("TAIPEI-1")
assert results[0].source_provider == "searxng"
assert results[0].metadata["engine"] == "duckduckgo"
@pytest.mark.asyncio
async def test_location_search_evidence_returns_failure_on_empty_results(monkeypatch):
class EmptySearchClient:
async def search(self, *args, **kwargs):
return []
result = await collect_location_search_evidence(
web_search_client=EmptySearchClient(),
query=LocationQuery(name="TAIPEI-1", country="Taiwan"),
entity_type="compute_center",
)
assert result.evidence == []
assert "no usable" in result.failure_reason
@pytest.mark.asyncio
async def test_llm_location_fallback_skips_when_search_evidence_empty():
class ExplodingAIClient:
async def analyze(self, *_args, **_kwargs):
raise AssertionError("LLM should not be called without evidence")
result = await collect_llm_location_fallback_candidate(
provider_client=ExplodingAIClient(),
query=LocationQuery(name="TAIPEI-1", country="Taiwan"),
entity_type="compute_center",
search_evidence=[],
)
assert result.candidates == []
assert "no WebSearch evidence" in result.failure_reason
@pytest.mark.asyncio
async def test_credential_guide_keeps_default_without_search_evidence():
class EmptySearchClient:
async def search(self, *args, **kwargs):
return []
class ExplodingAIClient:
async def analyze(self, *_args, **_kwargs):
raise AssertionError("AI should not be called without search evidence")
async def fake_get_store(_db):
return None, {}
import app.services.credential_guides as credential_guides
original = credential_guides._get_guide_store
credential_guides._get_guide_store = fake_get_store
try:
guide = await generate_credential_guide(
object(),
"barentswatch",
ExplodingAIClient(),
EmptySearchClient(),
)
finally:
credential_guides._get_guide_store = original
assert guide["source"] == "default"
assert guide["verification_status"] == "unverified_no_search_evidence"
@pytest.mark.asyncio
async def test_credential_guide_uses_search_evidence(monkeypatch):
class SearchClient:
async def search(self, *args, **kwargs):
from app.services.ai_tools.schemas import SearchEvidence
return [
SearchEvidence(
title="Official docs",
url="https://docs.example.test",
snippet="Create an AIS client.",
source_provider="tavily",
)
]
class AIClient:
async def analyze(self, payload):
assert payload.context["search_evidence"]
return SimpleNamespace(content="## Generated\n\nSources included.")
saved = {}
async def fake_get_store(_db):
return None, saved
async def fake_save(db, provider, title, markdown, **metadata):
return {
"provider": provider,
"title": title,
"markdown": markdown,
"source": "ai",
**metadata,
}
import app.services.credential_guides as credential_guides
monkeypatch.setattr(credential_guides, "_get_guide_store", fake_get_store)
monkeypatch.setattr(credential_guides, "save_credential_guide", fake_save)
guide = await generate_credential_guide(
object(),
"barentswatch",
AIClient(),
SearchClient(),
)
assert guide["source"] == "ai"
assert guide["verification_status"] == "verified_with_search_evidence"
assert guide["sources"][0]["url"] == "https://docs.example.test"

View File

@@ -0,0 +1,46 @@
import pytest
from app.core.websocket.manager import ConnectionManager
class FakeWebSocket:
def __init__(self):
self.accepted = False
self.sent = []
self.closed = False
async def accept(self):
self.accepted = True
async def send_json(self, message):
self.sent.append(message)
async def close(self):
self.closed = True
@pytest.mark.asyncio
async def test_channel_subscribers_receive_channel_broadcasts():
manager = ConnectionManager()
socket = FakeWebSocket()
await manager.connect(socket, "user-1")
manager.subscribe(socket, ["dashboard"])
await manager.broadcast({"type": "data_frame", "channel": "dashboard"}, channel="dashboard")
assert socket.accepted is True
assert socket.sent == [{"type": "data_frame", "channel": "dashboard"}]
@pytest.mark.asyncio
async def test_disconnect_removes_channel_subscriptions():
manager = ConnectionManager()
socket = FakeWebSocket()
await manager.connect(socket, "user-1")
manager.subscribe(socket, ["dashboard"])
manager.disconnect(socket, "user-1")
await manager.broadcast({"type": "data_frame", "channel": "dashboard"}, channel="dashboard")
assert socket.sent == []
assert "dashboard" not in manager.channel_subscriptions

View File

@@ -86,6 +86,12 @@ CREATE TABLE collection_tasks (
id BIGSERIAL PRIMARY KEY,
datasource_id INTEGER NOT NULL REFERENCES data_sources(id) ON DELETE CASCADE,
status task_status NOT NULL DEFAULT 'pending',
phase VARCHAR(30) DEFAULT 'queued',
phase_progress FLOAT,
phase_message VARCHAR(255),
phase_current BIGINT,
phase_total BIGINT,
phase_unit VARCHAR(30),
started_at TIMESTAMP WITH TIME ZONE,
completed_at TIMESTAMP WITH TIME ZONE,
records_processed INTEGER DEFAULT 0,

View File

@@ -8,6 +8,9 @@ services:
args:
PYTHON_IMAGE: ${PYTHON_IMAGE:-python:3.14-slim}
UV_IMAGE: ${UV_IMAGE:-ghcr.io/astral-sh/uv:latest}
env_file:
- ./aiprovider/.env
- ${PLANET_AI_PROVIDER_RUNTIME_ENV_FILE:-./aiprovider/.env}
container_name: planet_aiprovider
ports:
- "8010:8010"

View File

@@ -10,6 +10,7 @@ services:
UV_IMAGE: ${UV_IMAGE:-ghcr.io/astral-sh/uv:latest}
env_file:
- ./aiprovider/.env
- ${PLANET_AI_PROVIDER_RUNTIME_ENV_FILE:-./aiprovider/.env}
container_name: planet_aiprovider
ports:
- "8010:8010"

View File

@@ -8,6 +8,209 @@ This project follows the repository versioning rule:
- `improvement` -> `+0.0.1`bugfix + 小功能混合)
- `bugfix` -> `+0.0.1`
## [0.51.1] — 2026-05-11
Released: 2026-05-11
### Highlights
- 修复 Earth 静态资源引用方式,图标和国家边界数据改为模块相对 URL避免部署路径变化时资源加载失败。
- 将 Material Symbols Rounded 字体切换为本地资源,减少 Earth 页面首屏对外部字体服务的依赖。
### Fixed
- 修复 BGP 广播图标、算力中心图标和国家边界 GeoJSON 在非固定 `/earth` 路径下可能失效的问题。
---
## [0.51.0] — 2026-05-11
Released: 2026-05-11
### ✨ Highlights
- 新增 AI Settings 控制台页面与 `backend/app/services/ai_tools/` 工具层,串通 Web Search Provider 与轻量 Agent orchestrator。
- 重写 Earth 算力中心候选「预览 / 保存」交互:单一委托 click + 内存 candidate Map新增空心呼吸圈预览保存后即时生成正式图标后台刷新失败不再误报为保存失败。
- 重写动作捕捉 zoom 识别mirror-safe 的 trend + pose hold 双通道,张开/合拢手势直接对应 zoom_in/out 并支持持续触发;单臂 rotate 仅在另一只手明确静止时才允许。
### Improvements
- 同步中英文 `earth-frontend-context.md``frontend-admin-frontend-context.md``faq.md``manual.md``quickstart.md`
- Earth 模块多处优化bgp-cruise-adapter、interactable、satellites、presentation-controller、controls 调整与回归测试补全。
---
## [0.50.0] — 2026-05-10
Released: 2026-05-10
### ✨ Highlights
- 新增 Earth 动作捕捉双通道控制Browser Camera 本地识别与 Motion Agent WebSocket 高级接入,并补齐调试 HUD、骨架预览和手势冷却保护。
- 新增 Motion 目标展示的 `PresentationController` 接入,动捕聚焦复用巡航卡片和 connector同时保持 BGP/News 原巡航体验不变。
- 扩展位置候选管线与 AI Provider 兜底,支持算力中心和 BGP 观测站候选采集、保存、待定位队列与 LLM factcheck。
### Added / Fixed / Improved
- 改进 `planet.sh`:支持可选 Motion Agent 启动、摄像头 index/URL 参数、WSL 摄像头引导、端口清理细化和 AI Provider/Motion 依赖自动处理。
- Settings 与 Playground 支持多 provider AI 配置、密钥来源脱敏预览和运行时默认 provider 解析。
- Docs 新增 FAQ 入口并同步中英文手册、Earth 前端上下文、位置管线和启动脚本文档。
- Earth 媒体面板记录直播/新闻 tab 状态,刷新后恢复用户上次选择。
---
## [0.49.0] — 2026-05-08
Released: 2026-05-08
### ✨ Features
- 新增统一地理位置解析 Pipeline支持 SourceCoordinates / Nominatim / Registry / Inherit 多策略链式 resolver。
- 新增 BGP 采集站与算力中心地理定位服务(`bgp_collector_locations``compute_center_locations``bgp_event_locations`)。
- 新增 Docs Gatekeeper 带鉴权文档 API`/api/v1/docs`),按用户权限动态返回文档目录与内容。
- 新增 Earth 全球新闻栏(`/api/v1/news/earth-feed`),根据地球视角坐标推断地区并聚合多源 RSS 信息流。
- Earth 新增 Mobile 算力中心国家高亮(`mobile-center-country-highlight.js`)。
---
## [0.48.0] — 2026-05-07
Released: 2026-05-07
### ✨ Highlights
- 自定义数据源新增 REST / WebSocket 映射运行时,并提供本地 AIS mock WebSocket用于实时船只 upsert 链路验证。
- AIS 原始观测、聚合策略、字段来源、冲突记录与船舶 enrichment 继续完善Earth 船只实时展示链路更接近生产数据形态。
- Earth 全球态势 summary 改为轻量 SQL 聚合,并在卫星 current 异常时回退到最近有效 TLE 批次,避免统计接口被大规模明细读取拖慢。
### 🔧 Improvements
- 修复 `/geo/summary``/geo/satellites` 在大表下加载慢或超时的问题,并补充 `collected_data` 与 AIS raw 相关索引。
- WebSocket 管理器支持匿名连接、频道订阅清理和更稳的连接生命周期测试,前端 WebSocket candidates / fallback 更可靠。
- `planet.sh` 强化端口释放、端口诊断和前端启动流程mock AIS server 提供 Bun 脚本入口。
---
## [0.47.0] — 2026-04-30
Released: 2026-04-30
### ✨ Highlights
- 新增 AISStream WebSocket 船只采集器,并将 AIS 多源数据写入原始观测层,由聚合接口统一去重、合并和解释字段来源。
- 设置页新增 AISStream API Key、采集范围 preset、运行状态、连接验证和凭证教程入口让全球 AIS 采集链路可配置、可观察。
- Earth 船只图层默认不再限制 5000 艘,并统一 marker 颜色、详情卡、hover 和搜索结果的船型归一化显示。
### 🔧 Improvements
- 聚合接口新增 `field_sources``selected_reasons``source_summary``quality_flags` 和冲突记录调试接口,动态字段默认优先采用更新的实时流观测。
- AISStream 标准化支持 `MetaData.ShipName` 船名兜底,并将 AIS 数字船型映射为 Cargo / Tanker / Passenger / Fishing / Military。
- 将仓库 docs 技能改为通用文档工作流Planet 专属白名单、双语、裸文件标题和凭证教程规则迁移到 `docs/documentation-coverage-rules.md`
- 更新 AIS v4/v5 TODO 与计划文档,明确后续聚合策略配置、船舶资料 enrichment 和媒体缓存边界。
---
## [0.46.3] — 2026-04-30
Released: 2026-04-30
### 🐛 Fixes
- 优化 Starlink footprint 显示后的地球拖拽性能,避免旋转地球时每帧重建 footprint 大网格,同时保持现有视觉效果不变。
- 恢复点击线缆后的呼吸透明度动画,让 locked / hover 线缆重新使用既有 pulse 配置。
---
## [0.46.2] — 2026-04-30
Released: 2026-04-30
### 🐛 Fixes
- 修复 Earth 启动时高清材质、云图和图层可见性绕过 `startupPriority` 的问题,统一由启动队列按文档顺序加载。
- 修复保存为关闭的高清材质/图层仍会先加载再关闭的问题,并保持海陆基座作为国界线图层的常驻底图。
- 修复搜索跳转会误关媒体面板、船只轨迹末端不贴合当前船只、Iridium footprint 被地表层遮挡等 Earth 交互问题。
### 📝 Documentation
- 更新 Earth 图层顺序、样式参考、使用手册和 AIS 聚合计划,补齐中英文说明与后续接入策略。
---
## [0.46.1] — 2026-04-30
Released: 2026-04-30
### 🐛 Fixes
- 修复新增 technical docs 文件存在但未进入 Docs 前端白名单时,侧栏不显示且 Markdown 链接无法解析到 `/docs/<slug>` 的问题。
- 补齐数据源/采集器连接验证与 Earth Interactable 使用说明的英文文档,保证公开 Docs 切换 EN 时同名页面可访问。
- 清理中英文 technical docs 中裸 `.md` 文件名链接标题,改为面向读者的语义标题。
### 📝 Documentation
- 将 Docs 前端白名单、公开文档双语配对、裸文件名链接标题三项检查写入 Claude 与 Codex 的 docs 技能流程。
---
## [0.46.0] — 2026-04-30
Released: 2026-04-30
### ✨ Highlights
- Earth 新增通用 Interactable 图标层船只、算力中心、BGP 事件与观测站统一使用批量 Points、屏幕拾取、状态 glow 和状态缩放。
- BGP 事件保留向外扩散圈,观测站保留雷达扫描层,并与 Interactable 主图标解耦到稳定的地表渲染层级。
- 登陆点回归黄色球形 Sprite贴近海缆层级并保持更稳定的地表显示和遮挡表现。
### 🔧 Improvements
- 新增 SVG asset 到 canvas texture 的 Interactable 资产加载路径,支持统一图标资源、缓存和可选染色。
- 同坐标 Interactable 自动做地表切向避让,降低重叠物件无法选择的问题。
- 优化 Earth toolbar 初始尺寸注入,避免首次显示原始尺寸后再跳到缩放尺寸。
- 补充 Interactable 计划、使用说明、图层顺序和 Earth 前端上下文文档。
- 修复船只 hover/locked 状态仅发光但放大反馈不明显的问题,将已有状态缩放接入通用图标层。
---
## [0.45.0] — 2026-04-29
### ✨ Highlights
- 采集任务新增阶段级量化进度,`fetching` 可展示百分比、阶段说明和字节下载量。
- AI Provider 启动链路支持从 `aiprovider/.env``~/.zshrc` 注入运行期配置,并避免密钥/模型变化触发镜像重建。
- AI Provider Docker build context 收敛到服务必需文件,`uv sync` 接入 BuildKit 缓存以减少重复下载。
### 🔧 Improvements
- IPtoASN、OpenGeoFeed、NRO delegated 下载型采集器接入真实字节进度上报。
- 数据源页、采集中任务弹窗和任务历史页展示阶段摘要,并在 tooltip 中保留完整进度细节。
- 调整 Earth 船只默认高度偏移,进一步贴近地表展示。
---
## [0.44.2] — 2026-04-29
### 📝 Documentation
- 补充 Earth 船只图层技术文档,记录分桶 `THREE.Points` 批量渲染、同尺寸交互 overlay 和屏幕空间 picking 的设计约束。
- 同步 Earth 渲染图层顺序和样式参考,明确 AIS 船只 renderOrder、depthTest、图标尺寸、航向分桶与 hover 命中半径。
- 更新船只渲染性能计划状态,标注 `0.44.1` 已落地的实现与后续全球 AIS / LOD 演进方向。
---
## [0.44.1] — 2026-04-29
### 🐛 Fixes
- 修复 Earth 船只图层拖动不跟手的问题,将普通船只从独立 Sprite 切换为按航向分桶的批量 Points 渲染。
- 修正船只 hover/click 拾取错位,改为屏幕空间命中检测并在拖拽/惯性期间跳过 hover 拾取。
- 统一 AIS 船只普通态与交互态方向,并让 hover/locked glow 与普通图标保持同尺寸覆盖。
- 恢复船只深度测试并收敛默认图标尺寸,避免北部岛屿/冰面附近出现明显压盖陆地的视觉问题。
---
## [0.44.0] — 2026-04-29
### ✨ Highlights
- 重构数据源与采集器设置边界:数据源页回归目录和采集触发,采集器 endpoint、请求头、凭证与连接验证统一进入设置页
- BarentsWatch AIS 完整接入凭证解析、连接检查、默认教程、AI 生成教程和船只采集/可视化链路
- Earth 新增船只图例、缩放反馈胶囊、缩放感知拖拽灵敏度,并将船只渲染性能优化方案沉淀到 plans
- 仪表盘重启服务新增前端重启 action并让 runner 通过 `~/.zshrc` 继承本地环境变量
### 🔧 Improvements
- 采集器连接状态改为基于成功采集或手动连接校验 checksum 判断,避免只依赖前端样式状态
- 数据源页新增采集中任务标签和任务进度弹窗,内置与自定义数据源统一展示
- Earth 国界线进一步贴近地表,并补充船只图层渲染顺序、样式和用户手册说明
- docs skill 与 Claude/Codex 文档流程补齐技术文档和 plans 的职责边界
---
## [0.43.1] — 2026-04-28
### 🐛 Fixes
- 修正 `planet.sh` 在全量 restart 后启动 AI Provider 时的提示语义,避免把预期内未就绪描述成异常不健康
---
## [0.43.0] — 2026-04-28
### ✨ Highlights

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