pk-api / app.py
parthislive's picture
Update app.py
2a55e8a verified
Raw History Blame Contribute Delete
12.2 kB
# ============================================================
# HfFolder monkey-patch (MUST be before gradio import)
# ============================================================
import os
try:
import huggingface_hub
if not hasattr(huggingface_hub, "HfFolder"):
class _HfFolder:
@staticmethod
def get_token():
return os.getenv("HF_TOKEN", "")
huggingface_hub.HfFolder = _HfFolder
except ImportError:
pass
# ============================================================
import asyncio
import json
import threading
from concurrent.futures import ThreadPoolExecutor
from typing import Any
import duckdb
import gradio as gr
import httpx
from fastapi import FastAPI, HTTPException, Query, Response
from pydantic import BaseModel
# ── Config ──────────────────────────────────────────────────────────────────
HF_INDEX_BASE = os.environ.get(
"PK_HF_INDEX_BASE",
"https://huggingface.co/datasets/parthislive/pk-db/resolve/main/parts",
).rstrip("/")
PARALLELISM = int(os.environ.get("PK_PARALLEL", "1"))
THREADS_PER_CONN = int(os.environ.get("PK_THREADS_PER_CONN", "2"))
DUPLICATE_CAP = 2
SEARCH_FIELDS = ["NUMBER", "CNIC", "NAME", "ADDRESS"]
NUMBER_FIELDS = ["NUMBER"]
REMOTE_INDEXES = {
"main": [f"{HF_INDEX_BASE}/data_{i}.parquet" for i in range(17)],
}
# ── DuckDB Connection Pool ──────────────────────────────────────────────────
_conns: list = []
_conns_lock = threading.Lock()
_thread_local = threading.local()
pool = ThreadPoolExecutor(max_workers=PARALLELISM, thread_name_prefix="duck")
def _idx_ready(kind: str) -> bool:
return kind in REMOTE_INDEXES
def _new_conn():
con = duckdb.connect()
con.execute("SET home_directory='/tmp'")
con.execute("SET extension_directory='/tmp/duckdb_extensions'")
con.execute("INSTALL parquet; LOAD parquet;")
con.execute("INSTALL httpfs; LOAD httpfs;")
con.execute("SET memory_limit='400MB'")
con.execute("SET enable_object_cache=true")
con.execute("SET http_keep_alive=true")
con.execute("SET http_timeout=60000")
con.execute("SET threads=2")
for kind, urls in REMOTE_INDEXES.items():
view = f"people_{kind}"
lst = ", ".join(f"'{u}'" for u in urls)
con.execute(
f"CREATE OR REPLACE VIEW {view} AS "
f"SELECT NUMBER, CNIC, NAME, ADDRESS FROM read_parquet([{lst}])"
)
return con
def _thread_id() -> int:
tid = getattr(_thread_local, "id", None)
if tid is None:
with _conns_lock:
tid = len(_conns)
_thread_local.id = tid
return tid
def _get_conn():
ident = _thread_id()
with _conns_lock:
while len(_conns) <= ident:
_conns.append(_new_conn())
return _conns[ident]
# ── Dedup ───────────────────────────────────────────────────────────────────
def _person_key(row: dict) -> tuple:
cnic = (row.get("CNIC") or "").strip()
num = str(row.get("NUMBER") or "").strip()
return (cnic, num)
def _connected_numbers(row: dict) -> list:
connected, seen = [], set()
for field in NUMBER_FIELDS:
raw = row.get(field)
if raw is None:
continue
value = str(raw).strip()
if not value or value in seen:
continue
seen.add(value)
connected.append({"field": field, "value": value})
return connected
def _cap_duplicates(rows: list) -> list:
seen, out = {}, []
for r in rows:
k = _person_key(r)
n = seen.get(k, 0)
if n < DUPLICATE_CAP:
seen[k] = n + 1
record = dict(r)
record["connected_numbers"] = _connected_numbers(record)
out.append(record)
return out
# ── Fetch Without Pandas ────────────────────────────────────────────────────
def _fetch_rows(con, sql):
cur = con.execute(sql)
cols = [d[0] for d in cur.description]
rows = cur.fetchall()
return [dict(zip(cols, r)) for r in rows]
# ── Search Logic ────────────────────────────────────────────────────────────
def _run_field_search(field: str, value: str, mode: str, limit: int) -> dict:
if field not in SEARCH_FIELDS:
raise ValueError(f"Unknown field: {field}")
v = str(value).replace("'", "''")
view = "people_main"
cols = "NUMBER, CNIC, NAME, ADDRESS"
if mode == "exact":
if field == "NUMBER":
try:
sql = f"SELECT {cols} FROM {view} WHERE NUMBER = {int(value)} LIMIT {limit * DUPLICATE_CAP + 20}"
except:
return {"field": field, "value": value, "mode": mode, "count": 0, "results": []}
else:
sql = f"SELECT {cols} FROM {view} WHERE {field} = '{v}' LIMIT {limit * DUPLICATE_CAP + 20}"
elif mode == "contains":
v2 = v.replace("%", r"\%").replace("_", r"\_")
sql = f"SELECT {cols} FROM {view} WHERE {field} ILIKE '%{v2}%' ESCAPE '\\' LIMIT {limit * DUPLICATE_CAP + 20}"
else:
raise ValueError(f"Unknown mode: {mode}")
con = _get_conn()
raw_rows = _fetch_rows(con, sql)
results = _cap_duplicates(raw_rows)[:limit]
return {"field": field, "value": value, "mode": mode, "count": len(results), "results": results}
def _unified_search(q: str, limit: int = 10) -> dict:
q = q.strip()
if not q:
return {"query": q, "searched_fields": [], "count": 0, "results": []}
is_num = q.isdigit() and len(q) >= 8
all_rows, searched = [], []
if is_num:
r = _run_field_search("NUMBER", q, "exact", limit)
all_rows.extend(r["results"])
searched.append("NUMBER")
if not all_rows:
r = _run_field_search("CNIC", q, "exact", limit)
all_rows.extend(r["results"])
searched.append("CNIC")
if not all_rows and not is_num:
r = _run_field_search("NAME", q, "contains", limit)
all_rows.extend(r["results"])
searched.append("NAME")
all_rows = _cap_duplicates(all_rows)[:limit]
return {"query": q, "searched_fields": searched, "count": len(all_rows), "results": all_rows}
# ── FastAPI ─────────────────────────────────────────────────────────────────
fastapi_app = FastAPI(title="PK DB Search API")
class BatchRequest(BaseModel):
queries: list
limit: int = 10
@fastapi_app.get("/")
def root():
return {"app": "PK DB Search API", "records": 218_697_206, "columns": SEARCH_FIELDS}
@fastapi_app.get("/health")
def health():
return {"status": "ok"}
@fastapi_app.get("/search")
async def search(
q: str | None = Query(None),
number: str | None = Query(None),
field: str | None = Query(None),
mode: str = Query("exact"),
limit: int = Query(10, ge=1, le=1000),
pretty: bool = Query(True),
):
q_val = (q or number or "").strip()
if not q_val:
raise HTTPException(422, "Provide q or number")
loop = asyncio.get_running_loop()
if field:
data = await loop.run_in_executor(pool, _run_field_search, field, q_val, mode, limit)
else:
data = await loop.run_in_executor(pool, _unified_search, q_val, limit)
result = {"success": bool(data["count"]), **data, "number": q_val, "total": data["count"]}
content = json.dumps(result, indent=2 if pretty else None, ensure_ascii=False)
return Response(content=content, media_type="application/json")
@fastapi_app.get("/api/number")
async def api_number(number: str = Query(...)):
loop = asyncio.get_running_loop()
data = await loop.run_in_executor(pool, _run_field_search, "NUMBER", number, "exact", 10)
return {"status": "success" if data["count"] else "not_found", **data}
@fastapi_app.get("/api/cnic")
async def api_cnic(cnic: str = Query(...)):
loop = asyncio.get_running_loop()
data = await loop.run_in_executor(pool, _run_field_search, "CNIC", cnic, "exact", 10)
return {"status": "success" if data["count"] else "not_found", **data}
@fastapi_app.get("/api/name")
async def api_name(name: str = Query(...)):
loop = asyncio.get_running_loop()
data = await loop.run_in_executor(pool, _run_field_search, "NAME", name, "contains", 50)
return {"status": "success" if data["count"] else "not_found", **data}
@fastapi_app.get("/api/address")
async def api_address(address: str = Query(...)):
loop = asyncio.get_running_loop()
data = await loop.run_in_executor(pool, _run_field_search, "ADDRESS", address, "contains", 50)
return {"status": "success" if data["count"] else "not_found", **data}
# ── Pinger ──────────────────────────────────────────────────────────────────
async def pinger():
port = os.getenv("PORT", "7860")
url = f"http://localhost:{port}/health"
async with httpx.AsyncClient(timeout=10) as client:
while True:
await asyncio.sleep(120)
try:
await client.get(url)
except Exception:
pass
@fastapi_app.on_event("startup")
async def startup_event():
asyncio.create_task(pinger())
# ── Gradio UI ───────────────────────────────────────────────────────────────
def format_result(row: dict) -> str:
lines = []
for field in SEARCH_FIELDS:
val = row.get(field, "")
if val:
lines.append(f"**{field}:** {val}")
return "\n\n".join(lines)
def search_ui(query: str, limit: int) -> str:
if not query or not query.strip():
return "⚠️ Kuch toh search karo β€” NUMBER, CNIC, NAME ya ADDRESS daalo."
q = query.strip()
try:
data = _unified_search(q, int(limit))
except Exception as e:
return f"❌ Error: {str(e)}"
count = data["count"]
results = data["results"]
searched = ", ".join(data.get("searched_fields", []))
if not results:
return f"πŸ” **Query:** `{q}`\n**Searched:** {searched}\n\n❌ **No data found**."
header = f"πŸ” **Query:** `{q}` | **Found:** {count} results | **Searched:** {searched}\n\n---\n\n"
parts = []
for i, row in enumerate(results, 1):
parts.append(f"### Result {i}\n{format_result(row)}")
return header + "\n\n---\n\n".join(parts)
def build_ui():
with gr.Blocks(title="PK DB Search") as demo:
gr.Markdown("# πŸ” PK DB Search")
gr.Markdown("Search **218M Pakistan records**")
with gr.Row():
with gr.Column(scale=3):
query_input = gr.Textbox(label="Search Query", placeholder="NUMBER, CNIC, NAME ya ADDRESS", lines=1)
with gr.Column(scale=1):
limit_slider = gr.Slider(minimum=1, maximum=50, value=10, step=1, label="Max Results")
search_btn = gr.Button("πŸ” Search", variant="primary", size="lg")
output = gr.Markdown(label="Results")
search_btn.click(fn=search_ui, inputs=[query_input, limit_slider], outputs=output)
query_input.submit(fn=search_ui, inputs=[query_input, limit_slider], outputs=output)
return demo
demo = build_ui()
app = gr.mount_gradio_app(fastapi_app, demo, path="/ui")
# ============================================================
# HF Space Gradio SDK Launch (module level, NOT in __main__)
# ============================================================
demo.launch(server_name="0.0.0.0", server_port=7860)