| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256 |
- """搜索结果的内部持久化流程,不暴露为 Agent 工具。"""
- from __future__ import annotations
- import hashlib
- import json
- import logging
- from typing import Any
- from supply_infra.video_discovery_gates import (
- normalize_duration_seconds,
- parse_datetime_value,
- )
- from supply_infra.services.video_discovery_service import (
- RunNotFoundError,
- format_db_error,
- get_video_discovery_service,
- )
- logger = logging.getLogger(__name__)
- _SOURCE_TYPES = {
- "demand",
- "seed",
- "point",
- "tag",
- "author",
- "pagination",
- "mixed",
- }
- _ROOT_SOURCE_TYPES = {"demand", "seed", "point", "mixed"}
- def _load_payload(payload_json: str) -> dict[str, Any]:
- try:
- payload = json.loads(payload_json)
- except (TypeError, ValueError):
- return {"error": "搜索接口返回了无效 JSON", "raw_result": str(payload_json)}
- if not isinstance(payload, dict):
- return {"error": "搜索接口返回值不是对象", "raw_result": payload}
- return payload
- def _clean_text(value: Any, *, max_length: int | None = None) -> str | None:
- if value is None:
- return None
- text = str(value).strip()
- if not text:
- return None
- return text[:max_length] if max_length else text
- def _nonnegative_int(value: Any) -> int | None:
- if value is None or value == "":
- return None
- try:
- return max(0, int(value))
- except (TypeError, ValueError):
- return None
- def _search_key(values: dict[str, Any]) -> str:
- identity = {
- key: values.get(key)
- for key in (
- "provider",
- "keyword",
- "content_type",
- "sort_type",
- "publish_time",
- "cursor",
- )
- }
- raw = json.dumps(identity, ensure_ascii=False, sort_keys=True)
- return hashlib.sha256(raw.encode("utf-8")).hexdigest()
- def _candidate_from_search_result(
- item: dict[str, Any],
- keyword: str,
- ) -> dict[str, Any] | None:
- aweme_id = _clean_text(
- item.get("aweme_id") or item.get("content_id"),
- max_length=64,
- )
- if not aweme_id:
- return None
- author = item.get("author") if isinstance(item.get("author"), dict) else {}
- stats = item.get("statistics") if isinstance(item.get("statistics"), dict) else {}
- topics = item.get("topics") if isinstance(item.get("topics"), list) else []
- duration = normalize_duration_seconds(
- item.get("duration_seconds")
- if item.get("duration_seconds") not in (None, "")
- else item.get("duration_ms"),
- unit=(
- "seconds"
- if item.get("duration_seconds") not in (None, "")
- else "milliseconds"
- ),
- )
- publish_at = parse_datetime_value(
- item.get("publish_at")
- or item.get("create_time")
- or item.get("create_timestamp")
- or item.get("publish_timestamp")
- )
- return {
- "aweme_id": aweme_id,
- "title": _clean_text(item.get("desc") or item.get("title"), max_length=512),
- "content_link": _clean_text(
- item.get("url") or item.get("content_link"),
- max_length=1024,
- ),
- "author_name": _clean_text(
- author.get("nickname") or item.get("author_name"),
- max_length=256,
- ),
- "author_sec_uid": _clean_text(
- author.get("sec_uid") or item.get("author_sec_uid"),
- max_length=256,
- ),
- "like_count": _nonnegative_int(
- stats.get("digg_count") or item.get("like_count")
- ),
- "comment_count": _nonnegative_int(
- stats.get("comment_count") or item.get("comment_count")
- ),
- "share_count": _nonnegative_int(
- stats.get("share_count") or item.get("share_count")
- ),
- "collect_count": _nonnegative_int(
- stats.get("collect_count") or item.get("collect_count")
- ),
- "play_count": _nonnegative_int(
- stats.get("play_count") or item.get("play_count")
- ),
- "duration_seconds": duration,
- "publish_at": (
- publish_at.replace(tzinfo=None) if publish_at is not None else None
- ),
- "tags_json": (
- json.dumps(topics, ensure_ascii=False) if topics else None
- ),
- "_source_keyword": keyword,
- }
- def persist_search_payload(
- payload_json: str,
- *,
- run_id: str,
- keyword: str,
- query_reason: str,
- source_type: str,
- cursor: str,
- page_no: int,
- provider: str,
- source_value: str | None = None,
- parent_search_id: int | None = None,
- content_type: str = "视频",
- sort_type: str = "综合排序",
- publish_time: str = "不限",
- provider_state: dict[str, Any] | None = None,
- ) -> str:
- """新增搜索记录和本页全部候选,并把数据库 ID 拼回搜索结果。"""
- payload = _load_payload(payload_json)
- run_text = _clean_text(run_id, max_length=64)
- keyword_text = _clean_text(keyword, max_length=256)
- reason_text = _clean_text(query_reason)
- if not run_text or not keyword_text or not reason_text:
- payload["error"] = "run_id、keyword、query_reason 不能为空"
- payload["input_error"] = True
- return json.dumps(payload, ensure_ascii=False, default=str)
- if source_type not in _SOURCE_TYPES:
- payload["error"] = f"source_type 必须是: {sorted(_SOURCE_TYPES)}"
- payload["input_error"] = True
- return json.dumps(payload, ensure_ascii=False, default=str)
- raw_results = payload.get("search_results")
- results = raw_results if isinstance(raw_results, list) else []
- candidate_rows = [
- row
- for item in results
- if isinstance(item, dict)
- if (row := _candidate_from_search_result(dict(item), keyword_text)) is not None
- ]
- normalized_page_no = max(1, int(page_no))
- normalized_source_type = (
- "pagination" if normalized_page_no > 1 else source_type
- )
- normalized_parent_search_id = (
- None
- if normalized_source_type in _ROOT_SOURCE_TYPES
- else parent_search_id
- )
- merged_provider_state = dict(provider_state or {})
- for key in ("search_id", "backtrace"):
- value = payload.get(key)
- if value not in (None, ""):
- merged_provider_state[key] = value
- provider_search_id = payload.get("search_id")
- search_values: dict[str, Any] = {
- "run_id": run_text,
- "keyword": keyword_text,
- "query_reason": reason_text,
- "source_type": normalized_source_type,
- "source_value": _clean_text(source_value),
- "parent_search_id": normalized_parent_search_id,
- "provider": _clean_text(provider, max_length=32) or "internal_keyword",
- "provider_state_json": (
- json.dumps(merged_provider_state, ensure_ascii=False)
- if merged_provider_state
- else None
- ),
- "content_type": _clean_text(content_type, max_length=16) or "视频",
- "sort_type": _clean_text(sort_type, max_length=32) or "综合排序",
- "publish_time": _clean_text(publish_time, max_length=32) or "不限",
- "cursor": _clean_text(cursor, max_length=128) or "0",
- "page_no": normalized_page_no,
- "results_count": len(results),
- "new_candidate_count": 0,
- "has_more": int(bool(payload.get("has_more"))),
- "next_cursor": (
- _clean_text(payload.get("next_cursor"), max_length=128)
- ),
- "result_ids_json": None,
- "status": "failed" if payload.get("error") else "success",
- "error_message": _clean_text(payload.get("error")),
- }
- search_values["search_key"] = _search_key(search_values)
- try:
- saved = get_video_discovery_service().save_search_page(
- run_text,
- search_values,
- candidate_rows,
- )
- except RunNotFoundError as exc:
- payload["error"] = str(exc)
- payload["input_error"] = True
- return json.dumps(payload, ensure_ascii=False, default=str)
- except Exception as exc:
- logger.error("persist search payload failed: %s", exc, exc_info=True)
- payload["error"] = format_db_error(exc)
- return json.dumps(payload, ensure_ascii=False, default=str)
- payload.pop("search_results", None)
- payload.pop("user_videos", None)
- if provider_search_id not in (None, ""):
- payload["provider_search_id"] = provider_search_id
- payload.update(saved)
- payload["persisted"] = True
- return json.dumps(payload, ensure_ascii=False, default=str)
|