Files
asmr-vector-search/scripts/search_server.py
T
2026-06-11 03:43:59 +09:00

570 lines
21 KiB
Python

#!/usr/bin/env python3
"""Serve the WebUI and provide local natural-language vector search."""
from __future__ import annotations
import argparse
import http.client
import json
import math
import os
from functools import partial
from http import HTTPStatus
from http.server import SimpleHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Any
from urllib.error import HTTPError, URLError
from urllib.parse import parse_qs, urlparse
from urllib.request import Request, urlopen
DEFAULT_HOST = "127.0.0.1"
DEFAULT_PORT = 8000
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Serve built Astro WebUI and expose /api/search for natural-language ASMR search."
)
parser.add_argument("--host", default=DEFAULT_HOST, help="Bind host.")
parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="Bind port.")
parser.add_argument("--web-dir", default="web/dist", help="Directory to serve as static WebUI.")
parser.add_argument("--data-dir", default="web/data", help="Directory containing manifest/works/embeddings.")
parser.add_argument("--batch-size", type=int, default=1, help="SentenceTransformer encode batch size.")
parser.add_argument(
"--api-key-env",
default="EMBEDDING_API_KEY",
help="Environment variable that contains the remote embeddings API key.",
)
parser.add_argument(
"--remote-base-url",
default=None,
help="Override manifest baseUrl for remote-openai-compatible search.",
)
parser.add_argument(
"--remote-query-prefix",
default="",
help="Optional prefix prepended to remote search queries, for example 'query: '.",
)
parser.add_argument("--timeout", type=float, default=120.0, help="Remote request timeout seconds.")
return parser.parse_args()
def normalize(value: Any) -> str:
return str(value or "").lower().strip()
def parse_limit(value: Any, default: int = 48, maximum: int = 200) -> int:
return max(1, min(maximum, int(value or default)))
def searchable_text(work: dict[str, Any]) -> str:
values = [
work.get("title"),
work.get("sourceId"),
work.get("circle"),
*(work.get("vas") or []),
*(work.get("tags") or []),
]
return " ".join(normalize(value) for value in values)
def recent_sort_value(work: dict[str, Any]) -> str:
return str(work.get("createDate") or work.get("release") or work.get("sourceId") or "")
def embeddings_url(base_url: str) -> str:
return f"{base_url.rstrip('/')}/v1/embeddings"
def normalize_vector(vector: list[float]) -> list[float]:
norm = math.sqrt(sum(value * value for value in vector))
if norm == 0:
return vector
return [value / norm for value in vector]
def request_remote_embedding(
*,
url: str,
model: str,
query: str,
api_key: str,
timeout: float,
) -> list[float]:
body = json.dumps({"model": model, "input": [query]}).encode("utf-8")
headers = {"Content-Type": "application/json"}
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
request = Request(url, data=body, headers=headers, method="POST")
try:
with urlopen(request, timeout=timeout) as response:
payload = json.load(response)
except HTTPError as exc:
detail = exc.read(1000).decode("utf-8", errors="replace")
raise RuntimeError(f"remote embeddings API returned HTTP {exc.code}: {detail}") from exc
except (
URLError,
TimeoutError,
json.JSONDecodeError,
http.client.IncompleteRead,
http.client.RemoteDisconnected,
ConnectionResetError,
OSError,
) as exc:
raise RuntimeError(f"remote embeddings request failed: {exc}") from exc
data = payload.get("data") if isinstance(payload, dict) else None
if not isinstance(data, list) or not data:
raise RuntimeError("remote embeddings response does not contain data[]")
first = data[0]
if not isinstance(first, dict):
raise RuntimeError("remote embeddings response data[0] is not an object")
embedding = first.get("embedding")
if not isinstance(embedding, list) or not embedding:
raise RuntimeError("remote embeddings response does not contain embedding")
try:
return [float(value) for value in embedding]
except (TypeError, ValueError) as exc:
raise RuntimeError("remote embedding contains non-numeric values") from exc
class LocalSentenceTransformerQueryEncoder:
def __init__(self, model_name: str, batch_size: int) -> None:
try:
from sentence_transformers import SentenceTransformer
except ImportError as exc:
raise SystemExit(
"Local natural-language search requires semantic dependencies. "
"Run `uv run --extra semantic python scripts/search_server.py`."
) from exc
self.model_name = model_name
self.batch_size = batch_size
self.model = SentenceTransformer(model_name)
def encode(self, query: str) -> list[float]:
text = f"query: {query}" if "e5" in self.model_name.lower() else query
embedding = self.model.encode(
[text],
batch_size=self.batch_size,
normalize_embeddings=True,
show_progress_bar=False,
)[0]
return [float(value) for value in embedding]
class RemoteOpenAICompatibleQueryEncoder:
def __init__(
self,
*,
base_url: str,
model: str,
api_key: str,
timeout: float,
query_prefix: str,
should_normalize: bool,
) -> None:
self.url = embeddings_url(base_url)
self.model = model
self.api_key = api_key
self.timeout = timeout
self.query_prefix = query_prefix
self.should_normalize = should_normalize
def encode(self, query: str) -> list[float]:
vector = request_remote_embedding(
url=self.url,
model=self.model,
query=f"{self.query_prefix}{query}",
api_key=self.api_key,
timeout=self.timeout,
)
return normalize_vector(vector) if self.should_normalize else vector
class SearchIndex:
def __init__(self, args: argparse.Namespace, data_dir: Path) -> None:
try:
import numpy as np
except ImportError as exc:
raise SystemExit(
"Natural-language search requires numpy. "
"Run `uv run --extra remote python scripts/search_server.py` or "
"`uv run --extra semantic python scripts/search_server.py`."
) from exc
self.np = np
self.manifest = self._read_json(data_dir / "manifest.json")
self.works = self._read_json(data_dir / "works.json")
self.method = str(self.manifest.get("method") or "")
self.dimensions = int(self.manifest["dimensions"])
self.model_name = str(self.manifest["model"])
embeddings_path = data_dir / str(self.manifest["embeddingFile"])
self.embeddings = np.fromfile(embeddings_path, dtype="<f4")
expected = len(self.works) * self.dimensions
if self.embeddings.size != expected:
raise SystemExit(
f"embedding size mismatch: expected {expected}, got {self.embeddings.size}"
)
self.embeddings = self.embeddings.reshape((len(self.works), self.dimensions))
score_config = self.manifest.get("score") if isinstance(self.manifest, dict) else None
self.vector_weight = float((score_config or {}).get("vectorWeight", 0.8))
self.tag_weight = float((score_config or {}).get("tagWeight", 0.2))
self.query_encoder = self.create_query_encoder(args)
@staticmethod
def _read_json(path: Path) -> Any:
with path.open(encoding="utf-8") as file:
return json.load(file)
def create_query_encoder(self, args: argparse.Namespace):
if self.method == "sentence-transformers":
return LocalSentenceTransformerQueryEncoder(self.model_name, args.batch_size)
if self.method == "remote-openai-compatible":
base_url = args.remote_base_url or self.manifest.get("baseUrl")
if not base_url:
raise SystemExit(
"remote-openai-compatible manifest does not contain baseUrl. "
"Pass --remote-base-url."
)
return RemoteOpenAICompatibleQueryEncoder(
base_url=str(base_url),
model=self.model_name,
api_key=os.environ.get(args.api_key_env, ""),
timeout=args.timeout,
query_prefix=args.remote_query_prefix,
should_normalize=bool(self.manifest.get("normalized", True)),
)
raise SystemExit(
f"unsupported embedding method for natural-language search: {self.method}. "
"Use sentence-transformers or remote-openai-compatible data."
)
def query_embedding(self, query: str):
embedding = self.query_encoder.encode(query)
if len(embedding) != self.dimensions:
raise RuntimeError(
f"query embedding dimensions mismatch: expected {self.dimensions}, got {len(embedding)}"
)
return self.np.asarray(embedding, dtype="<f4")
def search(self, payload: dict[str, Any]) -> dict[str, Any]:
query = str(payload.get("query") or "").strip()
if not query:
raise ValueError("query is required")
limit = parse_limit(payload.get("limit"))
ages = {normalize(value) for value in payload.get("ages") or [] if normalize(value)}
required_tags = [normalize(value) for value in payload.get("requiredTags") or [] if normalize(value)]
try:
query_vector = self.query_embedding(query)
except RuntimeError as exc:
raise ValueError(str(exc)) from exc
scores = self.embeddings @ query_vector
candidates = self.filtered_indices(ages, required_tags)
if candidates.size == 0:
return {"results": []}
candidate_scores = scores[candidates]
take = min(limit, candidate_scores.size)
top_positions = self.np.argpartition(-candidate_scores, take - 1)[:take]
top_positions = top_positions[self.np.argsort(-candidate_scores[top_positions])]
results = [
{
"index": int(candidates[position]),
"score": float(candidate_scores[position]),
"work": self.works[int(candidates[position])],
}
for position in top_positions
]
return {"results": results}
def list_works(self, payload: dict[str, Any]) -> dict[str, Any]:
limit = parse_limit(payload.get("limit"), default=100)
query = str(payload.get("query") or payload.get("q") or "").strip()
ages = {normalize(value) for value in payload.get("ages") or [] if normalize(value)}
required_tags = [
normalize(value)
for value in (payload.get("requiredTags") or payload.get("tags") or [])
if normalize(value)
]
entries = [
work
for work in self.works
if self.matches_work(work, ages=ages, required_tags=required_tags, query=query)
]
entries.sort(key=recent_sort_value, reverse=True)
return {"results": [{"work": work} for work in entries[:limit]]}
def get_work(self, identifier: Any) -> dict[str, Any]:
normalized_id = normalize(identifier)
if not normalized_id:
raise ValueError("id is required")
for work in self.works:
if (
normalize(work.get("sourceId")) == normalized_id
or normalize(work.get("id")) == normalized_id
or normalize(work.get("embeddingIndex")) == normalized_id
):
return {"work": work}
raise ValueError("work not found")
def similar(self, payload: dict[str, Any]) -> dict[str, Any]:
try:
source_index = int(payload.get("index"))
except (TypeError, ValueError) as exc:
raise ValueError("index is required") from exc
if not 0 <= source_index < len(self.works):
raise ValueError("index is out of range")
limit = parse_limit(payload.get("limit"))
filters = payload.get("filters") if isinstance(payload.get("filters"), dict) else {}
ages = {normalize(value) for value in filters.get("ages") or [] if normalize(value)}
exclude_same_group = bool(filters.get("excludeSameGroup"))
source_work = self.works[source_index]
source_group = source_work.get("groupId")
source_tags = set(source_work.get("tags") or [])
scores = self.embeddings @ self.embeddings[source_index]
candidates: list[int] = []
combined_scores: list[float] = []
vector_scores: list[float] = []
tag_scores: list[float] = []
for index, work in enumerate(self.works):
if index == source_index:
continue
if exclude_same_group and source_group and source_group == work.get("groupId"):
continue
if ages and normalize(work.get("ageCategory") or "unknown") not in ages:
continue
vector_score = float(scores[index])
tag_score = tag_jaccard(source_tags, work.get("tags") or [])
score = self.vector_weight * vector_score + self.tag_weight * tag_score
candidates.append(index)
combined_scores.append(score)
vector_scores.append(vector_score)
tag_scores.append(tag_score)
if not candidates:
return {"results": []}
order = sorted(range(len(candidates)), key=lambda position: combined_scores[position], reverse=True)
results = [
{
"index": int(candidates[position]),
"score": float(combined_scores[position]),
"vector": float(vector_scores[position]),
"tag": float(tag_scores[position]),
"work": self.works[int(candidates[position])],
}
for position in order[:limit]
]
return {"results": results}
def filtered_indices(self, ages: set[str], required_tags: list[str]):
indices: list[int] = []
for index, work in enumerate(self.works):
if ages and normalize(work.get("ageCategory")) not in ages:
continue
if required_tags:
tags = " ".join(normalize(tag) for tag in work.get("tags") or [])
if not all(tag in tags for tag in required_tags):
continue
indices.append(index)
return self.np.asarray(indices, dtype="int64")
@staticmethod
def matches_work(
work: dict[str, Any],
*,
ages: set[str],
required_tags: list[str],
query: str,
) -> bool:
if ages and normalize(work.get("ageCategory") or "unknown") not in ages:
return False
if required_tags:
tags = " ".join(normalize(tag) for tag in work.get("tags") or [])
if not all(tag in tags for tag in required_tags):
return False
if query and normalize(query) not in searchable_text(work):
return False
return True
def tag_jaccard(source_tags: set[str], candidate_tags: list[Any]) -> float:
if not source_tags and not candidate_tags:
return 0.0
intersection = 0
seen: set[Any] = set()
for tag in candidate_tags:
if tag in seen:
continue
seen.add(tag)
if tag in source_tags:
intersection += 1
union = len(source_tags) + len(seen) - intersection
return intersection / union if union else 0.0
class SearchRequestHandler(SimpleHTTPRequestHandler):
search_index: SearchIndex
def end_headers(self) -> None:
self.send_header("Access-Control-Allow-Origin", "*")
self.send_header("Access-Control-Allow-Headers", "content-type")
self.send_header("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
super().end_headers()
def do_OPTIONS(self) -> None:
self.send_response(HTTPStatus.NO_CONTENT)
self.end_headers()
def do_GET(self) -> None:
parsed = urlparse(self.path)
if parsed.path == "/api/status":
self.write_json(
{
"ok": True,
"count": len(self.search_index.works),
"model": self.search_index.model_name,
}
)
return
if parsed.path == "/api/search":
params = parse_qs(parsed.query)
payload = {
"query": (params.get("q") or [""])[0],
"limit": (params.get("limit") or [48])[0],
}
self.handle_search(payload)
return
if parsed.path == "/api/works":
params = parse_qs(parsed.query)
payload = {
"query": (params.get("q") or [""])[0],
"limit": (params.get("limit") or [100])[0],
"ages": params.get("age") or [],
"requiredTags": params.get("tag") or [],
}
self.handle_works(payload)
return
if parsed.path == "/api/work":
params = parse_qs(parsed.query)
self.handle_work((params.get("id") or [""])[0])
return
if parsed.path == "/api/similar":
params = parse_qs(parsed.query)
payload = {
"index": (params.get("index") or [None])[0],
"limit": (params.get("limit") or [48])[0],
"filters": {
"ages": params.get("age") or [],
"excludeSameGroup": (params.get("excludeSameGroup") or [""])[0] in {"1", "true"},
},
}
self.handle_similar(payload)
return
super().do_GET()
def do_POST(self) -> None:
parsed = urlparse(self.path)
if parsed.path not in {"/api/search", "/api/similar", "/api/works", "/api/work"}:
self.send_error(HTTPStatus.NOT_FOUND)
return
try:
length = int(self.headers.get("content-length") or 0)
body = self.rfile.read(length).decode("utf-8") if length else "{}"
payload = json.loads(body)
if not isinstance(payload, dict):
raise ValueError("request body must be a JSON object")
except (UnicodeDecodeError, json.JSONDecodeError, ValueError) as exc:
self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST)
return
if parsed.path == "/api/search":
self.handle_search(payload)
elif parsed.path == "/api/similar":
self.handle_similar(payload)
elif parsed.path == "/api/works":
self.handle_works(payload)
else:
self.handle_work(payload.get("id"))
def handle_search(self, payload: dict[str, Any]) -> None:
try:
self.write_json(self.search_index.search(payload))
except ValueError as exc:
self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST)
def handle_works(self, payload: dict[str, Any]) -> None:
try:
self.write_json(self.search_index.list_works(payload))
except ValueError as exc:
self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST)
def handle_work(self, identifier: Any) -> None:
try:
self.write_json(self.search_index.get_work(identifier))
except ValueError as exc:
status = HTTPStatus.NOT_FOUND if str(exc) == "work not found" else HTTPStatus.BAD_REQUEST
self.write_json({"error": str(exc)}, status)
def handle_similar(self, payload: dict[str, Any]) -> None:
try:
self.write_json(self.search_index.similar(payload))
except ValueError as exc:
self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST)
def write_json(self, payload: dict[str, Any], status: HTTPStatus = HTTPStatus.OK) -> None:
body = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
self.send_response(status)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def main() -> int:
args = parse_args()
web_dir = Path(args.web_dir).resolve()
data_dir = Path(args.data_dir).resolve()
if not web_dir.exists():
raise SystemExit(f"web directory not found: {web_dir}")
if not data_dir.exists():
raise SystemExit(f"data directory not found: {data_dir}")
print("Loading natural-language search index...")
SearchRequestHandler.search_index = SearchIndex(args, data_dir)
handler = partial(SearchRequestHandler, directory=str(web_dir))
server = ThreadingHTTPServer((args.host, args.port), handler)
print(f"Serving http://{args.host}:{args.port}")
print(f"Search model: {SearchRequestHandler.search_index.model_name}")
print(f"Search method: {SearchRequestHandler.search_index.method}")
try:
server.serve_forever()
except KeyboardInterrupt:
print("\nShutting down...")
finally:
server.server_close()
return 0
if __name__ == "__main__":
raise SystemExit(main())