SPB Git forge
28commits 1branches 0releases
7.7 MBsize
maindefault branch
10 days agolast push
Python 66.3% TypeScript 22.7% JavaScript 8.6% HTML 1.4% CSS 0.7%
14.8 KB · 255 lines python
Raw Blame History
1"""Search: `/search` (FTS + trigram), `/search/suggest` (fast prefix), `/ask` (deterministic parser, LLM optional)."""2from __future__ import annotations34import inspect5import logging6import time7from dataclasses import dataclass, field8from typing import Any910from fastapi import APIRouter, Query, Response1112from companyatlas.api import aggregates as agg13from companyatlas.api import ask_fallback14from companyatlas.api import queries as q15from companyatlas.api import serializers as ser16from companyatlas.api.common import cached, public_cache_value17from companyatlas.db import connection, fetch_all18from companyatlas.ids import normalize_alias, slugify19from companyatlas.taxonomy import EventType2021log = logging.getLogger("companyatlas.api.search")22ORDER = 1023router = APIRouter(prefix="/api/v1", tags=["search"])24ALL_TYPES = ("companies", "events", "industries", "countries", "people", "products")25SUGGEST_MAX = 10262728def _clean(qs: str) -> str:29    return " ".join(qs.replace("%", " ").replace("_", " ").split())[:200]303132async def _search_companies(conn: Any, qs: str, limit: int) -> list[dict[str, Any]]:33    if len(qs) >= 3:34        rows = await fetch_all(conn, """35            select c.id, greatest(coalesce(ts_rank(c.search, websearch_to_tsquery('simple', :q)), 0), similarity(c.display_name, :q),36                                  similarity(c.canonical_domain, :q), case when a.company_id is not null then 1.0 else 0 end) as score37            from companies c left join (select distinct company_id from company_aliases where alias_norm = :alias) a on a.company_id = c.id38            where c.search @@ websearch_to_tsquery('simple', :q) or c.display_name % :q or c.canonical_domain % :q or a.company_id is not null39            order by score desc, c.importance desc limit :lim""", q=qs, alias=normalize_alias(qs), lim=limit)40    else:41        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 "42                                     "order by c.importance desc limit :lim", p=qs + "%", lim=limit)43    cards = await q.fetch_cards_by_ids(conn, [r["id"] for r in rows])44    return [ser.company_card(c) for c in cards]454647async def _search_events(conn: Any, qs: str, limit: int) -> list[dict[str, Any]]:48    if len(qs) < 2:49        return []50    rows = await fetch_all(conn, f"{q.EVENT_SELECT} where e.status = 'active' and e.search @@ websearch_to_tsquery('english', :q) "51                                 "order by ts_rank(e.search, websearch_to_tsquery('english', :q)) desc, e.detected_at desc limit :lim", q=qs, lim=limit)52    return [ser.event(r) for r in rows]535455async def _search_people(conn: Any, qs: str, limit: int) -> list[dict[str, Any]]:56    if len(qs) < 3:57        return []58    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, "59                                 "c.country as company_country, c.logo_url as company_logo_url from people p join companies c on c.id = p.company_id "60                                 "where p.name % :q or p.name ilike :like order by similarity(p.name, :q) desc, p.is_executive desc limit :lim",61                           q=qs, like=f"%{qs}%", lim=limit)62    out = []63    for r in rows:64        item = ser.person(r)65        item["company"] = ser.company_ref(r)66        out.append(item)67    return out686970async def _search_products(conn: Any, qs: str, limit: int) -> list[dict[str, Any]]:71    if len(qs) < 3:72        return []73    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, "74                                 "c.country as company_country, c.logo_url as company_logo_url from products p join companies c on c.id = p.company_id "75                                 "where p.name % :q or p.name ilike :like order by similarity(p.name, :q) desc, p.last_seen_at desc limit :lim",76                           q=qs, like=f"%{qs}%", lim=limit)77    out = []78    for r in rows:79        item = ser.product(r)80        item["company"] = ser.company_ref(r)81        out.append(item)82    return out838485def _match_rows(rows: list[dict[str, Any]], qs: str, keys: tuple[str, ...], limit: int) -> list[dict[str, Any]]:86    ql = qs.lower()87    starts = [r for r in rows if any(str(r.get(k) or "").lower().startswith(ql) for k in keys)]88    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)]89    return (starts + contains)[:limit]909192@router.get("/search", summary="Search companies, events, industries, countries, people and products")93async def search(response: Response, q_: str = Query(..., alias="q", min_length=1, max_length=200), types: str = Query(",".join(ALL_TYPES)),94                 limit: int = Query(10, ge=1, le=50)) -> dict[str, Any]:95    response.headers["cache-control"] = public_cache_value(30)96    t0 = time.perf_counter()97    qs = _clean(q_)98    wanted = {t for t in q.csv_list(types) if t in ALL_TYPES} or set(ALL_TYPES)99    out: dict[str, Any] = {"query": qs, "companies": [], "events": [], "industries": [], "countries": [], "people": [], "products": []}100    if qs:101        async with connection() as conn:102            if "companies" in wanted:103                out["companies"] = await _search_companies(conn, qs, limit)104            if "events" in wanted:105                out["events"] = await _search_events(conn, qs, limit)106            if "people" in wanted:107                out["people"] = await _search_people(conn, qs, limit)108            if "products" in wanted:109                out["products"] = await _search_products(conn, qs, limit)110        if "industries" in wanted:111            out["industries"] = _match_rows(await agg.cached_industry_rows(), qs, ("name", "slug"), limit)112        if "countries" in wanted:113            out["countries"] = _match_rows(await agg.cached_country_rows(), qs, ("name", "code"), limit)114    out["took_ms"] = round((time.perf_counter() - t0) * 1000, 1)115    return out116117118@router.get("/search/suggest", summary="Typeahead suggestions (≤ 10)")119async def suggest(response: Response, q_: str = Query(..., alias="q", min_length=1, max_length=100)) -> dict[str, Any]:120    response.headers["cache-control"] = public_cache_value(30)121    qs = _clean(q_)122    if not qs:123        return {"items": []}124125    async def produce() -> dict[str, Any]:126        items: list[dict[str, Any]] = []127        async with connection() as conn:128            rows = await fetch_all(conn, "select c.slug, c.display_name, c.canonical_domain, c.country from companies c "129                                         "where c.display_name ilike :p or c.canonical_domain ilike :p "130                                         "or c.id in (select company_id from company_aliases where alias_norm like :n) "131                                         "order by c.importance desc, c.display_name limit 6", p=qs + "%", n=normalize_alias(qs) + "%")132        for r in rows:133            items.append({"kind": "company", "label": r["display_name"], "sublabel": r["canonical_domain"], "href": f"/company/{r['slug']}", "slug": r["slug"],134                          "country": r["country"]})135        for r in _match_rows(await agg.cached_industry_rows(), qs, ("name", "slug"), 3):136            items.append({"kind": "industry", "label": r["name"], "sublabel": f"{r['companies']} companies", "href": f"/industry/{r['slug']}", "slug": r["slug"]})137        for r in _match_rows(await agg.cached_country_rows(), qs, ("name", "code"), 3):138            items.append({"kind": "country", "label": r["name"], "sublabel": f"{r['companies']} companies", "href": f"/country/{r['slug']}", "code": r["code"]})139        ql = qs.upper().replace(" ", "_")140        for t in EventType:141            if t.value.startswith(ql) or ql in t.value:142                items.append({"kind": "event_type", "label": t.value.replace("_", " ").title(), "sublabel": "event type", "href": f"/events?event_type={t.value}",143                              "event_type": t.value})144        return {"items": items[:SUGGEST_MAX]}145    return await cached(f"suggest:{qs.lower()}", 30, produce)146147148@dataclass149class AskPlan:150    """Engine-independent query plan for `/ask` (filled from the intelligence parser when present, else from `ask_fallback`)."""151    window: str = "30d"152    event_types: list[str] = field(default_factory=list)153    event_subtypes: list[str] = field(default_factory=list)154    country: str | None = None155    industry: str | None = None156    ai: bool = False157    ranking_kind: str | None = None158    company_terms: list[str] = field(default_factory=list)159    keywords: list[str] = field(default_factory=list)160    min_importance: float | None = None161    interpretation: dict[str, Any] = field(default_factory=dict)162    engine: str = "deterministic"163    answer_fn: Any = None164165166def _window_from_days(days: int | None) -> str:167    if not days:168        return "30d"169    for name, td in sorted(q.WINDOWS.items(), key=lambda kv: kv[1]):170        if days <= td.days or (name == "24h" and days <= 1):171            return name172    return "1y"173174175async def _plan_with_intelligence(question: str, countries: list[dict[str, Any]], industries: list[dict[str, Any]]) -> AskPlan | None:176    try:177        from companyatlas.services.llm import ask as llm_ask  # intelligence agent's module — optional178    except ImportError:179        return None180    route = getattr(llm_ask, "route_question", None)181    if route is None:182        return None183    try:184        it = route(question, industries={i["slug"]: i["name"] for i in industries}, countries={c["code"]: c["name"] for c in countries})185        if inspect.isawaitable(it):186            it = await it187    except Exception:188        log.warning("intelligence ask parser failed; using fallback", exc_info=True)189        return None190    kind = None191    intent = getattr(it, "intent", "")192    if intent == "trend" or getattr(it, "answer_style", "") == "trend":193        kind = "most_active"194    plan = AskPlan(window=_window_from_days(getattr(it, "window_days", None)), event_types=list(getattr(it, "event_types", []) or []),195                   event_subtypes=list(getattr(it, "event_subtypes", []) or []), country=(getattr(it, "countries", None) or [None])[0],196                   industry=(getattr(it, "industries", None) or [None])[0], ai="ai" in (getattr(it, "tags", []) or []) or intent == "ai",197                   ranking_kind=kind, company_terms=list(getattr(it, "companies", []) or []), keywords=list(getattr(it, "keywords", []) or []),198                   min_importance=getattr(it, "min_importance", None), interpretation=it.to_dict() if hasattr(it, "to_dict") else {},199                   engine=getattr(it, "source", "deterministic"))200    build = getattr(llm_ask, "build_answer", None)201    if build is not None:202        plan.answer_fn = lambda companies, events, titles: build(it, companies=companies, events=events)203    plan.interpretation.setdefault("filters", getattr(it, "filters", {}))204    return plan205206207def _plan_with_fallback(question: str, countries: list[dict[str, Any]], industries: list[dict[str, Any]]) -> AskPlan:208    it = ask_fallback.interpret(question, countries=countries, industries=industries)209    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,210                   company_terms=it.company_terms, keywords=it.terms, interpretation=it.to_json(), engine="deterministic")211    plan.answer_fn = lambda companies, events, titles: ask_fallback.compose_answer(it, events_total=events, companies_count=companies, sample_titles=titles)212    return plan213214215@router.get("/ask", summary="Ask Company Atlas (deterministic parser; LLM refinement optional)")216async 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]:217    response.headers["cache-control"] = public_cache_value(30)218    question = _clean(q_)219    async with connection() as conn:220        countries = await fetch_all(conn, "select code, name from countries")221        industries = await fetch_all(conn, "select slug, name, keywords from industries")222        plan = await _plan_with_intelligence(question, countries, industries) or _plan_with_fallback(question, countries, industries)223        company_filter_id: str | None = None224        company_cards: list[dict[str, Any]] = []225        if plan.company_terms:226            probe = await _search_companies(conn, " ".join(plan.company_terms[:2]), 3)227            key = slugify(" ".join(plan.company_terms[:2]))228            if probe and key in (probe[0]["slug"], slugify(probe[0]["display_name"])):229                company_filter_id = probe[0]["id"]230                company_cards = probe[:1]231        subtypes = plan.event_subtypes or (["AI_HIRING", "AI_LAUNCH"] if (plan.ai and not plan.event_types) else None)232        where, params = q.event_filters(company_id=company_filter_id, event_types=plan.event_types or None, country=plan.country, industry=plan.industry,233                                        since=q.window_start(plan.window), event_subtypes=subtypes if not plan.event_types else None,234                                        min_importance=plan.min_importance)235        if plan.ai and plan.event_types:236            where.append("(e.event_subtype in ('AI_HIRING', 'AI_LAUNCH') or 'ai' = any(e.tags) or e.search @@ websearch_to_tsquery('english', 'ai'))")237        events = [ser.event(r) for r in await q.fetch_events(conn, where, params, sort="recent", limit=limit * 2)]238        total = await q.count_events(conn, where, params)239        if not company_cards:240            if plan.ranking_kind:241                company_cards = await agg.ranking_cards(conn, plan.ranking_kind, plan.window, country=plan.country, industry=plan.industry, limit=limit)242            else:243                free_text = " ".join(plan.keywords) if plan.keywords and not plan.event_types else None244                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)245                sort = "hiring" if "HIRING" in plan.event_types and not plan.ai else "activity"246                ids, _t = await q.company_page_ids(conn, cw, cp, sort=sort, limit=limit)247                company_cards = [ser.company_card(r) for r in await q.fetch_cards_by_ids(conn, ids, sparkline=True)]248                if plan.ai:249                    company_cards.sort(key=lambda c: -(c["metrics"].get("ai_adoption") or 0))250    distinct_companies = len({e["company"]["id"] for e in events}) if events else 0251    answer = plan.answer_fn(distinct_companies or len(company_cards), total, [e["title"] for e in events])252    sources = list(dict.fromkeys(e["source_url"] for e in events if e.get("source_url")))[:10]253    return {"interpretation": plan.interpretation, "answer": answer, "companies": company_cards[:limit], "events": events[:limit], "sources": sources,254            "events_total": total, "engine": plan.engine}255