"""Search: `/search` (FTS + trigram), `/search/suggest` (fast prefix), `/ask` (deterministic parser, LLM optional).""" from __future__ import annotations import inspect import logging import time from dataclasses import dataclass, field from typing import Any from fastapi import APIRouter, Query, Response from companyatlas.api import aggregates as agg from companyatlas.api import ask_fallback from companyatlas.api import queries as q from companyatlas.api import serializers as ser from companyatlas.api.common import cached, public_cache_value from companyatlas.db import connection, fetch_all from companyatlas.ids import normalize_alias, slugify from companyatlas.taxonomy import EventType log = logging.getLogger("companyatlas.api.search") ORDER = 10 router = APIRouter(prefix="/api/v1", tags=["search"]) ALL_TYPES = ("companies", "events", "industries", "countries", "people", "products") SUGGEST_MAX = 10 def _clean(qs: str) -> str: return " ".join(qs.replace("%", " ").replace("_", " ").split())[:200] async def _search_companies(conn: Any, qs: str, limit: int) -> list[dict[str, Any]]: if len(qs) >= 3: rows = await fetch_all(conn, """ select c.id, greatest(coalesce(ts_rank(c.search, websearch_to_tsquery('simple', :q)), 0), similarity(c.display_name, :q), similarity(c.canonical_domain, :q), case when a.company_id is not null then 1.0 else 0 end) as score from companies c left join (select distinct company_id from company_aliases where alias_norm = :alias) a on a.company_id = c.id where c.search @@ websearch_to_tsquery('simple', :q) or c.display_name % :q or c.canonical_domain % :q or a.company_id is not null order by score desc, c.importance desc limit :lim""", q=qs, alias=normalize_alias(qs), lim=limit) else: rows = await fetch_all(conn, "select c.id, 1.0 as score from companies c where c.display_name ilike :p or c.canonical_domain ilike :p " "order by c.importance desc limit :lim", p=qs + "%", lim=limit) cards = await q.fetch_cards_by_ids(conn, [r["id"] for r in rows]) return [ser.company_card(c) for c in cards] async def _search_events(conn: Any, qs: str, limit: int) -> list[dict[str, Any]]: if len(qs) < 2: return [] rows = await fetch_all(conn, f"{q.EVENT_SELECT} where e.status = 'active' and e.search @@ websearch_to_tsquery('english', :q) " "order by ts_rank(e.search, websearch_to_tsquery('english', :q)) desc, e.detected_at desc limit :lim", q=qs, lim=limit) return [ser.event(r) for r in rows] async def _search_people(conn: Any, qs: str, limit: int) -> list[dict[str, Any]]: if len(qs) < 3: return [] rows = await fetch_all(conn, "select p.*, c.slug as company_slug, c.display_name as company_display_name, c.canonical_domain as company_domain, " "c.country as company_country, c.logo_url as company_logo_url from people p join companies c on c.id = p.company_id " "where p.name % :q or p.name ilike :like order by similarity(p.name, :q) desc, p.is_executive desc limit :lim", q=qs, like=f"%{qs}%", lim=limit) out = [] for r in rows: item = ser.person(r) item["company"] = ser.company_ref(r) out.append(item) return out async def _search_products(conn: Any, qs: str, limit: int) -> list[dict[str, Any]]: if len(qs) < 3: return [] rows = await fetch_all(conn, "select p.*, c.slug as company_slug, c.display_name as company_display_name, c.canonical_domain as company_domain, " "c.country as company_country, c.logo_url as company_logo_url from products p join companies c on c.id = p.company_id " "where p.name % :q or p.name ilike :like order by similarity(p.name, :q) desc, p.last_seen_at desc limit :lim", q=qs, like=f"%{qs}%", lim=limit) out = [] for r in rows: item = ser.product(r) item["company"] = ser.company_ref(r) out.append(item) return out def _match_rows(rows: list[dict[str, Any]], qs: str, keys: tuple[str, ...], limit: int) -> list[dict[str, Any]]: ql = qs.lower() starts = [r for r in rows if any(str(r.get(k) or "").lower().startswith(ql) for k in keys)] contains = [r for r in rows if r not in starts and any(ql in str(r.get(k) or "").lower() for k in keys)] return (starts + contains)[:limit] @router.get("/search", summary="Search companies, events, industries, countries, people and products") async def search(response: Response, q_: str = Query(..., alias="q", min_length=1, max_length=200), types: str = Query(",".join(ALL_TYPES)), limit: int = Query(10, ge=1, le=50)) -> dict[str, Any]: response.headers["cache-control"] = public_cache_value(30) t0 = time.perf_counter() qs = _clean(q_) wanted = {t for t in q.csv_list(types) if t in ALL_TYPES} or set(ALL_TYPES) out: dict[str, Any] = {"query": qs, "companies": [], "events": [], "industries": [], "countries": [], "people": [], "products": []} if qs: async with connection() as conn: if "companies" in wanted: out["companies"] = await _search_companies(conn, qs, limit) if "events" in wanted: out["events"] = await _search_events(conn, qs, limit) if "people" in wanted: out["people"] = await _search_people(conn, qs, limit) if "products" in wanted: out["products"] = await _search_products(conn, qs, limit) if "industries" in wanted: out["industries"] = _match_rows(await agg.cached_industry_rows(), qs, ("name", "slug"), limit) if "countries" in wanted: out["countries"] = _match_rows(await agg.cached_country_rows(), qs, ("name", "code"), limit) out["took_ms"] = round((time.perf_counter() - t0) * 1000, 1) return out @router.get("/search/suggest", summary="Typeahead suggestions (≤ 10)") async def suggest(response: Response, q_: str = Query(..., alias="q", min_length=1, max_length=100)) -> dict[str, Any]: response.headers["cache-control"] = public_cache_value(30) qs = _clean(q_) if not qs: return {"items": []} async def produce() -> dict[str, Any]: items: list[dict[str, Any]] = [] async with connection() as conn: rows = await fetch_all(conn, "select c.slug, c.display_name, c.canonical_domain, c.country from companies c " "where c.display_name ilike :p or c.canonical_domain ilike :p " "or c.id in (select company_id from company_aliases where alias_norm like :n) " "order by c.importance desc, c.display_name limit 6", p=qs + "%", n=normalize_alias(qs) + "%") for r in rows: items.append({"kind": "company", "label": r["display_name"], "sublabel": r["canonical_domain"], "href": f"/company/{r['slug']}", "slug": r["slug"], "country": r["country"]}) for r in _match_rows(await agg.cached_industry_rows(), qs, ("name", "slug"), 3): items.append({"kind": "industry", "label": r["name"], "sublabel": f"{r['companies']} companies", "href": f"/industry/{r['slug']}", "slug": r["slug"]}) for r in _match_rows(await agg.cached_country_rows(), qs, ("name", "code"), 3): items.append({"kind": "country", "label": r["name"], "sublabel": f"{r['companies']} companies", "href": f"/country/{r['slug']}", "code": r["code"]}) ql = qs.upper().replace(" ", "_") for t in EventType: if t.value.startswith(ql) or ql in t.value: items.append({"kind": "event_type", "label": t.value.replace("_", " ").title(), "sublabel": "event type", "href": f"/events?event_type={t.value}", "event_type": t.value}) return {"items": items[:SUGGEST_MAX]} return await cached(f"suggest:{qs.lower()}", 30, produce) @dataclass class AskPlan: """Engine-independent query plan for `/ask` (filled from the intelligence parser when present, else from `ask_fallback`).""" window: str = "30d" event_types: list[str] = field(default_factory=list) event_subtypes: list[str] = field(default_factory=list) country: str | None = None industry: str | None = None ai: bool = False ranking_kind: str | None = None company_terms: list[str] = field(default_factory=list) keywords: list[str] = field(default_factory=list) min_importance: float | None = None interpretation: dict[str, Any] = field(default_factory=dict) engine: str = "deterministic" answer_fn: Any = None def _window_from_days(days: int | None) -> str: if not days: return "30d" for name, td in sorted(q.WINDOWS.items(), key=lambda kv: kv[1]): if days <= td.days or (name == "24h" and days <= 1): return name return "1y" async def _plan_with_intelligence(question: str, countries: list[dict[str, Any]], industries: list[dict[str, Any]]) -> AskPlan | None: try: from companyatlas.services.llm import ask as llm_ask # intelligence agent's module — optional except ImportError: return None route = getattr(llm_ask, "route_question", None) if route is None: return None try: it = route(question, industries={i["slug"]: i["name"] for i in industries}, countries={c["code"]: c["name"] for c in countries}) if inspect.isawaitable(it): it = await it except Exception: log.warning("intelligence ask parser failed; using fallback", exc_info=True) return None kind = None intent = getattr(it, "intent", "") if intent == "trend" or getattr(it, "answer_style", "") == "trend": kind = "most_active" plan = AskPlan(window=_window_from_days(getattr(it, "window_days", None)), event_types=list(getattr(it, "event_types", []) or []), event_subtypes=list(getattr(it, "event_subtypes", []) or []), country=(getattr(it, "countries", None) or [None])[0], industry=(getattr(it, "industries", None) or [None])[0], ai="ai" in (getattr(it, "tags", []) or []) or intent == "ai", ranking_kind=kind, company_terms=list(getattr(it, "companies", []) or []), keywords=list(getattr(it, "keywords", []) or []), min_importance=getattr(it, "min_importance", None), interpretation=it.to_dict() if hasattr(it, "to_dict") else {}, engine=getattr(it, "source", "deterministic")) build = getattr(llm_ask, "build_answer", None) if build is not None: plan.answer_fn = lambda companies, events, titles: build(it, companies=companies, events=events) plan.interpretation.setdefault("filters", getattr(it, "filters", {})) return plan def _plan_with_fallback(question: str, countries: list[dict[str, Any]], industries: list[dict[str, Any]]) -> AskPlan: it = ask_fallback.interpret(question, countries=countries, industries=industries) plan = AskPlan(window=it.window, event_types=list(it.event_types), country=it.country, industry=it.industry, ai=it.ai, ranking_kind=it.ranking_kind, company_terms=it.company_terms, keywords=it.terms, interpretation=it.to_json(), engine="deterministic") plan.answer_fn = lambda companies, events, titles: ask_fallback.compose_answer(it, events_total=events, companies_count=companies, sample_titles=titles) return plan @router.get("/ask", summary="Ask Company Atlas (deterministic parser; LLM refinement optional)") async def ask(response: Response, q_: str = Query(..., alias="q", min_length=2, max_length=300), limit: int = Query(10, ge=1, le=50)) -> dict[str, Any]: response.headers["cache-control"] = public_cache_value(30) question = _clean(q_) async with connection() as conn: countries = await fetch_all(conn, "select code, name from countries") industries = await fetch_all(conn, "select slug, name, keywords from industries") plan = await _plan_with_intelligence(question, countries, industries) or _plan_with_fallback(question, countries, industries) company_filter_id: str | None = None company_cards: list[dict[str, Any]] = [] if plan.company_terms: probe = await _search_companies(conn, " ".join(plan.company_terms[:2]), 3) key = slugify(" ".join(plan.company_terms[:2])) if probe and key in (probe[0]["slug"], slugify(probe[0]["display_name"])): company_filter_id = probe[0]["id"] company_cards = probe[:1] subtypes = plan.event_subtypes or (["AI_HIRING", "AI_LAUNCH"] if (plan.ai and not plan.event_types) else None) where, params = q.event_filters(company_id=company_filter_id, event_types=plan.event_types or None, country=plan.country, industry=plan.industry, since=q.window_start(plan.window), event_subtypes=subtypes if not plan.event_types else None, min_importance=plan.min_importance) if plan.ai and plan.event_types: where.append("(e.event_subtype in ('AI_HIRING', 'AI_LAUNCH') or 'ai' = any(e.tags) or e.search @@ websearch_to_tsquery('english', 'ai'))") events = [ser.event(r) for r in await q.fetch_events(conn, where, params, sort="recent", limit=limit * 2)] total = await q.count_events(conn, where, params) if not company_cards: if plan.ranking_kind: company_cards = await agg.ranking_cards(conn, plan.ranking_kind, plan.window, country=plan.country, industry=plan.industry, limit=limit) else: free_text = " ".join(plan.keywords) if plan.keywords and not plan.event_types else None cw, cp = q.company_filters(country=plan.country, industry=plan.industry, status="ACTIVE", has_events=True if not free_text else None, q=free_text) sort = "hiring" if "HIRING" in plan.event_types and not plan.ai else "activity" ids, _t = await q.company_page_ids(conn, cw, cp, sort=sort, limit=limit) company_cards = [ser.company_card(r) for r in await q.fetch_cards_by_ids(conn, ids, sparkline=True)] if plan.ai: company_cards.sort(key=lambda c: -(c["metrics"].get("ai_adoption") or 0)) distinct_companies = len({e["company"]["id"] for e in events}) if events else 0 answer = plan.answer_fn(distinct_companies or len(company_cards), total, [e["title"] for e in events]) sources = list(dict.fromkeys(e["source_url"] for e in events if e.get("source_url")))[:10] return {"interpretation": plan.interpretation, "answer": answer, "companies": company_cards[:limit], "events": events[:limit], "sources": sources, "events_total": total, "engine": plan.engine}