manual_queries.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378
  1. """Manual query batch submission API."""
  2. from __future__ import annotations
  3. import os
  4. import subprocess
  5. import sys
  6. from datetime import datetime
  7. from pathlib import Path
  8. from typing import Any
  9. from uuid import UUID
  10. from fastapi import APIRouter, Body, Depends, HTTPException
  11. from pydantic import BaseModel, Field
  12. from acquisition.repositories.postgres import PostgresAcquisitionRepository
  13. from app.dependencies import _env_file, get_creation_db_config
  14. from app.routes.acquisition import _model_dump
  15. from core.config import CreationDbConfig
  16. from core.db_session import transaction
  17. from pipeline.postgres import PostgresPipelineRepository
  18. from pipeline.tracing import TraceContext, new_trace_writer
  19. from query_planning import (
  20. GenerationRequest,
  21. GeneratorKind,
  22. QueryBatchWriteSpec,
  23. QueryBatchWriter,
  24. UnifiedQueryGenerationService,
  25. planning_store_for_repository,
  26. )
  27. router = APIRouter(prefix="/api/query-batches", tags=["manual-query-batches"])
  28. ROOT = Path(__file__).resolve().parents[2]
  29. RUNTIME_MANUAL_DIR = ROOT / "runtime" / "manual_query_runs"
  30. DEFAULT_PLATFORMS = ("xiaohongshu", "weixin", "douyin")
  31. QUERY_MAX_COUNT = 500
  32. QUERY_MAX_CHARS = 500
  33. class ManualQueryEntry(BaseModel):
  34. query_text: str
  35. axes: dict[str, Any] = Field(default_factory=dict)
  36. metadata: dict[str, Any] = Field(default_factory=dict)
  37. filter_reason: str | None = None
  38. class ManualQueryBatchRequest(BaseModel):
  39. name: str | None = None
  40. queries: list[Any] | None = None
  41. families: list[dict[str, Any]] | None = None
  42. target_platforms: list[str] | None = None
  43. metadata: dict[str, Any] = Field(default_factory=dict)
  44. search_limit: int = Field(default=10, ge=1, le=50)
  45. display_limit: int = Field(default=5, ge=1, le=50)
  46. decode_limit: int = Field(default=100, ge=0, le=500)
  47. def _as_request(payload: Any) -> ManualQueryBatchRequest:
  48. if isinstance(payload, list):
  49. payload = {"queries": payload}
  50. if not isinstance(payload, dict):
  51. raise HTTPException(status_code=422, detail="请求体必须是对象或 query 数组")
  52. return ManualQueryBatchRequest.model_validate(payload)
  53. def _query_text(value: Any) -> str:
  54. if value is None:
  55. return ""
  56. return str(value).strip()
  57. def _entry_from_value(value: Any, *, family: dict[str, Any] | None = None) -> ManualQueryEntry | None:
  58. if isinstance(value, str):
  59. query_text = _query_text(value)
  60. axes: dict[str, Any] = {}
  61. metadata: dict[str, Any] = {}
  62. filter_reason = None
  63. elif isinstance(value, dict):
  64. query_text = _query_text(value.get("query_text") or value.get("query"))
  65. axes = value.get("axes") or value.get("parts") or {}
  66. if not isinstance(axes, dict):
  67. axes = {}
  68. metadata = value.get("metadata") or {}
  69. if not isinstance(metadata, dict):
  70. metadata = {}
  71. filter_reason = value.get("filter_reason") or value.get("reason")
  72. else:
  73. return None
  74. if not query_text:
  75. return None
  76. if len(query_text) > QUERY_MAX_CHARS:
  77. raise HTTPException(status_code=422, detail=f"query 过长,最多 {QUERY_MAX_CHARS} 字符")
  78. merged_metadata = dict(metadata)
  79. merged_metadata["family_key"] = "manual"
  80. if family:
  81. merged_metadata["source_family_key"] = family.get("key") or family.get("family_key")
  82. merged_metadata["source_family_name"] = family.get("name") or family.get("title")
  83. return ManualQueryEntry(
  84. query_text=query_text,
  85. axes=axes,
  86. metadata=merged_metadata,
  87. filter_reason=_query_text(filter_reason) or None,
  88. )
  89. def _normalize_entries(request: ManualQueryBatchRequest) -> list[ManualQueryEntry]:
  90. entries: list[ManualQueryEntry] = []
  91. for value in request.queries or []:
  92. entry = _entry_from_value(value)
  93. if entry:
  94. entries.append(entry)
  95. for family in request.families or []:
  96. if not isinstance(family, dict):
  97. continue
  98. for value in family.get("items") or []:
  99. entry = _entry_from_value(value, family=family)
  100. if entry:
  101. entries.append(entry)
  102. if not entries:
  103. raise HTTPException(status_code=422, detail="没有可提交的 query")
  104. return entries
  105. def _normalize_platforms(platforms: list[str] | None) -> list[str]:
  106. values = platforms or list(DEFAULT_PLATFORMS)
  107. if not values:
  108. raise HTTPException(status_code=422, detail="至少选择一个平台")
  109. seen: set[str] = set()
  110. out: list[str] = []
  111. for platform in values:
  112. key = str(platform).strip()
  113. if key not in DEFAULT_PLATFORMS:
  114. raise HTTPException(status_code=422, detail=f"不支持的平台:{key}")
  115. if key not in seen:
  116. seen.add(key)
  117. out.append(key)
  118. return out
  119. def _pipeline_command(
  120. *,
  121. batch_id: str,
  122. run_key: str,
  123. platforms: list[str],
  124. search_limit: int,
  125. display_limit: int,
  126. decode_limit: int,
  127. ) -> list[str]:
  128. cmd = [
  129. sys.executable,
  130. str(ROOT / "scripts" / "run_creation_pipeline.py"),
  131. "--batch-id",
  132. batch_id,
  133. "--search-limit",
  134. str(search_limit),
  135. "--display-limit",
  136. str(display_limit),
  137. "--decode-limit",
  138. str(decode_limit),
  139. "--run-key",
  140. run_key,
  141. "--env-file",
  142. _env_file(),
  143. ]
  144. for platform in platforms:
  145. cmd.extend(["--platform", platform])
  146. return cmd
  147. def _create_pipeline_trace(
  148. db_config: CreationDbConfig,
  149. *,
  150. batch_id: Any,
  151. run_key: str,
  152. platforms: list[str],
  153. request: ManualQueryBatchRequest,
  154. log_path: Path,
  155. ) -> str | None:
  156. try:
  157. with transaction(db_config) as conn:
  158. run = PostgresPipelineRepository(conn).create_pipeline_run(
  159. run_key=run_key,
  160. batch_id=batch_id,
  161. status="pending",
  162. current_stage="query",
  163. config={
  164. "platforms": platforms,
  165. "search_limit": request.search_limit,
  166. "display_limit": request.display_limit,
  167. "decode_limit": request.decode_limit,
  168. },
  169. metadata={
  170. "source": "manual_api",
  171. "log_path": str(log_path),
  172. },
  173. )
  174. return str(run.id) if run.id else None
  175. except Exception:
  176. return None
  177. @router.post("/manual")
  178. def create_manual_query_batch(
  179. payload: Any = Body(...),
  180. db_config: CreationDbConfig = Depends(get_creation_db_config),
  181. ) -> dict[str, Any]:
  182. request = _as_request(payload)
  183. entries = _normalize_entries(request)
  184. platforms = _normalize_platforms(request.target_platforms)
  185. timestamp = datetime.now().strftime("%Y%m%d-%H%M%S")
  186. name = request.name or f"manual-query-{timestamp}"
  187. generation_result = UnifiedQueryGenerationService().generate(
  188. GenerationRequest(
  189. generator_kind=GeneratorKind.MANUAL,
  190. name=name,
  191. target_platforms=tuple(platforms),
  192. payload={
  193. "candidates": [
  194. {
  195. "query_text": entry.query_text,
  196. "axes": entry.axes,
  197. "filter_reason": entry.filter_reason,
  198. "metadata": entry.metadata,
  199. "source_refs": [
  200. {
  201. "generator_kind": GeneratorKind.MANUAL.value,
  202. "original_position": index,
  203. "axes": entry.axes,
  204. "source_family_key": entry.metadata.get("source_family_key"),
  205. "source_family_name": entry.metadata.get("source_family_name"),
  206. }
  207. ],
  208. }
  209. for index, entry in enumerate(entries)
  210. ],
  211. "generator_config": {
  212. "request_shape": "manual_query_api_v1",
  213. },
  214. "input_snapshot": {
  215. "submitted_count": len(entries),
  216. "target_platforms": platforms,
  217. },
  218. },
  219. metadata={"source": "manual_api"},
  220. )
  221. )
  222. if generation_result.stats.selected_count > QUERY_MAX_COUNT:
  223. raise HTTPException(status_code=422, detail=f"一次最多提交 {QUERY_MAX_COUNT} 条 query")
  224. batch_metadata = {
  225. "family_key": "manual",
  226. "source": "manual_api",
  227. "query_count": generation_result.stats.selected_count,
  228. "platforms": platforms,
  229. **request.metadata,
  230. }
  231. with transaction(db_config) as conn:
  232. repo = PostgresAcquisitionRepository(conn)
  233. write_result = QueryBatchWriter(
  234. legacy_sink=repo,
  235. planning_store=planning_store_for_repository(repo),
  236. ).write(
  237. generation_result,
  238. QueryBatchWriteSpec(
  239. name=name,
  240. source_type="manual",
  241. generation_method="manual_query_api_v1",
  242. target_platforms=tuple(platforms),
  243. metadata=batch_metadata,
  244. ),
  245. )
  246. batch = write_result.batch
  247. queries = write_result.queries
  248. run_key = f"manual-api:{batch.id}:{timestamp}"
  249. log_path = RUNTIME_MANUAL_DIR / f"{run_key.replace(':', '-')}.log"
  250. cmd = _pipeline_command(
  251. batch_id=str(batch.id),
  252. run_key=run_key,
  253. platforms=platforms,
  254. search_limit=request.search_limit,
  255. display_limit=request.display_limit,
  256. decode_limit=request.decode_limit,
  257. )
  258. run_metadata = {
  259. "source": "manual_api",
  260. "batch_id": str(batch.id),
  261. "platforms": platforms,
  262. "search_limit": request.search_limit,
  263. "display_limit": request.display_limit,
  264. "decode_limit": request.decode_limit,
  265. "dry_ingest_record": True,
  266. "log_path": str(log_path),
  267. "command": cmd,
  268. }
  269. run = repo.create_acquisition_run(
  270. batch_id=batch.id,
  271. run_key=run_key,
  272. status="pending",
  273. note="manual query API queued",
  274. metadata=run_metadata,
  275. )
  276. pipeline_run_id = _create_pipeline_trace(
  277. db_config,
  278. batch_id=batch.id,
  279. run_key=run_key,
  280. platforms=platforms,
  281. request=request,
  282. log_path=log_path,
  283. )
  284. if pipeline_run_id:
  285. trace_writer = new_trace_writer(db_config, env_file=_env_file())
  286. trace_writer.event(
  287. context=TraceContext(
  288. pipeline_run_id=UUID(pipeline_run_id),
  289. acquisition_run_id=run.id,
  290. stage="query",
  291. ),
  292. stage="query",
  293. event_type="manual_query_queued",
  294. status="pending",
  295. payload={"query_count": len(queries), "platforms": platforms},
  296. )
  297. RUNTIME_MANUAL_DIR.mkdir(parents=True, exist_ok=True)
  298. env = os.environ.copy()
  299. env["PYTHONPATH"] = f"{ROOT}:{env.get('PYTHONPATH', '')}".rstrip(":")
  300. try:
  301. with log_path.open("ab") as stream:
  302. process = subprocess.Popen(
  303. cmd,
  304. cwd=str(ROOT),
  305. env=env,
  306. stdout=stream,
  307. stderr=subprocess.STDOUT,
  308. start_new_session=True,
  309. )
  310. if pipeline_run_id:
  311. trace_writer = new_trace_writer(db_config, env_file=_env_file())
  312. trace_writer.event(
  313. context=TraceContext(
  314. pipeline_run_id=UUID(pipeline_run_id),
  315. acquisition_run_id=run.id,
  316. stage="query",
  317. ),
  318. stage="query",
  319. event_type="process_started",
  320. status="running",
  321. payload={"pid": process.pid, "log_path": str(log_path)},
  322. )
  323. except OSError as exc:
  324. with transaction(db_config) as conn:
  325. PostgresAcquisitionRepository(conn).update_acquisition_run(
  326. run.id,
  327. status="failed",
  328. error_message=f"启动手动 query pipeline 失败:{exc}",
  329. metadata={"log_path": str(log_path), "command": cmd},
  330. )
  331. raise HTTPException(status_code=500, detail="启动后台 pipeline 失败") from exc
  332. return {
  333. "status": "queued",
  334. "batch_id": str(batch.id),
  335. "pipeline_run_id": pipeline_run_id,
  336. "run_id": str(run.id),
  337. "run_key": run_key,
  338. "pid": process.pid,
  339. "query_count": len(queries),
  340. "queries": [_model_dump(query) for query in queries],
  341. "summary_url": f"/api/acquisition/runs/{run.id}/summary",
  342. "log_path": str(log_path),
  343. }