spb/company-atlas
Public
Python 66.3%
TypeScript 22.7%
JavaScript 8.6%
HTML 1.4%
CSS 0.7%
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