File size: 43,863 Bytes
a10a183
 
 
 
 
 
 
 
 
 
 
 
d0a25b4
a10a183
 
 
 
 
 
 
da95441
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e36772a
 
 
 
 
 
 
 
 
a10a183
 
 
 
9ac7863
 
 
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e48acc2
 
 
a10a183
 
 
 
 
 
e48acc2
 
 
a10a183
 
 
 
 
 
e48acc2
a10a183
04d67e9
a10a183
 
 
d0a25b4
a10a183
 
 
 
 
d0a25b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
 
34b0d20
 
 
 
 
9ac7863
 
34b0d20
 
 
9ac7863
 
34b0d20
9ac7863
 
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e36772a
 
 
 
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
e36772a
 
 
 
a10a183
 
 
 
 
 
 
e48acc2
 
 
 
 
 
e36772a
 
 
e48acc2
 
 
a10a183
 
 
 
 
 
 
 
 
 
e36772a
e48acc2
e36772a
 
e48acc2
 
 
 
e36772a
 
 
 
e48acc2
 
 
 
 
a10a183
 
 
 
 
 
b3a08a4
 
 
d0a25b4
 
 
 
b3a08a4
 
 
d0a25b4
 
 
 
 
 
 
 
 
 
 
b3a08a4
 
 
d0a25b4
b3a08a4
 
 
 
d0a25b4
 
 
 
 
 
 
 
 
b3a08a4
 
 
 
 
a10a183
 
 
 
 
 
 
34b0d20
 
 
9ac7863
 
34b0d20
9ac7863
34b0d20
9ac7863
34b0d20
 
9ac7863
 
 
 
 
 
d0a25b4
 
 
 
 
 
 
04aaad8
 
 
 
 
 
 
d0a25b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9ac7863
d0a25b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9ac7863
d0a25b4
 
9ac7863
 
d0a25b4
9ac7863
 
 
 
d0a25b4
 
9ac7863
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d0a25b4
 
 
 
 
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
7528321
 
 
 
a10a183
d0a25b4
a10a183
 
 
d0a25b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a10a183
 
e36772a
a10a183
e36772a
 
 
a10a183
 
 
 
 
e36772a
a10a183
 
e36772a
 
 
a10a183
 
 
e36772a
 
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9ac7863
a10a183
 
 
 
 
 
e48acc2
 
 
 
 
 
 
 
 
 
 
 
 
a10a183
9ac7863
a10a183
 
 
7528321
a10a183
 
e48acc2
 
 
 
9ac7863
a10a183
 
 
d0a25b4
e48acc2
 
 
 
9ac7863
 
d0a25b4
 
 
 
a10a183
 
 
 
 
 
 
 
e48acc2
a10a183
 
 
7528321
a10a183
 
e48acc2
 
 
 
a10a183
 
 
 
 
 
e48acc2
 
 
 
a10a183
 
 
 
 
 
 
 
 
 
 
9ac7863
a10a183
 
 
7528321
a10a183
 
 
 
e48acc2
 
 
 
9ac7863
a10a183
 
 
 
d0a25b4
e48acc2
 
 
 
9ac7863
 
d0a25b4
 
 
 
a10a183
 
 
 
 
 
 
 
e48acc2
a10a183
 
 
7528321
a10a183
 
 
 
e48acc2
 
 
 
a10a183
 
 
 
 
 
 
e48acc2
 
 
 
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e36772a
a10a183
 
 
 
 
 
e36772a
 
 
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e36772a
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e36772a
 
 
a10a183
 
 
 
 
0ac08e4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ae9ddcd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0ac08e4
 
da95441
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b3a08a4
da95441
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a10a183
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
1001
1002
1003
"""hub-search-v3 backend — FastAPI + in-memory DuckDB brute-force vector search.

Drop-in compatible with the old davanstrien/huggingface-datasets-search-v2 API
(same routes / params / response shapes) so the librarian-bots front-end can be
repointed with no code change. Replaces the ChromaDB boot-time HNSW rebuild
(which timed out) with an in-memory DuckDB FLOAT[dim] table loaded once at boot.

See bench: duckdb-vss-bench-2026-07-18. Query encoding replicates the old
backend so query vectors stay aligned with the seed document embeddings.
"""
import logging
import os
import re
import time
from contextlib import asynccontextmanager
from typing import List, Optional

import duckdb
import httpx
from cashews import cache
from fastapi import FastAPI, Header, HTTPException, Query
from fastapi.middleware.cors import CORSMiddleware
from huggingface_hub import HfApi, hf_hub_download, list_repo_files, login
from pydantic import BaseModel

os.environ.setdefault("HF_HUB_ENABLE_HF_TRANSFER", "1")
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
logger = logging.getLogger("hub-search-v3")

# --- config -----------------------------------------------------------------
EMBEDDING_MODEL = "Qwen/Qwen3-Embedding-0.6B"
EMBEDDINGS_REPO = os.getenv("EMBEDDINGS_REPO", "davanstrien/search-v2-embeddings")
EMB_DIM = int(os.getenv("EMB_DIM", "256"))          # MRL truncation dim; 1024 = full
FULL_DIM = 1024
THREADS = int(os.getenv("DUCKDB_THREADS", "2"))
CACHE_TTL = "24h"
TRENDING_CACHE_TTL = "1h"

DS_TASK = "Given a search query, retrieve relevant model and dataset summaries that match the query. "

# ISO-639-3/-2B codes that appear in source cards alongside their two-letter
# equivalents (e.g. datasets tagged `fra` vs `fr`); a two-letter filter matches both.
LANG_EQUIV = {
    "en": ["eng"], "fr": ["fra", "fre"], "de": ["deu", "ger"], "es": ["spa"],
    "zh": ["zho", "chi"], "pt": ["por"], "ru": ["rus"], "ja": ["jpn"],
    "ko": ["kor"], "ar": ["ara"], "it": ["ita"], "nl": ["nld", "dut"],
    "pl": ["pol"], "tr": ["tur"], "vi": ["vie"], "hi": ["hin"], "sv": ["swe"],
}

# config dir names inside the embeddings repo (datasets vs models slices)
DATASET_CFG = os.getenv("DATASET_CONFIG", "dataset_cards")
MODEL_CFG = os.getenv("MODEL_CONFIG", "model_cards")

# prebuilt full-card-text FTS database (hybrid search channel); optional at boot
FTS_FILE = os.getenv("FTS_FILE", "_fts/card-text.duckdb")


def config_files(cfg: str) -> List[str]:
    """Discover the parquet shards for a config dir (shard count varies across seed/v3)."""
    files = [f for f in list_repo_files(EMBEDDINGS_REPO, repo_type="dataset")
             if f.startswith(f"{cfg}/") and f.endswith(".parquet")]
    if not files:
        raise RuntimeError(f"no parquet files found under {cfg}/ in {EMBEDDINGS_REPO}")
    return sorted(files)

HF_TOKEN = os.getenv("HF_TOKEN")
if HF_TOKEN:
    login(token=HF_TOKEN)

cache.setup("mem://", size_limit="1gb")

# in-memory DuckDB (single connection, read-only after boot)
con = duckdb.connect(":memory:")
con.execute(f"SET threads={THREADS}")

_query_model = None
STATE = {"ready": False, "boot_seconds": None, "counts": {}, "revision": None,
         "dim": EMB_DIM, "model": EMBEDDING_MODEL}


def get_query_model():
    global _query_model
    if _query_model is None:
        from sentence_transformers import SentenceTransformer
        logger.info(f"Loading query model {EMBEDDING_MODEL} on CPU")
        _query_model = SentenceTransformer(EMBEDDING_MODEL, device="cpu")
    return _query_model


def embed_query(text: str) -> List[float]:
    """Replicate the old backend's query encoding (prompt_name='query')."""
    model = get_query_model()
    vec = model.encode(text, prompt_name="query", normalize_embeddings=False)
    vec = vec[:EMB_DIM]
    # L2 normalise (matches the stored-vector normalisation; makes cosine stable)
    import numpy as np
    n = np.linalg.norm(vec)
    if n > 0:
        vec = vec / n
    return vec.astype(float).tolist()


def load_table(cfg: str, table: str, has_param: bool):
    """Load one config into an in-memory FLOAT[EMB_DIM] table (MRL trunc + renorm).

    Filters NULL embeddings (v3 refusal rows have no embedding). Discovers shards
    dynamically so the seed and the ~1.17M-row search-v3 dataset both work.
    """
    paths = [hf_hub_download(EMBEDDINGS_REPO, f, repo_type="dataset") for f in config_files(cfg)]
    plist = "[" + ",".join(f"'{p}'" for p in paths) + "]"
    param_sel = "COALESCE(param_count,0)::BIGINT AS param_count" if has_param else "0::BIGINT AS param_count"
    # slice to EMB_DIM, L2-renormalise, cast to fixed-size FLOAT[] array (nested, no correlated subquery)
    # task/license/language are copied straight from the source cards; cast to
    # VARCHAR because the models config's all-null `language` reads back as a
    # NULL-typed column that would otherwise reject `=` filters at query time.
    con.execute(f"""
        CREATE OR REPLACE TABLE {table} AS
        SELECT id, summary,
               COALESCE(likes,0)::BIGINT AS likes,
               COALESCE(downloads,0)::BIGINT AS downloads,
               last_modified,
               CAST(task AS VARCHAR) AS task,
               CAST(license AS VARCHAR) AS license,
               CAST(language AS VARCHAR) AS language,
               {param_sel},
               list_transform(s, x -> (x / nrm))::FLOAT[{EMB_DIM}] AS emb
        FROM (
          SELECT *, sqrt(list_dot_product(s, s)) AS nrm
          FROM (
            SELECT *, embedding[1:{EMB_DIM}] AS s
            FROM read_parquet({plist}, union_by_name=True)
            WHERE embedding IS NOT NULL
            QUALIFY row_number() OVER (PARTITION BY id ORDER BY last_modified DESC NULLS LAST) = 1
          )
        )
    """)
    _refresh_metadata(cfg, table)
    n = con.execute(f"SELECT count(*) FROM {table}").fetchone()[0]
    STATE["counts"][table] = n
    logger.info(f"loaded {table} (from {cfg}): {n:,} rows @ {EMB_DIM}d")


def _refresh_metadata(cfg: str, table: str) -> None:
    """Overwrite likes/downloads from the pipeline's metadata snapshot, if present.

    Corpus rows carry the counts a repo had when its card was last summarised, and
    cards rarely change — so a repo can sit at 0 downloads long after it is widely
    used. Ranking now leans on those counts, so refresh them from the small
    (id, downloads, likes) file the delta pipeline publishes each run. Absent file
    or unreadable snapshot is not fatal: the corpus values still work, just stale.
    """
    try:
        path = hf_hub_download(EMBEDDINGS_REPO, f"_state/metadata/{cfg}.parquet", repo_type="dataset")
    except Exception as e:
        logger.warning(f"{table}: no metadata snapshot ({e}); using corpus counts")
        return
    try:
        n = con.execute(f"""
            UPDATE {table} AS t
            SET downloads = COALESCE(m.downloads, 0)::BIGINT,
                likes     = COALESCE(m.likes, 0)::BIGINT
            FROM read_parquet('{path}') AS m
            WHERE m.id = t.id
        """).fetchone()
        logger.info(f"{table}: refreshed metadata for {(n[0] if n else 0):,} rows")
    except Exception as e:
        logger.warning(f"{table}: metadata refresh failed ({e}); using corpus counts")


@asynccontextmanager
async def lifespan(app: FastAPI):
    t0 = time.time()
    logger.info(f"boot: loading model + embeddings (dim={EMB_DIM})")
    get_query_model()  # load model up front so first query isn't slow
    try:
        one = hf_hub_download(EMBEDDINGS_REPO, config_files(MODEL_CFG)[0], repo_type="dataset")
        cols = con.execute(f"DESCRIBE SELECT * FROM read_parquet('{one}')").fetchall()
        has_param = any(c[0] == "param_count" for c in cols)
    except Exception as e:
        logger.warning(f"param_count detection failed ({e}); assuming absent")
        has_param = False
    load_table(DATASET_CFG, "dataset_cards", has_param=False)
    load_table(MODEL_CFG, "model_cards", has_param=has_param)
    # Optional: open the prebuilt card-text FTS db on its OWN connection.
    # (Not ATTACHed: the fts match_bm25 macro references its internal schema
    # unqualified, which breaks across catalog boundaries.) Any failure
    # degrades hybrid search to vector-only rather than blocking boot.
    global fts_con
    try:
        fts_path = hf_hub_download(EMBEDDINGS_REPO, FTS_FILE, repo_type="dataset")
        fts_con = duckdb.connect(fts_path, read_only=True)
        fts_con.execute("INSTALL fts; LOAD fts;")
        STATE["fts_cards"] = fts_con.execute("SELECT count(*) FROM cards").fetchone()[0]
        logger.info(f"card-text FTS attached: {STATE['fts_cards']:,} cards")
    except Exception as e:
        fts_con = None
        STATE["fts_cards"] = None
        logger.warning(f"card-text FTS unavailable ({e}); hybrid=vector-only")
    try:
        info = HfApi().dataset_info(EMBEDDINGS_REPO)
        STATE["revision"] = info.sha[:8] if info.sha else None
        STATE["last_modified"] = str(getattr(info, "last_modified", None))
    except Exception:
        pass
    STATE["boot_seconds"] = round(time.time() - t0, 1)
    STATE["ready"] = True
    logger.info(f"boot complete in {STATE['boot_seconds']}s; ready")
    yield
    await cache.close()


app = FastAPI(title="hub-search-v3", lifespan=lifespan)
app.add_middleware(
    CORSMiddleware,
    allow_origin_regex=r"https://.*\.(hf\.space|huggingface\.co)",
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)


# --- response models (identical to old backend) -----------------------------
class QueryResult(BaseModel):
    dataset_id: str
    similarity: float
    summary: str
    likes: int
    downloads: int
    task: Optional[str] = None
    license: Optional[str] = None
    language: Optional[str] = None
    last_modified: Optional[str] = None


class QueryResponse(BaseModel):
    results: List[QueryResult]


class ModelQueryResult(BaseModel):
    model_id: str
    similarity: float
    summary: str
    likes: int
    downloads: int
    param_count: Optional[int] = None
    task: Optional[str] = None
    license: Optional[str] = None
    language: Optional[str] = None
    last_modified: Optional[str] = None


class ModelQueryResponse(BaseModel):
    results: List[ModelQueryResult]


# --- core search -------------------------------------------------------------
def _where(min_likes, min_downloads, min_param=0, max_param=None, use_param=False,
           task=None, license=None, language=None, modified_after=None):
    """Build a WHERE clause + its bind params.

    Numeric bounds are inlined (ints are injection-safe); the free-text metadata
    filters (task/license/language/modified_after) are bound as `?` params.
    task/language are comma-joined multi-value strings in the source cards, so
    they match by token membership, not string equality; two-letter language
    codes also match their ISO-639-3 equivalents (LANG_EQUIV).
    Returns ``(where_sql, params)`` for ``_knn``.
    """
    conds, params = [], []
    if min_likes > 0:
        conds.append(f"likes >= {int(min_likes)}")
    if min_downloads > 0:
        conds.append(f"downloads >= {int(min_downloads)}")
    if use_param:
        conds.append("param_count > 0")
        if min_param > 0:
            conds.append(f"param_count >= {int(min_param)}")
        if max_param is not None:
            conds.append(f"param_count <= {int(max_param)}")
    token_match = "list_contains(string_split(replace(lower({col}), ' ', ''), ','), ?)"
    if task:
        conds.append(token_match.format(col="task"))
        params.append(task.lower())
    if license:
        conds.append("license = ?")
        params.append(license)
    if language:
        variants = [language.lower()] + LANG_EQUIV.get(language.lower(), [])
        ors = " OR ".join(token_match.format(col="language") for _ in variants)
        conds.append(f"({ors})")
        params.extend(variants)
    if modified_after:
        # last_modified is an ISO-8601 string; lexicographic >= is chronological.
        conds.append("last_modified >= ?")
        params.append(modified_after)
    return (("WHERE " + " AND ".join(conds)) if conds else ""), params


def _vec_literal(vec):
    return "[" + ",".join(repr(float(x)) for x in vec) + f"]::FLOAT[{EMB_DIM}]"


DUP_SIM = 0.985  # near-identical cards (mirrors/re-uploads) collapse to the top-ranked copy


DUP_OVERFETCH = 6  # candidates per requested row, so dup-collapse can be backfilled
_DL_IDX = 3        # position of `downloads` within COLS (id, summary, likes, downloads, ...)


def _knn(table, qvec, k, where_sql, cols, params=None, return_emb=False):
    """Top-k by cosine, with near-duplicate collapse.

    Overfetches DUP_OVERFETCH x, then greedily drops rows whose embedding is
    ~identical (cos >= DUP_SIM) to a higher-ranked kept row — mirror repos
    otherwise eat adjacent result slots. When a duplicate is more downloaded than
    the copy already kept, it replaces it in place: the group keeps its ranking
    position but is represented by the canonical repo rather than whichever
    mirror happened to score a hair higher.

    The overfetch has to exceed the dup rate or collapse silently returns fewer
    than k (popular model families are dominated by near-identical conversion
    cards); the LIMIT is cheap because the full cosine scan dominates either way.
    Returns (cols..., sim) tuples; with return_emb=True, (cols..., emb, sim).
    """
    import numpy as np
    q = f"""SELECT {cols}, emb, array_cosine_similarity(emb, {_vec_literal(qvec)}) AS sim
            FROM {table} {where_sql} ORDER BY sim DESC LIMIT {int(k) * DUP_OVERFETCH}"""
    rows = con.execute(q, params or []).fetchall()
    kept, kept_vecs = [], []
    for r in rows:
        v = np.asarray(r[-2], dtype=np.float32)
        if kept_vecs:
            sims = np.stack(kept_vecs) @ v
            j = int(np.argmax(sims))
            if float(sims[j]) >= DUP_SIM:
                # Same card, different repo: keep whichever copy people actually use.
                if (r[_DL_IDX] or 0) > (kept[j][_DL_IDX] or 0):
                    kept[j] = r if return_emb else r[:-2] + (r[-1],)
                    kept_vecs[j] = v
                continue
        kept.append(r if return_emb else r[:-2] + (r[-1],))
        kept_vecs.append(v)
        if len(kept) >= int(k):
            break
    return kept


def _emb_of(table, repo_id):
    row = con.execute(f"SELECT emb FROM {table} WHERE id = ?", [repo_id]).fetchone()
    return row[0] if row else None


fts_con = None


def _bm25_ids(kind: str, text: str, n: int) -> list:
    """Top-n repo ids by BM25 over full card text (empty if FTS unavailable)."""
    if fts_con is None:
        return []
    rows = fts_con.execute(
        """SELECT id FROM (
               SELECT id, fts_main_cards.match_bm25(doc_id, ?) AS s
               FROM cards WHERE repo_type = ?
           ) WHERE s IS NOT NULL ORDER BY s DESC LIMIT ?""",
        [text, kind, int(n)],
    ).fetchall()
    return [r[0] for r in rows]


# A query that is short and has no spaces is a repo name, not a description
# ("iSAID", "rtdetr"). Summary embeddings cannot retrieve those: the summary of
# ariG23498/iSAID never says "iSAID", and "iSAID" alone retrieves Icelandic
# census records via ID-ish subword tokens. Matching the id directly is the only
# channel that can find them.
_NAME_MAX_TOKENS = 3   # "iSAID instance segmentation" still carries a name token
_NAME_MIN_CHARS = 4    # shorter tokens ("ocr", "sam") match far too much
# Equal to the vector channel. RRF scores a rank-r hit at weight/(60+r), so the
# vector pool (k*10) reaches ~1/140 at its tail while a name-only hit sits at
# weight/60. At 0.5 such a hit loses to every vector result down to rank ~60 and so
# can never enter the top k — which is the entire point of the channel. False
# positives are held back at the trigger (identifier-shaped tokens, matched at an
# id boundary) rather than by de-weighting the whole channel.
NAME_CHANNEL_WEIGHT = 1.0

# Hyphenated descriptors that pass the identifier test but name no repo.
_NAME_STOPWORDS = frozenset({"real-time", "state-of-the-art", "open-source", "fine-tuned"})


def _is_identifier_like(tok: str) -> bool:
    """True for tokens shaped like a repo/model name rather than an English word.

    Names carry a digit (yolov10), an internal capital (iSAID, RF-DETR), or a
    separator (rf-detr). Ordinary description words ("herbarium", "element") carry
    none of these, and matching those against ids returns download-ranked noise.
    """
    if tok.lower() in _NAME_STOPWORDS:
        return False
    return bool(
        any(c.isdigit() for c in tok)
        or any(c in "-_." for c in tok)
        or re.search(r"[a-z][A-Z]|^[A-Z]{2,}$", tok)
    )


def _name_tokens(query: str) -> list:
    """Tokens from the query that plausibly name a repo rather than describe one."""
    if len(query.strip()) > 60 or len(query.split()) > _NAME_MAX_TOKENS:
        return []
    return [
        tok
        for raw in re.split(r"[\s,]+", query.strip())
        if len(tok := raw.strip("\"'()[]")) >= _NAME_MIN_CHARS and _is_identifier_like(tok)
    ]


def _name_ids(table, tokens, n, where_sql, wparams):
    """Top-n ids whose name contains any of these tokens, most-downloaded first.

    The token must start at an id boundary (start, or after / _ . -). A bare
    substring match puts `RASSAISAID/finance-deepseek-prompts` above the real
    `ariG23498/iSAID` for the query "iSAID". The trailing side is deliberately left
    open so name variants still match — "yolov10" has to find `jameslahm/yolov10x`.

    Ordering by downloads (not cosine) is deliberate: this channel exists to find
    the repo actually named, and RRF only uses the rank, so the canonical copy of
    a name should lead its own channel.
    """
    if not tokens:
        return []
    cond = f"AND {where_sql[6:]}" if where_sql else ""
    ors = " OR ".join("regexp_matches(id, ?, 'i')" for _ in tokens)
    pats = [f"(^|[/_.-]){re.escape(t)}" for t in tokens]
    rows = con.execute(
        f"""SELECT id FROM {table}
            WHERE ({ors}) {cond}
            ORDER BY downloads DESC LIMIT {int(n)}""",
        pats + list(wparams or []),
    ).fetchall()
    return [r[0] for r in rows]


def _fuse(table, qvec, raw, extra_ids, k, where_sql, wparams, weight=1.0):
    """RRF-fuse the vector KNN rows with a second channel's id list.

    Channel-only hits are fetched from the in-memory table (with their true cosine
    similarity) and must pass the same metadata filters as the vector channel.
    `weight` scales the second channel's contribution, so a speculative channel can
    rescue a miss without outvoting the vector ranking.
    Returns rows in fused order, same tuple shape as _knn output.
    """
    if not extra_ids:
        return raw
    rrf: dict = {}
    for rank, rid in enumerate(r[0] for r in raw):
        rrf[rid] = rrf.get(rid, 0.0) + 1.0 / (60 + rank)
    for rank, rid in enumerate(extra_ids):
        rrf[rid] = rrf.get(rid, 0.0) + weight / (60 + rank)
    order = sorted(rrf, key=rrf.get, reverse=True)
    known = {r[0]: r for r in raw}
    missing = [rid for rid in order if rid not in known][: k * 2]
    if missing:
        ph = ",".join("?" for _ in missing)
        cond = f"AND {where_sql[6:]}" if where_sql else ""
        extra = con.execute(
            f"""SELECT {COLS}, array_cosine_similarity(emb, {_vec_literal(qvec)}) AS sim
                FROM {table} WHERE id IN ({ph}) {cond}""",
            missing + list(wparams or []),
        ).fetchall()
        known.update({r[0]: r for r in extra})
    return [known[rid] for rid in order if rid in known]


def _hybrid_fuse(table, kind, query, qvec, raw, k, where_sql, wparams):
    """RRF-fuse vector KNN rows with a BM25 channel over full card text."""
    return _fuse(table, qvec, raw, _bm25_ids(kind, query, k * 2), k, where_sql, wparams)


async def _sort_results(rows, id_field, sort_by, k):
    """rows: list of dicts with keys id, similarity, summary, likes, downloads, [param_count]."""
    if sort_by == "trending":
        scores = {}
        tasks = [get_trending_score(r["_id"], id_field) for r in rows]
        import asyncio
        vals = await asyncio.gather(*tasks)
        scores = {r["_id"]: v for r, v in zip(rows, vals)}
        rows.sort(key=lambda r: scores.get(r["_id"], 0), reverse=True)
        rows = rows[:k]
    elif sort_by in ("likes", "downloads"):
        rows.sort(key=lambda r: r[sort_by], reverse=True)
        rows = rows[:k]
    elif sort_by == "updated":
        # last_modified is an ISO-8601 string; lexicographic sort is chronological
        rows.sort(key=lambda r: r.get("last_modified") or "", reverse=True)
        rows = rows[:k]
    else:
        rows = _apply_popularity_prior(rows)[:k]
    return rows


# Optional log-scaled download term, added to cosine to break ties toward the copy
# people actually use (a 0-download mirror otherwise outranks the canonical repo,
# since their cards read the same).
#
# OFF by default, deliberately. On hf-find's eval/queries.jsonl it looks like a
# clear win (mean recall 0.363 -> 0.435, zero-download slots 28.3% -> 6.2% at 0.08),
# but that eval's expected repos were chosen *because* they are canonical, so it
# cannot help but reward a popularity prior. A blind 3-judge A/B over 24 unseen
# queries split by domain instead:
#
#     mainstream queries          prior 14 - 2 baseline
#     niche / historical / GLAM   prior  3 - 5 baseline
#
# The mechanism: a 150k-download repo gets the full weight while a 50-download
# exact match gets ~a third of it, so any similarity gap under ~0.054 flips — and
# niche queries live in exactly that regime. Concretely, at 0.08 the query "OCR for
# historical printed text" ranks generic dots.ocr-1.5 above
# wjbmattingly/LightOnOCR-2-1B-cultural-heritage-english.
#
# So this is a domain trade, not a free improvement. Enable per-deployment with
# POPULARITY_WEIGHT=0.08 if the audience is mainstream model/dataset discovery.
POPULARITY_WEIGHT = float(os.getenv("POPULARITY_WEIGHT", "0"))


def _rerank_pool(k: int, sort_by: str) -> int:
    """Candidate rows to pull before re-ranking down to k.

    The relevance path re-ranks (popularity prior, optional name fusion) and gains
    from a deeper pool: on eval/queries.jsonl mean recall is 0.363 at k*1, 0.399 at
    k*4 and 0.435 at k*10, flat thereafter. The cosine scan touches every row
    regardless, so a larger LIMIT is close to free.

    The explicit sort modes keep the old k*4: they sort hard by likes/downloads, so
    widening the pool would trade relevance for popularity rather than break ties.
    """
    return k * (10 if sort_by == "similarity" else 4)


def _apply_popularity_prior(rows):
    """Re-rank by similarity + a log-scaled, per-result-set-normalised download term."""
    if POPULARITY_WEIGHT <= 0 or not rows:
        return rows
    import math
    mx = max((r.get("downloads") or 0) for r in rows)
    if mx <= 0:
        return rows
    denom = math.log1p(mx)
    return sorted(
        rows,
        key=lambda r: r["similarity"] + POPULARITY_WEIGHT * (math.log1p(r.get("downloads") or 0) / denom),
        reverse=True,
    )


def _rows_to_dataset(raw):
    out = []
    for id_, summary, likes, downloads, param, task, license_, language, last_modified, sim in raw:
        out.append({"_id": id_, "dataset_id": id_, "similarity": float(sim),
                    "summary": summary, "likes": int(likes), "downloads": int(downloads),
                    "task": task, "license": license_, "language": language,
                    "last_modified": last_modified})
    return out


def _rows_to_model(raw):
    out = []
    for id_, summary, likes, downloads, param, task, license_, language, last_modified, sim in raw:
        out.append({"_id": id_, "model_id": id_, "similarity": float(sim),
                    "summary": summary, "likes": int(likes), "downloads": int(downloads),
                    "param_count": int(param),
                    "task": task, "license": license_, "language": language,
                    "last_modified": last_modified})
    return out


META_COLS = "task, license, language, CAST(last_modified AS VARCHAR) AS last_modified"
COLS = f"id, summary, likes, downloads, param_count, {META_COLS}"


# --- routes ------------------------------------------------------------------
@app.get("/")
async def root():
    from fastapi.responses import RedirectResponse
    return RedirectResponse(url="/docs")


@app.get("/health")
async def health():
    return {
        "ready": STATE["ready"],
        "boot_seconds": STATE["boot_seconds"],
        "counts": STATE["counts"],
        "index_revision": STATE.get("revision"),
        "index_last_modified": STATE.get("last_modified"),
        "fts_cards": STATE.get("fts_cards"),
        "embedding_model": STATE["model"],
        "dim": STATE["dim"],
        "embeddings_repo": EMBEDDINGS_REPO,
    }


@app.get("/lookup/{repo_id:path}")
async def lookup(repo_id: str):
    """Resolve a repo id to its type in one call: {"id":..., "type": "dataset"|"model"}.

    Saves clients the 404-fallthrough (try datasets, then models). 404 with a JSON
    detail if the id is in neither index. `:path` so `org/name` ids keep their slash.
    """
    for table, typ in (("dataset_cards", "dataset"), ("model_cards", "model")):
        if con.execute(f"SELECT 1 FROM {table} WHERE id = ? LIMIT 1", [repo_id]).fetchone():
            return {"id": repo_id, "type": typ}
    raise HTTPException(status_code=404, detail=f"'{repo_id}' not found as a dataset or model in the index")


@app.get("/search/datasets", response_model=QueryResponse)
@cache(ttl=CACHE_TTL, key="search:d:{query}:{k}:{sort_by}:{min_likes}:{min_downloads}:{task}:{license}:{language}:{modified_after}:{hybrid}")
async def search_datasets(
    query: str,
    k: int = Query(default=5, ge=1, le=100),
    sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending", "updated"]),
    min_likes: int = Query(default=0, ge=0),
    min_downloads: int = Query(default=0, ge=0),
    task: Optional[str] = Query(default=None),
    license: Optional[str] = Query(default=None),
    language: Optional[str] = Query(default=None),
    modified_after: Optional[str] = Query(default=None),
    hybrid: bool = Query(default=False),
):
    try:
        qvec = embed_query(f"Instruct: {DS_TASK}\nQuery:{query}")
        n = _rerank_pool(k, sort_by)
        where_sql, wparams = _where(min_likes, min_downloads,
                                    task=task, license=license, language=language,
                                    modified_after=modified_after)
        raw = _knn("dataset_cards", qvec, n, where_sql, COLS, wparams)
        if hybrid:
            raw = _hybrid_fuse("dataset_cards", "datasets", query, qvec, raw, k, where_sql, wparams)
        elif (name_toks := _name_tokens(query)):
            raw = _fuse("dataset_cards", qvec, raw,
                        _name_ids("dataset_cards", name_toks, k, where_sql, wparams),
                        k, where_sql, wparams, weight=NAME_CHANNEL_WEIGHT)
        rows = await _sort_results(_rows_to_dataset(raw), "dataset", sort_by, k)
        return QueryResponse(results=[QueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
    except Exception as e:
        logger.error(f"search datasets: {e}")
        raise HTTPException(status_code=500, detail="Search failed")


@app.get("/similarity/datasets", response_model=QueryResponse)
@cache(ttl=CACHE_TTL, key="sim:d:{dataset_id}:{k}:{sort_by}:{min_likes}:{min_downloads}:{task}:{license}:{language}:{modified_after}")
async def find_similar_datasets(
    dataset_id: str,
    k: int = Query(default=5, ge=1, le=100),
    sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending", "updated"]),
    min_likes: int = Query(default=0, ge=0),
    min_downloads: int = Query(default=0, ge=0),
    task: Optional[str] = Query(default=None),
    license: Optional[str] = Query(default=None),
    language: Optional[str] = Query(default=None),
    modified_after: Optional[str] = Query(default=None),
):
    emb = _emb_of("dataset_cards", dataset_id)
    if emb is None:
        raise HTTPException(status_code=404, detail=f"Dataset ID '{dataset_id}' not found")
    try:
        n = k * 4 if sort_by != "similarity" else k + 1
        where_sql, wparams = _where(min_likes, min_downloads,
                                    task=task, license=license, language=language,
                                    modified_after=modified_after)
        raw = _knn("dataset_cards", emb, n, where_sql, COLS, wparams)
        rows = [r for r in _rows_to_dataset(raw) if r["_id"] != dataset_id]
        rows = await _sort_results(rows, "dataset", sort_by, k)
        return QueryResponse(results=[QueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
    except HTTPException:
        raise
    except Exception as e:
        logger.error(f"similarity datasets: {e}")
        raise HTTPException(status_code=500, detail="Similarity search failed")


@app.get("/search/models", response_model=ModelQueryResponse)
@cache(ttl=CACHE_TTL, key="search:m:{query}:{k}:{sort_by}:{min_likes}:{min_downloads}:{min_param_count}:{max_param_count}:{task}:{license}:{language}:{modified_after}:{hybrid}")
async def search_models(
    query: str,
    k: int = Query(default=5, ge=1, le=100),
    sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending", "updated"]),
    min_likes: int = Query(default=0, ge=0),
    min_downloads: int = Query(default=0, ge=0),
    min_param_count: int = Query(default=0, ge=0),
    max_param_count: Optional[int] = Query(default=None, ge=0),
    task: Optional[str] = Query(default=None),
    license: Optional[str] = Query(default=None),
    language: Optional[str] = Query(default=None),
    modified_after: Optional[str] = Query(default=None),
    hybrid: bool = Query(default=False),
):
    try:
        use_param = min_param_count > 0 or max_param_count is not None
        qvec = embed_query(f"search_query: {query}")
        n = _rerank_pool(k, sort_by)
        where_sql, wparams = _where(min_likes, min_downloads, min_param_count, max_param_count, use_param,
                                    task=task, license=license, language=language,
                                    modified_after=modified_after)
        raw = _knn("model_cards", qvec, n, where_sql, COLS, wparams)
        if hybrid:
            raw = _hybrid_fuse("model_cards", "models", query, qvec, raw, k, where_sql, wparams)
        elif (name_toks := _name_tokens(query)):
            raw = _fuse("model_cards", qvec, raw,
                        _name_ids("model_cards", name_toks, k, where_sql, wparams),
                        k, where_sql, wparams, weight=NAME_CHANNEL_WEIGHT)
        rows = await _sort_results(_rows_to_model(raw), "model", sort_by, k)
        return ModelQueryResponse(results=[ModelQueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
    except Exception as e:
        logger.error(f"search models: {e}")
        raise HTTPException(status_code=500, detail="Model search failed")


@app.get("/similarity/models", response_model=ModelQueryResponse)
@cache(ttl=CACHE_TTL, key="sim:m:{model_id}:{k}:{sort_by}:{min_likes}:{min_downloads}:{min_param_count}:{max_param_count}:{task}:{license}:{language}:{modified_after}")
async def find_similar_models(
    model_id: str,
    k: int = Query(default=5, ge=1, le=100),
    sort_by: str = Query(default="similarity", enum=["similarity", "likes", "downloads", "trending", "updated"]),
    min_likes: int = Query(default=0, ge=0),
    min_downloads: int = Query(default=0, ge=0),
    min_param_count: int = Query(default=0, ge=0),
    max_param_count: Optional[int] = Query(default=None, ge=0),
    task: Optional[str] = Query(default=None),
    license: Optional[str] = Query(default=None),
    language: Optional[str] = Query(default=None),
    modified_after: Optional[str] = Query(default=None),
):
    emb = _emb_of("model_cards", model_id)
    if emb is None:
        raise HTTPException(status_code=404, detail=f"Model ID '{model_id}' not found")
    try:
        use_param = min_param_count > 0 or max_param_count is not None
        n = k * 4 if sort_by != "similarity" else k + 1
        where_sql, wparams = _where(min_likes, min_downloads, min_param_count, max_param_count, use_param,
                                    task=task, license=license, language=language,
                                    modified_after=modified_after)
        raw = _knn("model_cards", emb, n, where_sql, COLS, wparams)
        rows = [r for r in _rows_to_model(raw) if r["_id"] != model_id]
        rows = await _sort_results(rows, "model", sort_by, k)
        return ModelQueryResponse(results=[ModelQueryResult(**{k2: v for k2, v in r.items() if k2 != "_id"}) for r in rows])
    except HTTPException:
        raise
    except Exception as e:
        logger.error(f"similarity models: {e}")
        raise HTTPException(status_code=500, detail="Model similarity search failed")


# --- trending (proxies HF API, joins summaries from the in-memory table) -----
@cache(ttl=TRENDING_CACHE_TTL, key="tscore:{item_type}:{item_id}")
async def get_trending_score(item_id: str, item_type: str) -> float:
    try:
        endpoint = "models" if item_type == "model" else "datasets"
        async with httpx.AsyncClient(timeout=10) as c:
            r = await c.get(f"https://huggingface.co/api/{endpoint}/{item_id}?expand=trendingScore")
            r.raise_for_status()
            return r.json().get("trendingScore", 0)
    except Exception:
        return 0


@app.get("/trending/datasets", response_model=QueryResponse)
@cache(ttl=TRENDING_CACHE_TTL, key="trend:d:{limit}:{min_likes}:{min_downloads}")
async def trending_datasets(
    limit: int = Query(default=10, ge=1, le=100),
    min_likes: int = Query(default=0, ge=0),
    min_downloads: int = Query(default=0, ge=0),
):
    try:
        async with httpx.AsyncClient(timeout=15) as c:
            r = await c.get("https://huggingface.co/api/datasets?sort=trendingScore&limit=200")
            r.raise_for_status()
            items = r.json()
    except Exception as e:
        logger.error(f"trending datasets fetch: {e}")
        raise HTTPException(status_code=502, detail="Failed to fetch trending datasets")
    items = [d for d in items if d.get("likes", 0) >= min_likes and d.get("downloads", 0) >= min_downloads]
    ids = [d["id"] for d in items[: limit * 3]]
    if not ids:
        return QueryResponse(results=[])
    placeholders = ",".join("?" for _ in ids)
    found = {r[0]: r for r in con.execute(
        f"SELECT id, summary, likes, downloads, {META_COLS} FROM dataset_cards WHERE id IN ({placeholders})", ids
    ).fetchall()}
    out = []
    for d in items:
        row = found.get(d["id"])
        if row:
            out.append(QueryResult(dataset_id=d["id"], similarity=1.0, summary=row[1],
                                   likes=d.get("likes", 0), downloads=d.get("downloads", 0),
                                   task=row[4], license=row[5], language=row[6],
                                   last_modified=row[7]))
        if len(out) >= limit:
            break
    return QueryResponse(results=out)


@app.get("/trending/models", response_model=ModelQueryResponse)
@cache(ttl=TRENDING_CACHE_TTL, key="trend:m:{limit}:{min_likes}:{min_downloads}:{min_param_count}:{max_param_count}")
async def trending_models(
    limit: int = Query(default=10, ge=1, le=100),
    min_likes: int = Query(default=0, ge=0),
    min_downloads: int = Query(default=0, ge=0),
    min_param_count: int = Query(default=0, ge=0),
    max_param_count: Optional[int] = Query(default=None, ge=0),
):
    try:
        async with httpx.AsyncClient(timeout=15) as c:
            r = await c.get("https://huggingface.co/api/models?sort=trendingScore&limit=200")
            r.raise_for_status()
            items = r.json()
    except Exception as e:
        logger.error(f"trending models fetch: {e}")
        raise HTTPException(status_code=502, detail="Failed to fetch trending models")
    items = [m for m in items if m.get("likes", 0) >= min_likes and m.get("downloads", 0) >= min_downloads]
    ids = [m["id"] for m in items[: limit * 3]]
    if not ids:
        return ModelQueryResponse(results=[])
    placeholders = ",".join("?" for _ in ids)
    found = {r[0]: r for r in con.execute(
        f"SELECT id, summary, likes, downloads, param_count, {META_COLS} FROM model_cards WHERE id IN ({placeholders})", ids
    ).fetchall()}
    use_param = min_param_count > 0 or max_param_count is not None
    out = []
    for m in items:
        row = found.get(m["id"])
        if not row:
            continue
        param = int(row[4])
        if use_param:
            if param == 0:
                continue
            if min_param_count > 0 and param < min_param_count:
                continue
            if max_param_count is not None and param > max_param_count:
                continue
        out.append(ModelQueryResult(model_id=m["id"], similarity=1.0, summary=row[1],
                                    likes=m.get("likes", 0), downloads=m.get("downloads", 0),
                                    param_count=param,
                                    task=row[5], license=row[6], language=row[7],
                                    last_modified=row[8]))
        if len(out) >= limit:
            break
    return ModelQueryResponse(results=out)


# --- embedding primitives (personal-feeds prototype + compare tooling) -------
class EmbBatchReq(BaseModel):
    repo_type: str
    ids: List[str]


@app.post("/embeddings_batch")
async def embeddings_batch(req: EmbBatchReq):
    """Return stored embeddings for up to 200 repo ids (e.g. a user's likes)."""
    if req.repo_type not in ("datasets", "models"):
        raise HTTPException(status_code=422, detail="repo_type must be datasets|models")
    table = "dataset_cards" if req.repo_type == "datasets" else "model_cards"
    ids = req.ids[:200]
    if not ids:
        return {"embeddings": {}}
    placeholders = ",".join("?" for _ in ids)
    rows = con.execute(
        f"SELECT id, emb FROM {table} WHERE id IN ({placeholders})", ids
    ).fetchall()
    return {"embeddings": {r[0]: [float(x) for x in r[1]] for r in rows}}


@app.get("/embed_query")
@cache(ttl=CACHE_TTL, key="embq:{repo_type}:{text}")
async def embed_query_route(
    text: str = Query(min_length=3, max_length=200),
    repo_type: str = Query(default="datasets"),
):
    """Embed a free-text topic phrase with the serving query encoder.

    repo_type picks the query prefix the corpus was embedded against
    (datasets use the instruct prompt, models the search_query prefix).
    """
    if repo_type not in ("datasets", "models"):
        raise HTTPException(status_code=422, detail="repo_type must be datasets|models")
    prompt = (f"Instruct: {DS_TASK}\nQuery:{text}" if repo_type == "datasets"
              else f"search_query: {text}")
    return {"embedding": embed_query(prompt)}


class VectorSearchReq(BaseModel):
    repo_type: str
    vector: List[float]
    k: int = 20
    exclude_ids: List[str] = []
    include_embeddings: bool = False


@app.post("/search_vector")
async def search_vector(req: VectorSearchReq):
    """KNN search with a caller-supplied (already 256-d, normalised) vector.

    Powers relevance-feedback loops: refine a preference vector client-side
    from votes, then pull fresh results for it. exclude_ids drops already-seen
    items server-side so refinement surfaces new material.
    """
    if req.repo_type not in ("datasets", "models"):
        raise HTTPException(status_code=422, detail="repo_type must be datasets|models")
    if len(req.vector) != EMB_DIM:
        raise HTTPException(status_code=422, detail=f"vector must be {EMB_DIM}-d")
    table = "dataset_cards" if req.repo_type == "datasets" else "model_cards"
    k = max(1, min(int(req.k), 100))
    n = k + len(req.exclude_ids[:300])
    raw = _knn(table, req.vector, n, "", COLS, return_emb=req.include_embeddings)
    excl = set(req.exclude_ids[:300])
    out = []
    for row in raw:
        emb_val = row[-2] if req.include_embeddings else None
        (id_, summary, likes, downloads, param, task, license_, language, last_modified) = row[:9]
        if id_ in excl:
            continue
        item = {"id": id_, "summary": summary, "likes": int(likes),
                "downloads": int(downloads), "param_count": int(param),
                "task": task, "license": license_, "language": language,
                "last_modified": last_modified, "similarity": float(row[-1])}
        if req.include_embeddings:
            item["embedding"] = [float(x) for x in emb_val]
        out.append(item)
        if len(out) >= k:
            break
    return {"results": out}


# --- per-user feeds store (OAuth token -> whoami -> one JSON doc per user) ---
FEEDS_STORE_REPO = os.getenv("FEEDS_STORE_REPO", "davanstrien/hub-feeds-store")
FEEDS_DOC_LIMIT = 300_000  # bytes


@cache(ttl="10m", key="whoami:{token}")
async def _resolve_user(token: str) -> str:
    async with httpx.AsyncClient(timeout=10) as c:
        r = await c.get("https://huggingface.co/api/whoami-v2",
                        headers={"Authorization": f"Bearer {token}"})
    if r.status_code != 200:
        raise HTTPException(status_code=401, detail="Invalid or expired HF token")
    name = r.json().get("name")
    if not name:
        raise HTTPException(status_code=401, detail="Could not resolve username")
    return name


async def _auth_user(authorization: str) -> str:
    if not authorization or not authorization.startswith("Bearer "):
        raise HTTPException(status_code=401, detail="Missing Bearer token")
    return await _resolve_user(authorization[7:])


@app.get("/feeds_store")
async def feeds_load(authorization: str = Header(default=None)):
    user = await _auth_user(authorization)
    import json as _json
    try:
        path = hf_hub_download(FEEDS_STORE_REPO, f"users/{user}.json",
                               repo_type="dataset", force_download=True)
        with open(path) as f:
            return {"user": user, "doc": _json.load(f)}
    except Exception:
        return {"user": user, "doc": None}


class FeedsSaveReq(BaseModel):
    doc: dict


@app.put("/feeds_store")
async def feeds_save(req: FeedsSaveReq, authorization: str = Header(default=None)):
    user = await _auth_user(authorization)
    import json as _json
    payload = _json.dumps(req.doc).encode()
    if len(payload) > FEEDS_DOC_LIMIT:
        raise HTTPException(status_code=413, detail="Feeds doc too large")
    HfApi().upload_file(path_or_fileobj=payload, path_in_repo=f"users/{user}.json",
                        repo_id=FEEDS_STORE_REPO, repo_type="dataset",
                        commit_message=f"feeds sync: {user}")
    return {"ok": True, "user": user, "bytes": len(payload)}


# --- suggest (autocomplete; old backend never implemented it — bonus) --------
@app.get("/suggest/{repo_type}")
@cache(ttl="1h", key="suggest:{repo_type}:{q}")
async def suggest(repo_type: str, q: str):
    if len(q) < 2:
        return {"suggestions": []}
    table = "dataset_cards" if repo_type == "datasets" else "model_cards"
    rows = con.execute(
        f"SELECT id FROM {table} WHERE id ILIKE ? ORDER BY downloads DESC LIMIT 10",
        [f"%{q}%"],
    ).fetchall()
    return {"suggestions": [r[0] for r in rows]}


if __name__ == "__main__":
    import uvicorn
    uvicorn.run(app, host="0.0.0.0", port=7860)