Files
asmr-vector-search/scripts/update_pulled_vector_data.py

254 lines
8.9 KiB
Python

#!/usr/bin/env python3
"""Update web/data after pulling JSONL data generated by actions."""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
from typing import Any, Iterable
ROOT_DIR = Path(__file__).resolve().parents[1]
if str(ROOT_DIR) not in sys.path:
sys.path.insert(0, str(ROOT_DIR))
from build_vector_data import normalize_work # noqa: E402
from update_daily_asmr_data import ( # noqa: E402
build_incremental_web_data,
dry_run_web_stats,
format_work,
identity_keys,
read_json,
read_jsonl,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Incrementally update web/data from the pulled filtered JSONL dataset."
)
parser.add_argument(
"--filtered-input",
default="asmr_works.filtered.jsonl",
help="Pulled filtered ASMR JSONL file.",
)
parser.add_argument("--data-dir", default="web/data", help="Directory containing web data assets.")
parser.add_argument("--dry-run", action="store_true", help="Show diffs without writing web/data.")
parser.add_argument("--force", action="store_true", help="Rewrite web/data even when no diff is detected.")
parser.add_argument(
"--diff-limit",
type=int,
default=50,
help="Maximum rows to print per diff section. Use 0 for all rows.",
)
parser.add_argument(
"--api-key-env",
default="EMBEDDING_API_KEY",
help="Environment variable that contains the remote embedding API key.",
)
parser.add_argument("--embedding-timeout", type=float, default=120.0, help="Embedding request timeout seconds.")
parser.add_argument("--embedding-retries", type=int, default=3, help="Embedding retries per failed batch.")
parser.add_argument("--embedding-retry-wait", type=float, default=2.0, help="Initial embedding retry wait seconds.")
return parser.parse_args()
def validate_args(args: argparse.Namespace) -> None:
if args.diff_limit < 0:
raise SystemExit("--diff-limit must be greater than or equal to 0")
if args.embedding_timeout <= 0:
raise SystemExit("--embedding-timeout must be greater than 0")
if args.embedding_retries < 0:
raise SystemExit("--embedding-retries must be greater than or equal to 0")
if args.embedding_retry_wait < 0:
raise SystemExit("--embedding-retry-wait must be greater than or equal to 0")
def normalized_filtered_works(filtered_works: Iterable[dict[str, Any]]) -> list[dict[str, Any]]:
works: list[dict[str, Any]] = []
for raw_work in filtered_works:
works.append(normalize_work(raw_work, len(works)))
return works
def key_map(works: Iterable[dict[str, Any]]) -> dict[str, dict[str, Any]]:
result: dict[str, dict[str, Any]] = {}
for work in works:
for key in identity_keys(work):
result.setdefault(key, work)
return result
def has_matching_key(work: dict[str, Any], works_by_key: dict[str, dict[str, Any]]) -> bool:
return any(key in works_by_key for key in identity_keys(work))
def matching_work(work: dict[str, Any], works_by_key: dict[str, dict[str, Any]]) -> dict[str, Any] | None:
for key in identity_keys(work):
match = works_by_key.get(key)
if match is not None:
return match
return None
def changed_fields(old: dict[str, Any], new: dict[str, Any]) -> list[str]:
fields = sorted((old.keys() | new.keys()) - {"embeddingIndex"})
return [field for field in fields if old.get(field) != new.get(field)]
def data_diffs(
*,
old_works: list[dict[str, Any]],
new_works: list[dict[str, Any]],
) -> dict[str, list[Any]]:
old_by_key = key_map(old_works)
new_by_key = key_map(new_works)
added: list[dict[str, Any]] = []
removed: list[dict[str, Any]] = []
changed: list[dict[str, Any]] = []
reordered: list[dict[str, Any]] = []
for work in new_works:
old_work = matching_work(work, old_by_key)
if old_work is None:
added.append(work)
elif old_work != work:
reasons = changed_fields(old_work, work)
if reasons:
changed.append({"work": work, "reasons": reasons})
else:
reordered.append(work)
for work in old_works:
if not has_matching_key(work, new_by_key):
removed.append(work)
return {"added": added, "removed": removed, "changed": changed, "reordered": reordered}
def limited_rows(rows: list[Any], limit: int) -> list[Any]:
return rows if limit == 0 else rows[:limit]
def remaining_count(rows: list[Any], limit: int) -> int:
return 0 if limit == 0 else max(0, len(rows) - limit)
def print_work_section(title: str, prefix: str, works: list[dict[str, Any]], limit: int) -> None:
print(f"\n{title}: {len(works)}")
for work in limited_rows(works, limit):
print(f" {prefix} {format_work(work)}")
remaining = remaining_count(works, limit)
if remaining:
print(f" ... {remaining} more")
def print_changed_section(changes: list[dict[str, Any]], limit: int) -> None:
print(f"\nChanged web works: {len(changes)}")
for change in limited_rows(changes, limit):
print(f" ~ {format_work(change['work'])} / {', '.join(change['reasons'])}")
remaining = remaining_count(changes, limit)
if remaining:
print(f" ... {remaining} more")
def print_embedding_section(diffs: list[dict[str, Any]], limit: int) -> None:
print(f"\nEmbedding requests: {len(diffs)}")
for diff in limited_rows(diffs, limit):
marker = "+" if diff["reason"] == "new work" else "~"
print(f" {marker} {format_work(diff['work'])} / {diff['reason']}")
remaining = remaining_count(diffs, limit)
if remaining:
print(f" ... {remaining} more")
def embeddings_file_in_sync(data_dir: Path, manifest: dict[str, Any], old_works: list[dict[str, Any]]) -> bool:
embeddings_path = data_dir / str(manifest.get("embeddingFile", "embeddings.f32"))
if not embeddings_path.exists():
return False
try:
dimensions = int(manifest["dimensions"])
except (KeyError, TypeError, ValueError):
return False
return embeddings_path.stat().st_size == len(old_works) * dimensions * 4
def needs_update(
*,
old_works: list[dict[str, Any]],
new_works: list[dict[str, Any]],
manifest: dict[str, Any],
data_dir: Path,
) -> bool:
if old_works != new_works:
return True
if int(manifest.get("count", -1)) != len(new_works):
return True
return not embeddings_file_in_sync(data_dir, manifest, old_works)
def main() -> int:
args = parse_args()
validate_args(args)
filtered_path = Path(args.filtered_input)
data_dir = Path(args.data_dir)
manifest_path = data_dir / "manifest.json"
works_path = data_dir / "works.json"
if not filtered_path.exists():
raise SystemExit(f"filtered input file not found: {filtered_path}")
if not manifest_path.exists():
raise SystemExit(f"manifest file not found: {manifest_path}")
if not works_path.exists():
raise SystemExit(f"works file not found: {works_path}")
filtered_works = read_jsonl(filtered_path)
new_works = normalized_filtered_works(filtered_works)
old_works = read_json(works_path)
if not isinstance(old_works, list):
raise SystemExit(f"{works_path} must contain a JSON array")
manifest = read_json(manifest_path)
if not isinstance(manifest, dict):
raise SystemExit(f"{manifest_path} must contain a JSON object")
diffs = data_diffs(old_works=old_works, new_works=new_works)
web_stats = dry_run_web_stats(args, filtered_works)
should_update = needs_update(
old_works=old_works,
new_works=new_works,
manifest=manifest,
data_dir=data_dir,
)
print(f"Filtered works: {len(new_works)}")
print(f"Existing web/data works: {len(old_works)}")
print(f"web/data would reuse {web_stats['reused']} embeddings and request {web_stats['embedded']} embeddings")
if args.dry_run:
print_work_section("Web data additions", "+", diffs["added"], args.diff_limit)
print_work_section("Web data removals", "-", diffs["removed"], args.diff_limit)
print_changed_section(diffs["changed"], args.diff_limit)
print(f"\nReordered existing works: {len(diffs['reordered'])}")
print_embedding_section(web_stats.get("embeddingDiffs", []), args.diff_limit)
print("\nDry run: no files were written.")
return 0
if not should_update and not args.force:
print("web/data is already in sync. update skipped.")
return 0
try:
stats = build_incremental_web_data(args, filtered_works)
except RuntimeError as exc:
raise SystemExit(f"web/data update failed: {exc}") from exc
print(
f"Updated {data_dir}: {stats['count']} works, "
f"reused {stats['reused']} embeddings, requested {stats['embedded']} embeddings"
)
return 0
if __name__ == "__main__":
raise SystemExit(main())