release: bump version to 0.71.0
Some checks failed
ci / backend (push) Has been cancelled
ci / frontend (push) Has been cancelled
release / images (push) Has been cancelled
ci / delivery (push) Has been cancelled

This commit is contained in:
linkong
2026-06-11 16:47:24 +08:00
parent 8c204717cd
commit 899e3bce43
56 changed files with 4618 additions and 260 deletions

View File

@@ -6,7 +6,7 @@ from pathlib import Path
from typing import Any
from uuid import uuid4
from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile, status
from fastapi import APIRouter, Depends, File, Form, HTTPException, Query, Request, UploadFile, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pydantic import BaseModel, Field
from sqlalchemy import delete, func, select, text
@@ -27,6 +27,20 @@ from app.services.earth_news import (
save_earth_news_sources_payload,
test_news_source_config,
)
from app.services.earth_news_manual import (
broadcast_manual_news_changed,
create_manual_news_group,
delete_manual_news_item,
get_news_record_or_404,
import_manual_news_items,
list_news_groups,
list_news_records,
parse_manual_news_import_upload,
rename_manual_news_group,
reprocess_manual_news_item,
serialize_news_record,
upsert_manual_news_item,
)
from app.services.earth_boundaries import (
EarthBoundaryBuildError,
get_boundary_build_status,
@@ -119,6 +133,26 @@ class EarthNewsSourceTestPayload(BaseModel):
source: dict[str, Any] = Field(default_factory=dict)
class EarthNewsManualItemPayload(BaseModel):
title: str = Field(default="", max_length=500)
summary: str = Field(default="", max_length=1200)
content: str = Field(default="", max_length=12000)
url: str = Field(default="", max_length=2000)
source: str = Field(default="", max_length=255)
region: str = Field(default="global", max_length=80)
published_at: str | None = None
category: str = Field(default="other", max_length=80)
tags: list[str] = Field(default_factory=list)
location: dict[str, Any] | None = None
homepage_url: str = Field(default="", max_length=2000)
content_language: str = Field(default="", max_length=32)
group_id: str | None = Field(default=None, max_length=120)
class EarthNewsManualGroupPayload(BaseModel):
name: str = Field(default="", max_length=120)
def _normalize_earth_brand_payload(payload: dict[str, Any] | None) -> dict[str, str]:
merged = DEFAULT_EARTH_BRAND.copy()
if payload:
@@ -375,6 +409,162 @@ async def test_earth_news_source(
return await test_news_source_config(payload.source, db=db)
@router.get("/news-groups")
async def list_earth_news_groups_admin(
_current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await list_news_groups(db)
@router.post("/news-groups")
async def create_earth_news_group_admin(
payload: EarthNewsManualGroupPayload,
_current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
group = await create_manual_news_group(db, payload.name)
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
await db.commit()
return {"status": "ok", "group": group}
@router.put("/news-groups/{group_id:path}")
async def rename_earth_news_group_admin(
group_id: str,
payload: EarthNewsManualGroupPayload,
_current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
group = await rename_manual_news_group(db, group_id, payload.name)
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
await db.commit()
await broadcast_manual_news_changed()
return {"status": "ok", "group": group}
@router.get("/news-items")
async def list_earth_news_items_admin(
page: int = Query(1, ge=1),
page_size: int = Query(50, ge=1, le=100),
source_type: str | None = Query(None),
region: str | None = Query(None),
category: str | None = Query(None),
status_filter: str | None = Query(None, alias="status"),
group_id: str | None = Query(None),
_current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await list_news_records(
db,
page=page,
page_size=page_size,
source_type=source_type,
region=region,
category=category,
status_filter=status_filter,
group_id=group_id,
)
@router.post("/news-items")
async def create_earth_news_item_admin(
payload: EarthNewsManualItemPayload,
_current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
result = await upsert_manual_news_item(db, payload.model_dump())
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
except PermissionError as exc:
raise HTTPException(status_code=403, detail=str(exc)) from exc
await db.commit()
await broadcast_manual_news_changed()
return {"status": "ok", "created": result.created, "queued": result.queued, "item": serialize_news_record(result.item)}
@router.post("/news-items/import")
async def import_earth_news_items_admin(
file: UploadFile = File(...),
group_id: str | None = Form(default=None),
_current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
payload = await parse_manual_news_import_upload(await file.read())
result = await import_manual_news_items(db, payload, group_id=group_id)
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
await db.commit()
await broadcast_manual_news_changed()
return {"status": "ok", **result}
@router.put("/news-items/{item_id:path}")
async def update_earth_news_item_admin(
item_id: str,
payload: EarthNewsManualItemPayload,
_current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
existing = await get_news_record_or_404(db, item_id)
if existing is None:
raise HTTPException(status_code=404, detail="News item not found.")
try:
result = await upsert_manual_news_item(
db,
payload.model_dump(),
item_id_override=item_id,
)
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
except PermissionError as exc:
raise HTTPException(status_code=403, detail=str(exc)) from exc
await db.commit()
await broadcast_manual_news_changed()
return {"status": "ok", "created": result.created, "queued": result.queued, "item": serialize_news_record(result.item)}
@router.delete("/news-items/{item_id:path}")
async def delete_earth_news_item_admin(
item_id: str,
_current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
deleted = await delete_manual_news_item(db, item_id)
except PermissionError as exc:
raise HTTPException(status_code=403, detail=str(exc)) from exc
if not deleted:
raise HTTPException(status_code=404, detail="News item not found.")
await db.commit()
await broadcast_manual_news_changed()
return {"status": "deleted", "id": item_id}
@router.post("/news-items/{item_id:path}/reprocess")
async def reprocess_earth_news_item_admin(
item_id: str,
_current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
existing = await get_news_record_or_404(db, item_id)
if existing is None:
raise HTTPException(status_code=404, detail="News item not found.")
try:
queued = await reprocess_manual_news_item(db, item_id)
except PermissionError as exc:
raise HTTPException(status_code=403, detail=str(exc)) from exc
await db.commit()
await broadcast_manual_news_changed()
return {"status": "queued" if queued else "not_queued", "queued": queued, "id": item_id}
@router.get("/oobe-status")
async def get_earth_oobe_status(
current_user: User | None = Depends(_get_optional_current_user),

View File

@@ -66,6 +66,7 @@ class NewsSourceType(StrEnum):
ATOM = "atom"
AGGREGATED = "aggregated"
REFERENCE = "reference"
MANUAL = "manual"
class NewsEnrichmentStatus(StrEnum):

View File

@@ -2282,6 +2282,50 @@ def _filter_news_items_by_source_ids(
]
def _news_item_source_id(item: ParsedNewsItem) -> str:
return item.id.split(":", 1)[0] if ":" in item.id else ""
def _is_news_item_display_ready(item: ParsedNewsItem, *, locale: str) -> bool:
return bool(
_get_locale_text(item, "title", locale=locale)
and _get_locale_text(item, "summary", locale=locale)
)
def _diversify_news_items_for_locale(
items: list[ParsedNewsItem],
*,
active_region: str,
limit: int,
locale: str,
) -> list[ParsedNewsItem]:
ranked = sorted(
_rank_and_trim_items(items, active_region=active_region, limit=max(len(items), limit)),
key=lambda item: (not _is_news_item_display_ready(item, locale=locale),),
)
buckets: dict[str, list[ParsedNewsItem]] = {}
order: list[str] = []
for item in ranked:
key = _news_item_source_id(item) or item.source or item.feed_name or item.id
if key not in buckets:
buckets[key] = []
order.append(key)
buckets[key].append(item)
diversified: list[ParsedNewsItem] = []
while len(diversified) < limit and order:
next_order: list[str] = []
for source_id in order:
bucket = buckets.get(source_id) or []
if bucket and len(diversified) < limit:
diversified.append(bucket.pop(0))
if bucket:
next_order.append(source_id)
order = next_order
return diversified
async def _call_store_list_items(list_fn, db: AsyncSession, **kwargs):
try:
return await list_fn(db, **kwargs)
@@ -2717,6 +2761,37 @@ async def get_earth_news_payload(
categories=categories,
source_ids=source_ids,
)
if not source_ids:
ready_sources = {
_news_item_source_id(item)
for item in items
if _is_news_item_display_ready(item, locale=locale)
}
missing_ready_sources = [
source.id
for source in sources
if source.id and source.id not in ready_sources
]
if missing_ready_sources:
extra_items: list[ParsedNewsItem] = []
for missing_source_id in missing_ready_sources:
extra_items.extend(
await _call_store_list_items(
list_earth_news_items,
db,
active_region=active_region,
limit=3,
categories=categories,
source_ids={missing_source_id},
)
)
if extra_items:
items = _diversify_news_items_for_locale(
[*items, *extra_items],
active_region=active_region,
limit=limit,
locale=locale,
)
if hasattr(db, "execute"):
cruise_items = await _call_store_list_items(
list_earth_news_cruise_items,

View File

@@ -0,0 +1,693 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import UTC, datetime
import hashlib
import html
import json
import re
from typing import Any
from bs4 import BeautifulSoup
from sqlalchemy import delete, func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.enums import NewsEnrichmentStatus, NewsSourceType, NewsTaggingSource
from app.core.websocket.broadcaster import broadcaster
from app.models.earth_news import EarthNewsItem
from app.models.system_setting import SystemSetting
from app.services.earth_news import (
ALLOWED_NEWS_CATEGORY_KEYS,
DEFAULT_NEWS_LOCALE,
REGION_ANCHORS,
NewsFeedEndpoint,
NewsFeedSource,
NewsTargetLocation,
ParsedNewsItem,
apply_news_classification,
build_anchor_location_patch,
build_target_location_job_payload,
build_target_location_patch,
_serialize_item,
)
from app.services.earth_news_queue import enqueue_target_location_job
from app.services.earth_news_store import record_to_parsed_news_item
MANUAL_NEWS_SOURCE_ID = "manual"
MANUAL_NEWS_SOURCE_LABEL = "手动添加"
MANUAL_NEWS_MAX_IMPORT_ITEMS = 500
MANUAL_NEWS_MAX_TITLE_LENGTH = 500
MANUAL_NEWS_MAX_SUMMARY_LENGTH = 1200
MANUAL_NEWS_MAX_CONTENT_LENGTH = 12000
EARTH_NEWS_MANUAL_GROUPS_CATEGORY = "earth_news_manual_groups"
DEFAULT_MANUAL_NEWS_GROUP_ID = "manual-default"
DEFAULT_MANUAL_NEWS_GROUP_NAME = "新建新闻组"
@dataclass(frozen=True)
class ManualNewsWriteResult:
item: EarthNewsItem
created: bool
queued: bool
@dataclass(frozen=True)
class ManualNewsGroup:
id: str
name: str
sort_order: int = 0
def _clean_text(value: object, *, max_length: int) -> str:
raw = "" if value is None else str(value)
text = BeautifulSoup(html.unescape(raw), "html.parser").get_text(" ", strip=True)
text = re.sub(r"\s+", " ", text).strip()
if len(text) > max_length:
return text[: max_length - 1].rstrip() + ""
return text
def _parse_datetime(value: object) -> datetime | None:
if value is None or str(value).strip() == "":
return None
if isinstance(value, datetime):
parsed = value
else:
try:
parsed = datetime.fromisoformat(str(value).strip().replace("Z", "+00:00"))
except ValueError as exc:
raise ValueError("published_at 必须是 ISO8601 时间。") from exc
if parsed.tzinfo is None:
return parsed.replace(tzinfo=UTC)
return parsed.astimezone(UTC)
def _detect_language(*parts: str) -> str:
text = " ".join(part for part in parts if part)
cjk_count = len(re.findall(r"[\u4e00-\u9fff]", text))
latin_count = len(re.findall(r"[A-Za-z]", text))
return "zh-CN" if cjk_count >= max(4, latin_count // 3) else "en-US"
def _manual_item_id(*, title: str, published_at: datetime | None, url: str, source: str) -> str:
published = published_at.isoformat() if published_at else ""
basis = "\n".join([title.strip().lower(), published, url.strip().lower(), source.strip().lower()])
return f"manual:{hashlib.sha1(basis.encode('utf-8')).hexdigest()[:16]}"
def _manual_group_id(name: str) -> str:
basis = f"{name.strip().lower()}\n{datetime.now(UTC).isoformat()}"
return f"manual-group:{hashlib.sha1(basis.encode('utf-8')).hexdigest()[:10]}"
def _news_meta(record: EarthNewsItem) -> dict[str, Any]:
location_meta = record.location_meta if isinstance(record.location_meta, dict) else {}
news_meta = location_meta.get("news_meta")
return dict(news_meta) if isinstance(news_meta, dict) else {}
def _record_source_type(record: EarthNewsItem) -> str:
return str(_news_meta(record).get("feed_type") or _news_meta(record).get("source_type") or "rss")
def _record_manual_group_id(record: EarthNewsItem) -> str:
return str(_news_meta(record).get("manual_group_id") or DEFAULT_MANUAL_NEWS_GROUP_ID)
def _rss_group_id(record: EarthNewsItem) -> str:
basis = "\n".join(
[
_record_source_type(record),
str(record.feed_name or ""),
str(record.source or ""),
]
)
return f"rss:{hashlib.sha1(basis.encode('utf-8')).hexdigest()[:12]}"
def _default_manual_group() -> dict[str, Any]:
return {
"id": DEFAULT_MANUAL_NEWS_GROUP_ID,
"name": DEFAULT_MANUAL_NEWS_GROUP_NAME,
"sort_order": 0,
}
def _normalize_manual_groups_payload(payload: Any) -> list[dict[str, Any]]:
raw_groups = payload.get("groups") if isinstance(payload, dict) else None
normalized: list[dict[str, Any]] = []
seen: set[str] = set()
for index, item in enumerate(raw_groups if isinstance(raw_groups, list) else []):
if not isinstance(item, dict):
continue
group_id = str(item.get("id") or "").strip()
name = _clean_text(item.get("name"), max_length=120)
if not group_id or not name or group_id in seen:
continue
normalized.append(
{
"id": group_id,
"name": name,
"sort_order": int(item.get("sort_order") or index),
}
)
seen.add(group_id)
if DEFAULT_MANUAL_NEWS_GROUP_ID not in seen:
normalized.insert(0, _default_manual_group())
return sorted(normalized, key=lambda item: (int(item.get("sort_order") or 0), str(item.get("name") or "")))
async def _get_manual_groups_record(db: AsyncSession) -> SystemSetting | None:
result = await db.execute(
select(SystemSetting).where(SystemSetting.category == EARTH_NEWS_MANUAL_GROUPS_CATEGORY)
)
return result.scalar_one_or_none()
async def get_manual_news_groups(db: AsyncSession) -> list[dict[str, Any]]:
record = await _get_manual_groups_record(db)
return _normalize_manual_groups_payload(record.payload if record else None)
async def _save_manual_news_groups(db: AsyncSession, groups: list[dict[str, Any]]) -> list[dict[str, Any]]:
normalized = _normalize_manual_groups_payload({"groups": groups})
record = await _get_manual_groups_record(db)
payload = {"groups": normalized}
if record is None:
db.add(SystemSetting(category=EARTH_NEWS_MANUAL_GROUPS_CATEGORY, payload=payload))
else:
record.payload = payload
await db.flush()
return normalized
async def resolve_manual_news_group(db: AsyncSession, group_id: str | None) -> ManualNewsGroup:
normalized_id = str(group_id or DEFAULT_MANUAL_NEWS_GROUP_ID).strip() or DEFAULT_MANUAL_NEWS_GROUP_ID
groups = await get_manual_news_groups(db)
match = next((item for item in groups if item.get("id") == normalized_id), None)
if match is None and normalized_id != DEFAULT_MANUAL_NEWS_GROUP_ID:
raise ValueError(f"手动新闻组不存在:{normalized_id}")
match = match or _default_manual_group()
return ManualNewsGroup(
id=str(match["id"]),
name=str(match["name"]),
sort_order=int(match.get("sort_order") or 0),
)
async def create_manual_news_group(db: AsyncSession, name: str) -> dict[str, Any]:
group_name = _clean_text(name, max_length=120)
if not group_name:
raise ValueError("新闻组名称不能为空。")
groups = await get_manual_news_groups(db)
group = {"id": _manual_group_id(group_name), "name": group_name, "sort_order": len(groups)}
groups.append(group)
await _save_manual_news_groups(db, groups)
return group
async def rename_manual_news_group(db: AsyncSession, group_id: str, name: str) -> dict[str, Any]:
group_name = _clean_text(name, max_length=120)
if not group_name:
raise ValueError("新闻组名称不能为空。")
groups = await get_manual_news_groups(db)
match = next((item for item in groups if item.get("id") == group_id), None)
if match is None:
raise ValueError(f"手动新闻组不存在:{group_id}")
match["name"] = group_name
await _save_manual_news_groups(db, groups)
result = await db.execute(
select(EarthNewsItem).where(
EarthNewsItem.location_meta.op("->")("news_meta").op("->>")("feed_type") == NewsSourceType.MANUAL.value
)
)
for record in result.scalars().all():
if _record_manual_group_id(record) != group_id:
continue
location_meta = dict(record.location_meta or {})
news_meta = dict(location_meta.get("news_meta") or {})
news_meta["manual_group_name"] = group_name
location_meta["news_meta"] = news_meta
record.location_meta = location_meta
await db.flush()
return match
def _normalize_region(value: object) -> str:
region = str(value or "global").strip().lower() or "global"
if region not in REGION_ANCHORS:
raise ValueError(f"region 不支持:{region}")
return region
def _normalize_tags(value: object) -> list[str]:
if value is None:
return []
if isinstance(value, str):
parts = re.split(r"[,\n]", value)
elif isinstance(value, list):
parts = [str(item) for item in value]
else:
raise ValueError("tags 必须是字符串数组或逗号分隔字符串。")
return [item.strip() for item in parts if item.strip()][:20]
def _normalize_location(value: object) -> NewsTargetLocation | None:
if value in (None, ""):
return None
if not isinstance(value, dict):
raise ValueError("location 必须是对象。")
lat = value.get("latitude")
lon = value.get("longitude")
if lat in (None, "") and lon in (None, ""):
return None
try:
latitude = float(lat)
longitude = float(lon)
except (TypeError, ValueError) as exc:
raise ValueError("location.latitude / longitude 必须是数字。") from exc
if not -90 <= latitude <= 90 or not -180 <= longitude <= 180:
raise ValueError("location 经纬度超出范围。")
label = _clean_text(value.get("label"), max_length=255)
if not label:
label = f"{latitude:.4f}, {longitude:.4f}"
return NewsTargetLocation(
latitude=latitude,
longitude=longitude,
label=label,
source="manual_location",
confidence=1.0,
country=_clean_text(value.get("country"), max_length=100) or None,
city=_clean_text(value.get("city"), max_length=100) or None,
)
def _manual_source(source_name: str, *, region: str) -> NewsFeedSource:
return NewsFeedSource(
id=MANUAL_NEWS_SOURCE_ID,
name=source_name or MANUAL_NEWS_SOURCE_LABEL,
region=region,
feed_url="",
homepage_url="",
source_type=NewsSourceType.MANUAL.value,
default_category="other",
source_tags=("manual",),
)
def _manual_feed(category: str) -> NewsFeedEndpoint:
return NewsFeedEndpoint(
id=MANUAL_NEWS_SOURCE_ID,
name=MANUAL_NEWS_SOURCE_LABEL,
url="",
type=NewsSourceType.MANUAL.value,
default_category=category or "other",
tags=("manual",),
priority=1,
)
def parsed_manual_news_item(
payload: dict[str, Any],
*,
item_id_override: str | None = None,
) -> tuple[ParsedNewsItem, NewsTargetLocation | None, str]:
title = _clean_text(payload.get("title"), max_length=MANUAL_NEWS_MAX_TITLE_LENGTH)
if not title:
raise ValueError("title 不能为空。")
content = _clean_text(payload.get("content"), max_length=MANUAL_NEWS_MAX_CONTENT_LENGTH)
summary = _clean_text(payload.get("summary"), max_length=MANUAL_NEWS_MAX_SUMMARY_LENGTH)
if not summary:
summary = _clean_text(content, max_length=240) if content else title
source = _clean_text(payload.get("source"), max_length=255) or MANUAL_NEWS_SOURCE_LABEL
region = _normalize_region(payload.get("region"))
published_at = _parse_datetime(payload.get("published_at")) or datetime.now(UTC)
url = str(payload.get("url") or "").strip()
category = str(payload.get("category") or "other").strip().lower() or "other"
if category not in ALLOWED_NEWS_CATEGORY_KEYS:
raise ValueError(f"category 不支持:{category}")
tags = _normalize_tags(payload.get("tags"))
target = _normalize_location(payload.get("location"))
language = str(payload.get("content_language") or "").strip() or _detect_language(title, summary, content)
localizations = {
language: {
"title": title,
"summary": summary,
}
}
item = ParsedNewsItem(
id=item_id_override
or _manual_item_id(title=title, published_at=published_at, url=url, source=source),
title=title,
summary=summary,
url=url,
source=source,
feed_name=MANUAL_NEWS_SOURCE_LABEL,
feed_region=region,
homepage_url=str(payload.get("homepage_url") or ""),
published_at=published_at,
content_language=language,
localizations=localizations,
enrichment_status=NewsEnrichmentStatus.PENDING.value,
source_tags=["manual"],
feed_id=MANUAL_NEWS_SOURCE_ID,
feed_type=NewsSourceType.MANUAL.value,
feed_default_category=category,
category=category,
item_tags=tags,
tagging_source=NewsTaggingSource.MANUAL.value if payload.get("category") else NewsTaggingSource.RULES.value,
tagging_confidence=0.9 if payload.get("category") else 0.0,
)
source_config = _manual_source(source, region=region)
feed = _manual_feed(category)
apply_news_classification(item, source_config, feed=feed)
if payload.get("category"):
item.category = category
item.tagging_source = NewsTaggingSource.MANUAL.value
item.tagging_confidence = 0.9
if tags:
item.item_tags = sorted(set([*item.item_tags, *tags]))
return item, target, content
def _manual_editable(record: EarthNewsItem) -> bool:
if record.id.startswith("manual:"):
return True
news_meta = (record.location_meta or {}).get("news_meta") if isinstance(record.location_meta, dict) else None
return isinstance(news_meta, dict) and news_meta.get("feed_type") == NewsSourceType.MANUAL.value
async def _broadcast_news_reload() -> None:
await broadcaster.broadcast_earth_update(
{
"action": "database_changed",
"source": "earth_news_items",
"layers": ["news"],
"refresh_strategy": "reload",
}
)
async def upsert_manual_news_item(
db: AsyncSession,
payload: dict[str, Any],
*,
item_id_override: str | None = None,
group_id: str | None = None,
) -> ManualNewsWriteResult:
item, target, content = parsed_manual_news_item(payload, item_id_override=item_id_override)
group = await resolve_manual_news_group(db, group_id or payload.get("group_id"))
existing = await db.get(EarthNewsItem, item.id)
created = existing is None
patch = build_target_location_patch(item, target) if target else build_anchor_location_patch(item)
patch_meta = dict(patch.get("location_meta") or {})
patch_news_meta = dict(patch_meta.get("news_meta") or {})
patch_news_meta["feed_type"] = NewsSourceType.MANUAL.value
patch_news_meta["source_type"] = NewsSourceType.MANUAL.value
patch_news_meta["manual_group_id"] = group.id
patch_news_meta["manual_group_name"] = group.name
patch_meta["news_meta"] = patch_news_meta
patch["location_meta"] = patch_meta
now = datetime.now(UTC)
record = existing or EarthNewsItem(
id=item.id,
title=item.title,
summary=item.summary,
content_language=item.content_language,
localizations=dict(item.localizations or {}),
url=item.url,
source=item.source,
feed_name=item.feed_name,
region=item.feed_region,
homepage_url=item.homepage_url,
published_at=item.published_at,
latitude=patch["latitude"],
longitude=patch["longitude"],
location_label=patch["location_label"],
location_source=patch["location_source"],
verified=patch["verified"],
location_meta=patch["location_meta"],
first_seen_at=now,
last_seen_at=now,
resolved_at=now if patch["verified"] else None,
enrichment_status=item.enrichment_status,
)
if existing is None:
db.add(record)
else:
if not _manual_editable(record):
raise PermissionError("RSS 新闻不允许通过手动新闻接口编辑。")
record.title = item.title
record.summary = item.summary
record.content_language = item.content_language
record.localizations = dict(item.localizations or {})
record.url = item.url
record.source = item.source
record.feed_name = item.feed_name
record.region = item.feed_region
record.homepage_url = item.homepage_url
record.published_at = item.published_at
record.last_seen_at = now
if target is None and record.location_source == "manual_location":
merged_meta = dict(record.location_meta or {})
patch_meta = patch.get("location_meta") if isinstance(patch, dict) else None
patch_news_meta = patch_meta.get("news_meta") if isinstance(patch_meta, dict) else None
if isinstance(patch_news_meta, dict):
merged_meta["news_meta"] = patch_news_meta
record.location_meta = merged_meta
else:
record.location_meta = patch["location_meta"]
if target:
record.latitude = patch["latitude"]
record.longitude = patch["longitude"]
record.location_label = patch["location_label"]
record.location_source = patch["location_source"]
record.verified = patch["verified"]
record.resolved_at = now
elif record.location_source != "manual_location":
record.latitude = patch["latitude"]
record.longitude = patch["longitude"]
record.location_label = patch["location_label"]
record.location_source = patch["location_source"]
record.verified = patch["verified"]
record.resolved_at = None
record.enrichment_status = NewsEnrichmentStatus.PENDING.value
record.enrichment_error = None
record.enriched_at = None
if content:
meta = dict(record.location_meta or {})
meta["manual_content"] = content
record.location_meta = meta
await db.flush()
queued = await enqueue_target_location_job(build_target_location_job_payload(item), force=True)
if queued:
record.enrichment_status = NewsEnrichmentStatus.QUEUED.value
await db.flush()
return ManualNewsWriteResult(item=record, created=created, queued=queued)
async def import_manual_news_items(
db: AsyncSession,
payload: list[Any],
*,
group_id: str | None = None,
) -> dict[str, Any]:
if len(payload) > MANUAL_NEWS_MAX_IMPORT_ITEMS:
raise ValueError(f"单次最多导入 {MANUAL_NEWS_MAX_IMPORT_ITEMS} 条。")
created = 0
updated = 0
queued = 0
errors: list[dict[str, Any]] = []
for index, raw_item in enumerate(payload):
if not isinstance(raw_item, dict):
errors.append({"index": index, "error": "条目必须是 JSON 对象。"})
continue
try:
result = await upsert_manual_news_item(db, raw_item, group_id=group_id)
created += 1 if result.created else 0
updated += 0 if result.created else 1
queued += 1 if result.queued else 0
except Exception as exc:
errors.append({"index": index, "error": str(exc)})
if errors and created == 0 and updated == 0:
raise ValueError("导入失败,未写入任何新闻。")
return {"created": created, "updated": updated, "queued": queued, "failed": len(errors), "errors": errors}
async def parse_manual_news_import_upload(raw_bytes: bytes) -> list[Any]:
try:
payload = json.loads(raw_bytes.decode("utf-8-sig"))
except UnicodeDecodeError as exc:
raise ValueError("JSON 文件必须使用 UTF-8 编码。") from exc
except json.JSONDecodeError as exc:
raise ValueError(f"JSON 解析失败:第 {exc.lineno} 行第 {exc.colno} 列。") from exc
if not isinstance(payload, list):
raise ValueError("JSON 顶层必须是数组。")
return payload
def serialize_news_record(record: EarthNewsItem, *, locale: str = DEFAULT_NEWS_LOCALE) -> dict[str, Any]:
item = record_to_parsed_news_item(record)
payload = _serialize_item(item, active_region=item.feed_region, locale=locale)
news_meta = _news_meta(record)
payload["editable"] = _manual_editable(record)
payload["source_type"] = payload.get("feed_type")
payload["status"] = record.enrichment_status
payload["translated"] = bool((record.localizations or {}).get("zh-CN") and (record.localizations or {}).get("en-US"))
payload["manual_content"] = (record.location_meta or {}).get("manual_content") if isinstance(record.location_meta, dict) else None
payload["manual_group_id"] = news_meta.get("manual_group_id")
payload["manual_group_name"] = news_meta.get("manual_group_name")
return payload
def _record_matches_group(record: EarthNewsItem, group_id: str) -> bool:
source_type = _record_source_type(record)
if source_type == NewsSourceType.MANUAL.value:
return _record_manual_group_id(record) == group_id
return _rss_group_id(record) == group_id
async def list_news_records(
db: AsyncSession,
*,
page: int,
page_size: int,
source_type: str | None = None,
region: str | None = None,
category: str | None = None,
status_filter: str | None = None,
group_id: str | None = None,
) -> dict[str, Any]:
page = max(page, 1)
page_size = min(max(page_size, 1), 100)
query = select(EarthNewsItem)
count_query = select(func.count(EarthNewsItem.id))
filters = []
if region and region != "all":
filters.append(EarthNewsItem.region == region)
if status_filter and status_filter != "all":
filters.append(EarthNewsItem.enrichment_status == status_filter)
if source_type and source_type != "all":
filters.append(EarthNewsItem.location_meta.op("->")("news_meta").op("->>")("feed_type") == source_type)
if category and category != "all":
filters.append(EarthNewsItem.location_meta.op("->")("news_meta").op("->>")("category") == category)
for clause in filters:
query = query.where(clause)
count_query = count_query.where(clause)
ordered_query = query.order_by(EarthNewsItem.published_at.desc().nullslast(), EarthNewsItem.last_seen_at.desc())
if group_id:
result = await db.execute(ordered_query)
all_records = [record for record in result.scalars().all() if _record_matches_group(record, group_id)]
total = len(all_records)
records = all_records[(page - 1) * page_size : page * page_size]
else:
total_result = await db.execute(count_query)
result = await db.execute(
ordered_query.offset((page - 1) * page_size).limit(page_size)
)
records = list(result.scalars().all())
total = int(total_result.scalar() or 0)
return {
"items": [serialize_news_record(record) for record in records],
"page": page,
"page_size": page_size,
"total": total,
}
async def list_news_groups(db: AsyncSession, *, locale: str = DEFAULT_NEWS_LOCALE) -> dict[str, Any]:
manual_groups = await get_manual_news_groups(db)
manual_by_id: dict[str, dict[str, Any]] = {
str(group["id"]): {
"id": str(group["id"]),
"name": str(group["name"]),
"group_type": "manual",
"source_type": NewsSourceType.MANUAL.value,
"editable": True,
"sort_order": int(group.get("sort_order") or 0),
"count": 0,
"items": [],
}
for group in manual_groups
}
rss_by_id: dict[str, dict[str, Any]] = {}
result = await db.execute(
select(EarthNewsItem).order_by(EarthNewsItem.published_at.desc().nullslast(), EarthNewsItem.last_seen_at.desc())
)
for record in result.scalars().all():
serialized = serialize_news_record(record, locale=locale)
source_type = _record_source_type(record)
if source_type == NewsSourceType.MANUAL.value:
group_id = _record_manual_group_id(record)
group = manual_by_id.setdefault(
group_id,
{
"id": group_id,
"name": str(_news_meta(record).get("manual_group_name") or DEFAULT_MANUAL_NEWS_GROUP_NAME),
"group_type": "manual",
"source_type": NewsSourceType.MANUAL.value,
"editable": True,
"sort_order": len(manual_by_id),
"count": 0,
"items": [],
},
)
else:
group_id = _rss_group_id(record)
group = rss_by_id.setdefault(
group_id,
{
"id": group_id,
"name": record.feed_name or record.source or "RSS 新闻",
"group_type": "rss",
"source_type": source_type,
"editable": False,
"region": record.region,
"source": record.source,
"feed_name": record.feed_name,
"count": 0,
"items": [],
},
)
group["count"] = int(group.get("count") or 0) + 1
group.setdefault("items", []).append(serialized)
manual_items = sorted(manual_by_id.values(), key=lambda item: (int(item.get("sort_order") or 0), str(item.get("name") or "")))
rss_items = sorted(rss_by_id.values(), key=lambda item: str(item.get("name") or ""))
return {"groups": [*manual_items, *rss_items], "manual_groups": manual_items, "rss_groups": rss_items}
async def get_news_record_or_404(db: AsyncSession, item_id: str) -> EarthNewsItem | None:
return await db.get(EarthNewsItem, item_id)
async def delete_manual_news_item(db: AsyncSession, item_id: str) -> bool:
record = await db.get(EarthNewsItem, item_id)
if record is None:
return False
if not _manual_editable(record):
raise PermissionError("RSS 新闻不允许通过手动新闻接口删除。")
await db.execute(delete(EarthNewsItem).where(EarthNewsItem.id == item_id))
await db.flush()
return True
async def reprocess_manual_news_item(db: AsyncSession, item_id: str) -> bool:
record = await db.get(EarthNewsItem, item_id)
if record is None:
return False
if not _manual_editable(record):
raise PermissionError("RSS 新闻不允许通过手动新闻接口重新处理。")
item = record_to_parsed_news_item(record)
queued = await enqueue_target_location_job(build_target_location_job_payload(item), force=True)
if queued:
record.enrichment_status = NewsEnrichmentStatus.QUEUED.value
record.enrichment_error = None
await db.flush()
return queued
async def broadcast_manual_news_changed() -> None:
await _broadcast_news_reload()

View File

@@ -133,42 +133,6 @@ def _source_filter_clause(source_ids: set[str] | None):
return EarthNewsItem.location_meta.op("->")("news_meta").op("->>")("source_id").in_(sorted(source_ids))
def _record_source_id(record: EarthNewsItem) -> str:
location_meta = dict(record.location_meta or {})
news_meta = location_meta.get("news_meta") if isinstance(location_meta.get("news_meta"), dict) else {}
source_id = str(news_meta.get("source_id") or "").strip()
if source_id:
return source_id
if isinstance(record.id, str) and ":" in record.id:
return record.id.split(":", 1)[0]
return record.feed_name or record.source or record.id
def _diversify_records_by_source(records: list[EarthNewsItem], *, limit: int) -> list[EarthNewsItem]:
if limit <= 0 or len(records) <= limit:
return records[:limit]
buckets: dict[str, list[EarthNewsItem]] = {}
order: list[str] = []
for record in records:
source_id = _record_source_id(record)
if source_id not in buckets:
buckets[source_id] = []
order.append(source_id)
buckets[source_id].append(record)
diversified: list[EarthNewsItem] = []
while len(diversified) < limit and order:
next_order: list[str] = []
for source_id in order:
bucket = buckets.get(source_id) or []
if bucket and len(diversified) < limit:
diversified.append(bucket.pop(0))
if bucket:
next_order.append(source_id)
order = next_order
return diversified
async def list_earth_news_items(
db: AsyncSession,
*,
@@ -177,7 +141,7 @@ async def list_earth_news_items(
categories: set[str] | None = None,
source_ids: set[str] | None = None,
) -> list[ParsedNewsItem]:
query_limit = limit if source_ids else min(max(limit * 8, limit), 200)
query_limit = limit if source_ids else min(max(limit * 20, limit), 500)
query = (
select(EarthNewsItem)
.order_by(*_query_sort_key(active_region))
@@ -382,13 +346,28 @@ async def update_earth_news_item_enrichment(
if record is None:
return False
if "latitude" in patch:
record.latitude = float(patch["latitude"])
record.longitude = float(patch["longitude"])
record.location_label = str(patch["location_label"])
record.location_source = str(patch["location_source"])
record.verified = bool(patch["verified"])
record.location_meta = dict(patch.get("location_meta") or {})
record.resolved_at = datetime.now(UTC) if record.verified else None
patch_meta = dict(patch.get("location_meta") or {})
if record.location_source == "manual_location":
current_meta = dict(record.location_meta or {})
patch_news_meta = patch_meta.get("news_meta")
if isinstance(patch_news_meta, dict):
current_meta["news_meta"] = patch_news_meta
current_meta["manual_enrichment"] = {
"resolution_stage": patch_meta.get("resolution_stage"),
"ai_attempted": patch_meta.get("ai_attempted"),
"ai_status": patch_meta.get("ai_status"),
"ai_error": patch_meta.get("ai_error"),
"debug_note": patch_meta.get("debug_note"),
}
record.location_meta = current_meta
else:
record.latitude = float(patch["latitude"])
record.longitude = float(patch["longitude"])
record.location_label = str(patch["location_label"])
record.location_source = str(patch["location_source"])
record.verified = bool(patch["verified"])
record.location_meta = patch_meta
record.resolved_at = datetime.now(UTC) if record.verified else None
if "content_language" in patch:
record.content_language = str(patch.get("content_language") or "en")
if "localizations" in patch:

View File

@@ -1,4 +1,5 @@
[pytest]
pythonpath = ..
asyncio_mode = auto
testpaths = tests
python_files = test_*.py

View File

@@ -12,6 +12,7 @@ from app.services.earth_news import (
default_earth_news_sources_payload,
normalize_earth_news_sources_payload,
_fetch_source,
_diversify_news_items_for_locale,
_enrich_items_with_target_locations,
_extract_target_location_from_text,
_parse_feed_entries,
@@ -130,6 +131,47 @@ def test_rank_and_trim_items_prioritizes_active_breaking():
assert [item.id for item in ranked] == ["global:critical", "europe:expired", "europe:regular"]
def test_diversify_news_items_prefers_display_ready_content_across_sources():
published_at = datetime(2026, 6, 11, 3, 0, tzinfo=UTC)
def make_item(source_id: str, suffix: str, *, zh_ready: bool) -> ParsedNewsItem:
return ParsedNewsItem(
id=f"{source_id}:{suffix}",
title=f"{source_id} title {suffix}",
summary=f"{source_id} summary {suffix}",
url=f"https://example.com/{source_id}/{suffix}",
source=source_id,
feed_name=source_id,
feed_region="global",
homepage_url="https://example.com",
published_at=published_at,
content_language="en",
localizations={
"zh-CN": {
"title": f"{source_id} 中文标题 {suffix}",
"summary": f"{source_id} 中文摘要 {suffix}",
}
} if zh_ready else {},
)
items = [
make_item("source-a", "1", zh_ready=False),
make_item("source-a", "2", zh_ready=False),
make_item("source-a", "3", zh_ready=False),
make_item("source-b", "1", zh_ready=True),
make_item("source-c", "1", zh_ready=True),
]
result = _diversify_news_items_for_locale(
items,
active_region="global",
limit=3,
locale="zh-CN",
)
assert [item.id.split(":", 1)[0] for item in result] == ["source-b", "source-c", "source-a"]
def test_serialize_item_falls_back_to_global_anchor():
item = ParsedNewsItem(
id="custom:test",
@@ -961,9 +1003,11 @@ async def test_earth_news_payload_uses_fresh_database_items_without_rss(monkeypa
async def fake_get_earth_news_freshness(_db, *, active_region):
return 12, datetime.now(UTC)
async def fake_list_earth_news_items(_db, *, active_region, limit, categories=None):
assert limit == 12
return [item]
async def fake_list_earth_news_items(_db, *, active_region, limit, categories=None, source_ids=None):
if source_ids is None:
assert limit == 12
return [item]
return []
async def fail_fetch(_sources):
raise AssertionError("fresh database items should not fetch RSS")
@@ -1062,8 +1106,8 @@ async def test_earth_news_payload_passes_region_and_category_filters_to_store(mo
async def fake_list_earth_news_items(_db, *, active_region, limit, categories=None, source_ids=None):
captured["items_region"] = active_region
captured["items_categories"] = categories
captured["items_source_ids"] = source_ids
return [item]
captured.setdefault("items_source_ids", []).append(source_ids)
return [item] if source_ids is None else []
async def fake_list_earth_news_cruise_items(_db, *, limit, categories=None, source_ids=None):
captured["cruise_categories"] = categories
@@ -1090,7 +1134,8 @@ async def test_earth_news_payload_passes_region_and_category_filters_to_store(mo
assert captured["freshness_region"] == "europe"
assert captured["items_region"] == "europe"
assert captured["items_categories"] == {"business", "ecommerce"}
assert captured["items_source_ids"] is None
assert captured["items_source_ids"][0] is None
assert any(source_ids for source_ids in captured["items_source_ids"][1:])
assert captured["cruise_categories"] == {"business", "ecommerce"}
assert captured["cruise_source_ids"] is None
assert payload["filters"] == {

View File

@@ -0,0 +1,252 @@
from datetime import UTC, datetime
import pytest
from app.core.enums import NewsSourceType
from app.models.earth_news import EarthNewsItem
from app.models.system_setting import SystemSetting
from app.services.earth_news import REGION_ANCHORS
from app.services.earth_news_manual import (
DEFAULT_MANUAL_NEWS_GROUP_ID,
create_manual_news_group,
import_manual_news_items,
list_news_groups,
list_news_records,
parse_manual_news_import_upload,
rename_manual_news_group,
upsert_manual_news_item,
)
class _FakeResult:
def __init__(self, rows=None, scalar=None):
self.rows = rows or []
self._scalar = scalar
def scalar_one_or_none(self):
return self._scalar
def scalar(self):
return self._scalar
def scalars(self):
return self
def all(self):
return self.rows
class _FakeNewsSession:
def __init__(self, records=None, setting=None):
self.records = dict(records or {})
self.setting = setting
async def get(self, _model, item_id):
return self.records.get(item_id)
def add(self, item):
if isinstance(item, SystemSetting):
self.setting = item
else:
self.records[item.id] = item
async def execute(self, stmt):
statement = str(stmt)
if "system_settings" in statement:
return _FakeResult(scalar=self.setting)
if "count" in statement.lower():
return _FakeResult(scalar=len(self.records))
return _FakeResult(rows=list(self.records.values()))
async def flush(self):
return None
@pytest.fixture
def fake_news_queue(monkeypatch):
queued = []
async def _enqueue(payload, force=False):
queued.append({"payload": payload, "force": force})
return True
monkeypatch.setattr("app.services.earth_news_manual.enqueue_target_location_job", _enqueue)
return queued
@pytest.mark.asyncio
async def test_manual_news_upsert_uses_region_anchor_and_manual_metadata(fake_news_queue):
db = _FakeNewsSession()
result = await upsert_manual_news_item(
db,
{
"title": "手动添加的新闻",
"summary": "一条用于测试的手动新闻。",
"source": "人工录入",
"region": "europe",
"published_at": "2026-05-15T03:00:00Z",
"tags": ["manual", "test"],
},
)
anchor = REGION_ANCHORS["europe"]
assert result.created is True
assert result.queued is True
assert result.item.id.startswith("manual:")
assert result.item.feed_name == "手动添加"
assert result.item.source == "人工录入"
assert result.item.latitude == anchor.latitude
assert result.item.longitude == anchor.longitude
assert result.item.location_source == "region_anchor"
assert result.item.verified is False
assert result.item.location_meta["news_meta"]["feed_type"] == NewsSourceType.MANUAL.value
assert result.item.location_meta["news_meta"]["source_type"] == NewsSourceType.MANUAL.value
assert result.item.location_meta["news_meta"]["manual_group_id"] == DEFAULT_MANUAL_NEWS_GROUP_ID
assert len(fake_news_queue) == 1
@pytest.mark.asyncio
async def test_manual_news_duplicate_import_upserts_without_duplicate_rows(fake_news_queue):
db = _FakeNewsSession()
payload = {
"title": "Same manual story",
"source": "Manual Desk",
"published_at": "2026-05-15T03:00:00Z",
"region": "global",
}
first = await upsert_manual_news_item(db, payload)
second = await upsert_manual_news_item(db, {**payload, "summary": "Updated summary"})
assert first.created is True
assert second.created is False
assert len(db.records) == 1
assert db.records[first.item.id].summary == "Updated summary"
@pytest.mark.asyncio
async def test_manual_news_edit_without_location_preserves_manual_coordinates(fake_news_queue):
db = _FakeNewsSession()
created = await upsert_manual_news_item(
db,
{
"title": "Taipei-1 data center update",
"summary": "Initial summary.",
"region": "asia-pacific",
"published_at": "2026-05-15T03:00:00Z",
"location": {"label": "Kaohsiung, Taiwan", "latitude": 22.6273, "longitude": 120.3014},
},
)
updated = await upsert_manual_news_item(
db,
{
"title": "Taipei-1 data center update",
"summary": "Edited summary only.",
"region": "asia-pacific",
"published_at": "2026-05-15T03:00:00Z",
},
item_id_override=created.item.id,
)
assert updated.created is False
assert updated.item.latitude == pytest.approx(22.6273)
assert updated.item.longitude == pytest.approx(120.3014)
assert updated.item.location_source == "manual_location"
assert updated.item.verified is True
@pytest.mark.asyncio
async def test_manual_news_api_service_rejects_rss_records(fake_news_queue):
rss_record = EarthNewsItem(
id="bbc-world:example",
title="RSS story",
summary="RSS summary",
source="BBC World",
feed_name="BBC World",
region="global",
latitude=20,
longitude=0,
location_label="全球",
location_source="region_anchor",
verified=False,
location_meta={"news_meta": {"feed_type": "rss"}},
first_seen_at=datetime.now(UTC),
last_seen_at=datetime.now(UTC),
)
db = _FakeNewsSession({rss_record.id: rss_record})
with pytest.raises(PermissionError):
await upsert_manual_news_item(
db,
{"title": "Edited title", "region": "global"},
item_id_override=rss_record.id,
)
@pytest.mark.asyncio
async def test_manual_news_import_reports_per_item_errors(fake_news_queue):
db = _FakeNewsSession()
result = await import_manual_news_items(
db,
[
{"title": "Valid manual news", "region": "global"},
{"summary": "missing title"},
],
)
assert result["created"] == 1
assert result["failed"] == 1
assert result["errors"][0]["index"] == 1
@pytest.mark.asyncio
async def test_manual_news_groups_default_create_and_rename(fake_news_queue):
db = _FakeNewsSession()
initial = await list_news_groups(db)
assert initial["manual_groups"][0]["id"] == DEFAULT_MANUAL_NEWS_GROUP_ID
assert initial["manual_groups"][0]["name"] == "新建新闻组"
group = await create_manual_news_group(db, "专题组")
assert group["name"] == "专题组"
assert db.setting is not None
await upsert_manual_news_item(db, {"title": "Grouped story", "region": "global"}, group_id=group["id"])
renamed = await rename_manual_news_group(db, group["id"], "重命名专题")
record = next(iter(db.records.values()))
assert renamed["name"] == "重命名专题"
assert record.location_meta["news_meta"]["manual_group_id"] == group["id"]
assert record.location_meta["news_meta"]["manual_group_name"] == "重命名专题"
@pytest.mark.asyncio
async def test_manual_news_list_filters_by_group_id(fake_news_queue):
db = _FakeNewsSession()
group = await create_manual_news_group(db, "导入组")
await import_manual_news_items(
db,
[
{"title": "In group", "region": "global"},
{"title": "Also in group", "region": "global"},
],
group_id=group["id"],
)
await upsert_manual_news_item(db, {"title": "Default group", "region": "global"})
grouped = await list_news_records(db, page=1, page_size=20, group_id=group["id"])
default_group = await list_news_records(db, page=1, page_size=20, group_id=DEFAULT_MANUAL_NEWS_GROUP_ID)
assert grouped["total"] == 2
assert {item["manual_group_id"] for item in grouped["items"]} == {group["id"]}
assert default_group["total"] == 1
@pytest.mark.asyncio
async def test_manual_news_import_parser_requires_json_array():
with pytest.raises(ValueError, match="顶层必须是数组"):
await parse_manual_news_import_upload(b'{"title":"not an array"}')

View File

@@ -28,7 +28,7 @@ def test_protocol_enum_values_remain_api_compatible() -> None:
assert [item.value for item in NewsImportanceLevel] == ["low", "medium", "high", "critical"]
assert [item.value for item in BreakingLevel] == ["none", "watch", "breaking", "critical"]
assert [item.value for item in BreakingScope] == ["regional", "global"]
assert [item.value for item in NewsSourceType] == ["rss", "atom", "aggregated", "reference"]
assert [item.value for item in NewsSourceType] == ["rss", "atom", "aggregated", "reference", "manual"]
assert [item.value for item in UserRole] == ["viewer", "admin", "super_admin"]
assert JobStatus.RUNNING.value == "running"
assert PlaygroundMessageKind.THINKING.value == "thinking"

View File

@@ -40,6 +40,9 @@ def test_gesture_event_serializes_stable_protocol_fields():
assert payload["seq"] == 7
assert payload["source"] == "motion-agent"
assert payload["mode"] == "single"
assert payload["protocol_version"] == "motion.v2"
assert payload["input_mode"] == "single"
assert payload["camera_id"] == "unknown"
assert payload["payload"] == {}
@@ -88,8 +91,13 @@ def test_motion_server_status_includes_dry_run_camera_and_heartbeat():
assert status["camera_count"] == 1
assert status["active_camera_ids"] == ["dry-run:null-camera"]
assert status["recognizer"] == "dry-run"
assert status["protocol_version"] == "motion.v2"
assert status["armed"] is False
assert status["paused"] is False
assert status["devices_open"] is False
assert heartbeat == {
"timestamp_ms": 123,
"protocol_version": "motion.v2",
"source": "motion-agent",
"type": "heartbeat",
}
@@ -109,6 +117,7 @@ def test_skeleton_event_serializes_without_raw_image_fields():
payload = json.loads(event.to_json())
assert payload["type"] == "skeleton"
assert payload["protocol_version"] == "motion.v2"
assert payload["matched_gesture"] == "rotate_left"
assert payload["confidence"] == 0.91
assert payload["camera_id"] == "usb:0"
@@ -120,6 +129,25 @@ def test_skeleton_event_serializes_without_raw_image_fields():
assert "frame" not in payload
def test_v2_gesture_set_accepts_frontend_motion_gestures():
state = GestureStateMachine(confidence_threshold=0.7, cooldown_ms=0)
for gesture in [
"rotate_up",
"rotate_down",
"focus_prev",
"focus_next",
"layer_prev",
"layer_next",
]:
event = state.accept(
GestureObservation(gesture, confidence=0.9, intensity=0.8, timestamp_ms=1000)
)
assert event is not None
assert event.gesture == gesture
def test_dry_run_recognizer_produces_debug_skeleton():
server = MotionAgentServer(MotionAgentConfig(dry_run=True))
@@ -240,3 +268,92 @@ async def test_motion_agent_cli_reports_dependency_error_without_traceback(monke
assert exit_code == 2
assert "Motion agent failed: missing cv stack" in captured.err
assert "Traceback" not in captured.err
@pytest.mark.asyncio
async def test_motion_agent_command_updates_control_state():
server = MotionAgentServer(MotionAgentConfig(dry_run=True))
armed = await server.handle_command(
json.dumps(
{
"type": "command",
"command": "set_armed",
"request_id": "req-armed",
"payload": {"armed": True},
}
)
)
paused = await server.handle_command(
{
"type": "command",
"command": "set_paused",
"request_id": "req-paused",
"payload": {"paused": True},
}
)
assert armed.ok is True
assert armed.request_id == "req-armed"
assert armed.status["armed"] is True
assert paused.ok is True
assert paused.status["paused"] is True
@pytest.mark.asyncio
async def test_motion_agent_open_devices_command_accepts_dual_mode():
server = MotionAgentServer(MotionAgentConfig(dry_run=True))
try:
result = await server.handle_command(
{
"type": "command",
"command": "open_devices",
"request_id": "req-open",
"payload": {"input_mode": "dual_redundant"},
}
)
assert result.ok is True
assert result.status["input_mode"] == "dual_redundant"
assert result.status["active_camera_ids"] == ("dry-run:null-camera",)
assert server._recognition_subprocess is not None
assert server._recognition_subprocess.returncode is None
finally:
await server.stop_recognition_subprocess()
def test_motion_agent_dual_fusion_merges_matching_observations():
server = MotionAgentServer(MotionAgentConfig(dry_run=True))
server.state.mode = "dual_redundant"
selected, fusion = server._fuse_observations(
[
GestureObservation("zoom_in", confidence=0.82, intensity=0.4, camera_id="usb:0"),
GestureObservation("zoom_in", confidence=0.86, intensity=0.8, camera_id="usb:1"),
]
)
assert selected.gesture == "zoom_in"
assert selected.camera_id == "fusion"
assert selected.confidence > 0.86
assert fusion == {
"source_cameras": ["usb:1", "usb:0"],
"window_ms": 120,
"reason": "matched_observations",
}
def test_motion_agent_dual_fusion_suppresses_close_conflict():
server = MotionAgentServer(MotionAgentConfig(dry_run=True, confidence_threshold=0.7))
selected, fusion = server._fuse_observations(
[
GestureObservation("zoom_in", confidence=0.82, intensity=0.5, camera_id="usb:0"),
GestureObservation("zoom_out", confidence=0.78, intensity=0.5, camera_id="usb:1"),
]
)
assert selected.gesture == "zoom_in"
assert selected.confidence == 0
assert fusion["reason"] == "conflict_ignored"