diff --git a/.gitignore b/.gitignore index ab04949..65107e1 100644 --- a/.gitignore +++ b/.gitignore @@ -3,6 +3,7 @@ *.jl.mem /Manifest*.toml /docs/Manifest*.toml +/test/Manifest*.toml /docs/build/ .env .env.example diff --git a/Project.toml b/Project.toml index f38ea99..f7a9cac 100644 --- a/Project.toml +++ b/Project.toml @@ -9,11 +9,10 @@ HuggingFaceHub = "d0076355-e2c0-48e6-a044-05906e51b7fc" JSON3 = "0f8b85d8-7281-11e9-16c2-39a750bddbf1" LibPQ = "194296ae-ab2e-5f79-8cd4-7183a0a5a0d1" LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" +OpenAI = "e9f21f70-7185-4079-aca2-91159181367c" PromptingTools = "670122d1-24a8-4d70-bfce-740807c42192" RAGTools = "16ddad29-bbe8-45a7-857d-3d9514eb0023" Serialization = "9e88b42a-f829-5b0c-bbe9-9e923198166b" -SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf" -Statistics = "10745b16-79ce-11e8-11f9-7d13ad32a3b2" URIs = "5c2747f8-b7ea-4ff2-ba2e-563bfd36b1d4" [compat] @@ -22,10 +21,9 @@ HuggingFaceHub = "0.1.2" JSON3 = "1.14.3" LibPQ = "1.18.0" LinearAlgebra = "1.10" +OpenAI = "0.11" PromptingTools = "0.82.1" RAGTools = "0.7.0" Serialization = "1.10" -SparseArrays = "1.10" -Statistics = "1.10" URIs = "1.6" julia = "1.10" diff --git a/docs/make.jl b/docs/make.jl index 634c9a8..ea8c3df 100644 --- a/docs/make.jl +++ b/docs/make.jl @@ -8,8 +8,8 @@ makedocs(; authors="ParamThakkar123 and TheCedarPrince ", sitename="HealthLLM.jl", format=Documenter.HTML(; - canonical="https://ParamThakkar123.github.io/HealthLLM.jl", - edit_link="master", + canonical="https://JuliaHealth.github.io/HealthLLM.jl", + edit_link="main", assets=String[], ), pages=[ @@ -17,10 +17,11 @@ makedocs(; "Getting Started" => "getting-started.md", "Document Ingestion" => "ingestion.md", "Building Embeddings" => "embeddings.md", + "Querying the RAG System" => "querying.md", ], ) deploydocs(; - repo="github.com/ParamThakkar123/HealthLLM.jl", - devbranch="master", + repo="github.com/JuliaHealth/HealthLLM.jl", + devbranch="main", ) diff --git a/docs/src/embeddings.md b/docs/src/embeddings.md index 9e92828..8e72f21 100644 --- a/docs/src/embeddings.md +++ b/docs/src/embeddings.md @@ -58,9 +58,73 @@ embedding_ref("all-minilm"; provider=:huggingface) # "hf:sentence-transf !!! note "Backend setup" For Ollama, pull the tag once with `ollama pull nomic-embed-text` and make - sure the server is running. For HuggingFace, ensure the corresponding schema - is configured in `PromptingTools`. Either way, register the model in the RAG - pipeline with `register_models(...)` as shown in [Getting Started](getting-started.md). + sure the server is running. For HuggingFace, set a token (see below). Either + way, register the model in the RAG pipeline with `register_models(...)` as + shown in [Getting Started](getting-started.md). + +### The HuggingFace backend + +PromptingTools ships schemas for a long list of OpenAI-compatible providers but +none for HuggingFace, so this package supplies one: +[`HuggingFaceOpenAISchema`](@ref). It is an `AbstractOpenAISchema`, so +`aigenerate`, `aiembed`, message rendering and retries all work unchanged — +`provider = :huggingface` and `hf:`-prefixed model names route through it +automatically. + +Set a token once per session, or let it come from the environment +(`HF_API_TOKEN`, `HF_TOKEN`, `HUGGINGFACE_API_KEY`, `HUGGING_FACE_HUB_TOKEN`, +checked in that order): + +```julia +set_huggingface_api_key!(ENV["HF_TOKEN"]) +huggingface_api_key() # what will actually be sent +``` + +!!! note "Pinning an inference provider" + The router auto-routes only to providers **enabled on your account**, so a + model can be live on HuggingFace and still be refused with + `model_not_supported`. Append the provider to pin it: + + ```julia + huggingface_providers("Qwen/Qwen2.5-7B-Instruct") # ["featherless-ai"] + + register_models("hf:Qwen/Qwen2.5-7B-Instruct:featherless-ai", "hf:BAAI/bge-m3") + ``` + + When routing is refused, the error raised here already names the live + providers and the exact model string to use, so you rarely need to look this + up yourself. Alternatively enable the provider at + [huggingface.co/settings/inference-providers](https://huggingface.co/settings/inference-providers). + +Chat and embeddings use different HuggingFace surfaces, because the router's +OpenAI-compatible API covers chat completions only — `/v1/embeddings` answers +404 there: + +| Call | Endpoint | +|--------------|---------------------------------------------------------------------------| +| `aigenerate` | [`HUGGINGFACE_ROUTER_URL`](@ref) + `/chat/completions` | +| `aiembed` | [`HUGGINGFACE_INFERENCE_URL`](@ref) + `//pipeline/feature-extraction` | + +The feature-extraction reply is reshaped into the OpenAI embeddings response +`aiembed` expects, and token-level output (from models that do not pool +internally) is mean-pooled to one vector per input — so a HuggingFace embedding +matrix is the same `dim × n` shape as an Ollama one. + +!!! note "Cold models" + A HuggingFace model that is not already warm loads while holding the + connection open — around 40–55s for `bge-m3` in practice. That exceeds + PromptingTools' 120s `aiembed` default under load, so this package raises the + read timeout to [`HUGGINGFACE_EMBED_TIMEOUT`](@ref) (300s) when you have not + chosen one yourself. Any `http_kwargs` you pass is left exactly as given. + +To use a deployment that *does* speak OpenAI embeddings — a Text Embeddings +Inference container or a dedicated Inference Endpoint — pass its URL and the +request is forwarded there instead: + +```julia +E = embed(chunks, "bge-m3"; provider = :huggingface, + api_kwargs = (; url = "https://my-endpoint.hf.space/v1")) +``` ## Generating embeddings @@ -186,7 +250,7 @@ conn = LibPQ.Connection("postgresql://user:pass@localhost/health") store = PgVectorStore(conn, embedding_dimension(); table = "omop_embeddings", metric = :cosine) add!(store, E, chunks) # creates the table + inserts in one transaction -hits = search(store, embed("count patients"), 5) # (; id, chunk, distance) +hits = search(store, embed("count patients"), 5) # Vector{Hit}: row id in `index`, raw `distance` kept ``` `metric` chooses the distance operator: `:cosine` (`<=>`), `:dot` (`<#>`), or @@ -196,7 +260,7 @@ wrapper: ```julia store_embeddings_pgvector(conn, E, chunks, embedding_dimension(); table = "omop_embeddings") -hits = search_embeddings_pgvector(conn, embed("count patients"), 5; table = "omop_embeddings") +hits = search_embeddings_pgvector(conn, embed("count patients"), 5; table = "omop_embeddings") # raw (; id, chunk, distance) ``` !!! note "pgvector prerequisites" @@ -217,7 +281,7 @@ using Faiss # optional; load before con store = FaissVectorStore(embedding_dimension()) # inner-product index (cosine after normalisation) add!(store, E, chunks) -hits = search(store, embed("count patients"), 5) # (; index, chunk, score) +hits = search(store, embed("count patients"), 5) # Vector{Hit}, as with every backend ``` As with the local store, vectors are normalised for the inner-product/cosine diff --git a/docs/src/getting-started.md b/docs/src/getting-started.md index 8472512..152f069 100644 --- a/docs/src/getting-started.md +++ b/docs/src/getting-started.md @@ -42,7 +42,10 @@ corpus = write_combined_file(files, "corpus.txt") ### 2. Register models -Register chat and embedding models for use with `PromptingTools`. Supports Ollama, HuggingFace, and other backends: +Register chat and embedding models for use with `PromptingTools`. Model names +prefixed with `hf:` are routed to HuggingFace via [`HuggingFaceOpenAISchema`](@ref) +(set a token with [`set_huggingface_api_key!`](@ref) or `HF_TOKEN`); everything +else defaults to Ollama. See [Building Embeddings](embeddings.md) for the details: ```julia register_models("llama3.2", "nomic-embed-text") diff --git a/docs/src/index.md b/docs/src/index.md index 28b4b04..ef57094 100644 --- a/docs/src/index.md +++ b/docs/src/index.md @@ -13,6 +13,7 @@ The package centers on five areas: - collecting source files and writing combined corpora - ingesting curated docs and web-search results into an index (see [Document Ingestion](ingestion.md)) - building retrieval indexes through `RAGTools` +- constructing grounded prompts from retrieved chunks (see [Querying the RAG System](querying.md)) - generating retrieval-backed answers for query construction - building and validating embeddings across Ollama and HuggingFace (see [Building Embeddings](embeddings.md)) - storing embeddings in a local file, PostgreSQL/`pgvector`, or FAISS @@ -32,5 +33,5 @@ index = build_index_rag(RAGTools.SimpleIndexer(), files) More detailed setup, testing commands, and the end-to-end walkthrough are in [Getting Started](getting-started.md). ```@autodocs -Modules = [HealthLLM, HealthLLM.Utils, HealthLLM.Database, HealthLLM.Query, HealthLLM.Ingestion, HealthLLM.Embeddings, HealthLLM.Storage] +Modules = [HealthLLM, HealthLLM.HuggingFace, HealthLLM.Utils, HealthLLM.Database, HealthLLM.Prompt, HealthLLM.Execution, HealthLLM.Query, HealthLLM.Ingestion, HealthLLM.Embeddings, HealthLLM.Storage] ``` diff --git a/docs/src/querying.md b/docs/src/querying.md new file mode 100644 index 0000000..92b448b --- /dev/null +++ b/docs/src/querying.md @@ -0,0 +1,327 @@ +```@meta +CurrentModule = HealthLLM +``` + +# Querying the RAG System + +Once documents are [ingested](ingestion.md) and [embedded](embeddings.md) into a +vector store, querying is a three-step loop: **retrieve** the chunks relevant to +a question, **construct** a grounded prompt from them, and **generate** a FunSQL +answer with a chat model. + +``` +question ─▶ retrieve ─▶ ranked chunks ─┐ + ├─▶ build_prompt ─▶ (; system, user, prompt) ─▶ chat model ─▶ FunSQL + (schema + docs + examples) ─┘ +``` + +The middle step — [`build_prompt`](@ref) — is what keeps generations honest: it +injects the retrieved schema/doc chunks alongside the natural-language question +and pairs them with a system message that forbids the model from using any table +or column name it cannot see in the context. + +## From question to prompt + +Assume you already have a populated store (see [Building Embeddings](embeddings.md)): + +```julia +using HealthLLM + +store = load(LocalVectorStore, "omop_index.jls") +question = "How many distinct patients had a diabetes diagnosis in 2020?" +``` + +Retrieve the most relevant chunks, then build the prompt from them: + +```julia +hits = retrieve(store, question, 6) # top-6 nearest chunks +p = build_prompt(question, hits) # (; system, user, prompt) +``` + +[`build_prompt`](@ref) returns a NamedTuple with three fields so it fits either +prompting style: + +| Field | What it is | Use it for | +|----------|-------------------------------------------------------|---------------------------------------------| +| `system` | The FunSQL/OMOP system message ([`FUNSQL_SYSTEM_PROMPT`](@ref)) | a proper system/user chat conversation | +| `user` | Context block + question + answer cue | the user turn of that conversation | +| `prompt` | `system` and `user` joined into one string | single-string / `{input_query}`-style calls | + +The `user` turn interleaves the retrieved context with the question: + +``` +# Retrieved context + +[Source 1] https://ohdsi.github.io/CommonDataModel/ › condition_occurrence +# condition_occurrence +condition_occurrence_id, person_id, condition_concept_id, condition_start_date, ... + +[Source 2] FunSQL-examples › patients with a condition +Query: patients with a condition +FunSQL: From(condition_occurrence) |> Group(Get.person_id) + +# Analytical question + +How many distinct patients had a diabetes diagnosis in 2020? + +# FunSQL query +``` + +## Generating the answer + +Feed the prompt to a registered chat model through `PromptingTools`. The split +`system`/`user` form is preferred — it puts the grounding rules in the system +role where models weight them most heavily: + +```julia +using PromptingTools + +register_models("llama3.2", "nomic-embed-text") # once per session + +msg = aigenerate( + [PromptingTools.SystemMessage(p.system), + PromptingTools.UserMessage(p.user)]; + model = "llama3.2", +) +println(msg.content) # the FunSQL query + a short explanation +``` + +For a store-free index built with RAGTools, [`generate_funsql_query`](@ref) +retrieves and generates in one call against a RAG `index` (see +[Document Ingestion](ingestion.md)): + +```julia +index = ingest_to_index(; sources = ["FunSQL.jl", "OMOP CDM"], + query = "OMOP condition_occurrence columns") +answer = generate_funsql_query(index, "nomic-embed-text", "llama3.2", + "Context: {input_query}", question) +``` + +## Grounding: why the model stays on-schema + +The single biggest failure mode for text-to-query over a large schema like OMOP +is **hallucinated identifiers** — the model confidently writes `patient_id` or +`diagnosis_date` because those names are plausible, even though the CDM uses +`person_id` and `condition_start_date`. [`FUNSQL_SYSTEM_PROMPT`](@ref) counters +this with strict grounding rules that make the retrieved context the *only* +admissible source of names and syntax: + +- **Names must appear verbatim in the context.** Every table and column the model + references has to be present in a retrieved chunk. Memory of the OMOP CDM is + explicitly overridden, and inventing, guessing, or "correcting" an identifier is + forbidden. +- **FunSQL constructs are copied, not recalled.** The model may use `From`, + `Where`, `Join`, `Group`, `Select`, `Get`, `Agg`, `Fun`, and `|>` only as the + retrieved examples demonstrate them — no raw SQL strings, no undemonstrated + features. +- **Missing schema is surfaced, not faked.** When the context lacks a table or + column the question needs, the model is told to say what is missing and generate + the closest correct query from what *is* available — never to paper over the gap + with a fabricated name. + +Because grounding is only as good as the context, retrieval quality matters: if +`retrieve` does not surface the `condition_occurrence` schema chunk, no prompt +rule can make the model use it correctly. Two levers help: + +1. **Retrieve enough chunks.** A larger `k` (or a larger `max_chunks`) raises the + chance the needed table definition is present — at the cost of a longer prompt. +2. **Index schema at table granularity.** The [`HeaderChunk`](@ref) strategy keeps + each OMOP table's full column list in one chunk, so a single hit carries the + whole definition rather than a fragment. + +## Tuning the prompt + +[`build_prompt`](@ref) is driven by a [`PromptTemplate`](@ref) — pass one to +change the system message, the section headers, per-chunk formatting, or the +context budget. [`DEFAULT_FUNSQL_TEMPLATE`](@ref) is the OMOP/FunSQL default +(provenance on, scores off, up to 8 chunks or 8000 characters). + +```julia +tmpl = PromptTemplate( + include_scores = true, # show each chunk's retrieval score + max_chunks = 5, # keep at most 5 chunks + max_context_chars = 6000, # ...and cap the context at 6000 chars +) + +p = build_prompt(question, hits; template = tmpl) +``` + +The two budgets protect the context window: chunks are admitted highest-ranked +first, and admission stops at a **chunk boundary** once either limit is reached — +a chunk is included whole or skipped, never cut mid-text (the sole exception is a +single chunk that exceeds the character budget on its own, which is truncated with +an ellipsis). This means the prompt never silently overflows the model's window, +and lower-ranked chunks are the ones dropped when space runs out. + +You can render just the context block — for inspection or a custom prompt layout — +with [`format_context`](@ref): + +```julia +println(format_context(DEFAULT_FUNSQL_TEMPLATE, hits)) +``` + +### Flexible chunk inputs + +`build_prompt` and `format_context` accept whatever your retrieval step produces, +so the same call works across store backends and index paths: + +| Input | Where it comes from | +|--------------------------------|----------------------------------------------| +| `search`/`retrieve` hit NamedTuples | [`LocalVectorStore`](@ref) / [`FaissVectorStore`](@ref) (`.chunk`, `.score`) and [`PgVectorStore`](@ref) (`.chunk`, `.distance`) | +| plain `String`s | any ad-hoc list of context snippets | +| [`Chunk`](@ref)s | the ingestion layer, carrying `:heading`/`:source`/`:url` provenance | + +When a chunk carries provenance (a URL or source plus its parent heading), it is +shown in the `[Source i]` header so the model — and you — can trace each fact back +to its origin. + +## From RAG response to a database query + +The model's reply is prose with a FunSQL code block embedded in it — not something +you can run directly. Turning that reply into an executed query is three moves: +**generate**, **extract**, **render to SQL**, then hand the SQL to a database. + +``` +grounded prompt ─▶ generate_funsql ─▶ .funsql (code) ─▶ eval ─▶ SQLNode ─▶ FunSQL.render ─▶ SQL ─▶ DB +``` + +[`generate_funsql`](@ref) sends the prompt to the chat model and pulls the FunSQL +code out of the reply for you (via [`extract_funsql`](@ref), which prefers the +first ` ```julia ` block): + +```julia +gen = generate_funsql(p; model = "llama3.2") +gen.funsql # just the FunSQL code, e.g. "From(:person) |> Group() |> Select(...)" +gen.answer # the full reply, including the model's explanation +``` + +FunSQL is a **query builder**: the code evaluates to a `SQLNode`, which +`FunSQL.render` serialises to SQL for a specific dialect. You then execute that SQL +with whatever driver your data lives behind — [DuckDB](https://duckdb.org) for the +JuliaHealth demo datasets, or `LibPQ` for PostgreSQL. FunSQL and the driver are +**driver-side dependencies** (not pulled in by HealthLLM), so load them yourself: + +```julia +using FunSQL, DuckDB, DataFrames + +conn = DuckDB.DB("synthea.duckdb") + +# 1. Evaluate the generated code into a FunSQL query node. +node = eval(Meta.parse(gen.funsql)) + +# 2. Render it to DuckDB SQL. +sql = FunSQL.render(node; dialect = :duckdb) + +# 3. Execute against the database. +result = DuckDB.execute(conn, String(sql)) |> DataFrame +``` + +!!! warning "You are evaluating model output" + `eval`-ing generated code runs it with full privileges. Only do this with a + model you trust, sandbox the database it can reach, and never point the + execution path at production data. This is why the [sanity + check](#Sanity-checking-the-generated-query) below renders (and optionally + runs) against a *disposable* schema/connection first. + +Rendering against a **schema catalog** rather than a bare dialect is what makes the +query trustworthy — FunSQL resolves every table and column against the catalog and +errors on anything it does not find. Build a catalog by hand, or reflect one from a +live connection: + +```julia +catalog = FunSQL.reflect(conn; dialect = :duckdb) # discovers the DB's tables +sql = FunSQL.render(catalog, node) # fails on unknown table/column +``` + +## Sanity-checking the generated query + +Before running generated SQL — or surfacing it to a user — check that it is +actually valid. [`sanity_check_funsql`](@ref) runs the query as far as it can +through a **parse → build → render → run** pipeline and reports the outcome as a +[`FunSQLCheck`](@ref), without ever throwing on a bad query: + +| Stage | What it proves | Needs | +|------------|-----------------------------------------------------|---------------------------| +| `parsed` | the code is syntactically valid Julia | nothing | +| `built` | it evaluates to a `FunSQL.SQLNode` | `using FunSQL` | +| `rendered` | it serialises to SQL — **resolved against a schema**| a `catalog` (recommended) | +| `executed` | the database accepts the SQL | an `executor` | + +```julia +using FunSQL + +catalog = FunSQL.SQLCatalog( + FunSQL.SQLTable(:person, columns = [:person_id, :year_of_birth]), + FunSQL.SQLTable(:condition_occurrence, + columns = [:condition_occurrence_id, :person_id, :condition_concept_id]), + dialect = :duckdb, +) + +check = sanity_check_funsql(gen; catalog = catalog) +check.ok # true if it parsed, built, and rendered +check.sql # the rendered SQL, when rendering succeeded +``` + +The `catalog` stage is the **execution-based grounding check**: it catches a +hallucinated OMOP name that slipped past the prompt's grounding rules, because +FunSQL refuses to resolve a column the schema does not contain. + +```julia +# The model invented a column that isn't in the schema: +bad = sanity_check_funsql("From(:person) |> Select(Get.made_up_column)"; catalog = catalog) +bad.ok # false +bad.stage # :built — it built, but failed to resolve against the schema +bad.error # message naming the unresolved reference +``` + +Pass an `executor` to go one stage further and confirm the database itself accepts +the SQL — the check stays failure-safe, capturing any driver error instead of +throwing: + +```julia +using DuckDB +conn = DuckDB.DB("synthea.duckdb") + +check = sanity_check_funsql(gen; catalog = catalog, + executor = sql -> DuckDB.execute(conn, sql)) +check.executed # true if the query ran +``` + +A natural loop is: generate → sanity-check → if `!check.ok`, feed `check.error` +back to the model (alongside the same context) and regenerate, so a hallucinated +name becomes a corrective signal rather than a silent failure. + +## End to end + +```julia +using HealthLLM, PromptingTools, FunSQL, DuckDB, DataFrames + +register_models("llama3.2", "nomic-embed-text") + +# 1. Load (or build) the index and open the database. +store = load(LocalVectorStore, "omop_index.jls") +conn = DuckDB.DB("synthea.duckdb") +catalog = FunSQL.reflect(conn; dialect = :duckdb) # schema to validate against + +# 2. Retrieve context for the question and build the grounded prompt. +question = "What is the average age of patients with a hypertension diagnosis?" +hits = retrieve(store, question, 6) +p = build_prompt(question, hits) + +# 3. Generate the FunSQL query from the grounded prompt. +gen = generate_funsql(p; model = "llama3.2") + +# 4. Sanity-check it against the schema before trusting it. +check = sanity_check_funsql(gen; catalog = catalog) +if !check.ok + error("Generated query failed at :$(check.stage) — $(check.error)") +end + +# 5. Execute the validated SQL. +result = DuckDB.execute(conn, check.sql) |> DataFrame +``` + +See the [API reference](index.md) for full docstrings of [`build_prompt`](@ref), +[`generate_funsql`](@ref), [`sanity_check_funsql`](@ref), [`FunSQLCheck`](@ref), +and [`FUNSQL_SYSTEM_PROMPT`](@ref). +``` diff --git a/src/HealthLLM.jl b/src/HealthLLM.jl index 10e5fa0..50e2a64 100644 --- a/src/HealthLLM.jl +++ b/src/HealthLLM.jl @@ -2,48 +2,41 @@ module HealthLLM using PromptingTools using RAGTools -using LinearAlgebra -using SparseArrays -using JSON3, Serialization -using Statistics +# Included in dependency order: each module may use the ones above it. +include("huggingface.jl") include("utils.jl") include("database.jl") include("embeddings.jl") include("storage.jl") +include("prompt.jl") +include("execution.jl") include("query.jl") include("ingestion.jl") -import .Utils: collect_files_with_extensions, write_combined_file, register_models, load_huggingface_model, HuggingFaceLoadResult, build_index_rag -import .Database: store_embeddings_pgvector, search_embeddings_pgvector, validate_embeddings_inputs -import .Embeddings: EmbeddingModel, EMBEDDING_MODELS, DEFAULT_EMBEDDING_MODEL, - embedding_model, embedding_ref, embedding_dimension, - embed, cosine_similarity, similarity_matrix, - validate_embeddings, embedding_sanity_check -import .Storage: AbstractVectorStore, LocalVectorStore, PgVectorStore, FaissVectorStore, - add!, search, retrieve, save, load -import .Query: generate_funsql_query -import .Ingestion: SourceDocument, SearchResult, - AbstractSearchProvider, DuckDuckGoProvider, - default_search_provider, web_search, - CURATED_SOURCES, fetch_url, html_to_text, fetch_curated, - ingest, ingest_to_index +""" + HealthLLM.PUBLIC_MODULES +The submodules whose exported names make up the public API of `HealthLLM`. +Everything a submodule exports is re-exported from the package, so each symbol is +declared in exactly one place — the `export` list of the module that defines it. +Adding a name to that list is all it takes to publish it; there is no second +mirrored list here to fall out of sync. +""" +const PUBLIC_MODULES = (HuggingFace, Utils, Database, Embeddings, Storage, Prompt, + Execution, Query, Ingestion) + +# Explicit `import` rather than `using`, because an explicit import takes +# precedence over names a `using` brings in: `Storage.retrieve` has to win over +# the `retrieve` that `using RAGTools` above also provides, or the name resolves +# to neither and disappears from the package's API. +for m in PUBLIC_MODULES, n in names(m) + n === nameof(m) && continue + Core.eval(@__MODULE__, Meta.parse("import .$(nameof(m)): $n")) + Core.eval(@__MODULE__, Expr(:export, n)) +end + +# The two upstream packages callers routinely need alongside HealthLLM. export PromptingTools, RAGTools -export collect_files_with_extensions, write_combined_file, generate_funsql_query, - build_index_rag, store_embeddings_pgvector, search_embeddings_pgvector, - validate_embeddings_inputs, - register_models, load_huggingface_model, HuggingFaceLoadResult -export EmbeddingModel, EMBEDDING_MODELS, DEFAULT_EMBEDDING_MODEL, - embedding_model, embedding_ref, embedding_dimension, - embed, cosine_similarity, similarity_matrix, - validate_embeddings, embedding_sanity_check -export AbstractVectorStore, LocalVectorStore, PgVectorStore, FaissVectorStore, - add!, search, retrieve, save, load -export SourceDocument, SearchResult, - AbstractSearchProvider, DuckDuckGoProvider, - default_search_provider, web_search, - CURATED_SOURCES, fetch_url, html_to_text, fetch_curated, - ingest, ingest_to_index end diff --git a/src/database.jl b/src/database.jl index eb3edb9..55a4ec3 100644 --- a/src/database.jl +++ b/src/database.jl @@ -1,5 +1,9 @@ module Database + using LibPQ +using ..Utils: check_dims, check_k + +export store_embeddings_pgvector, search_embeddings_pgvector, validate_embeddings_inputs function _vector_to_pgarray(v::AbstractVector{T}) where T<:Real string("[", join(v, ","), "]") @@ -30,12 +34,15 @@ _pg_create_sql(table::AbstractString, dim::Integer) = """ ) """ -_pg_search_sql(table::AbstractString, metric::Symbol) = """ - SELECT id, chunk, embedding $(_pg_metric_op(metric)) \$1 AS distance - FROM $(_check_identifier(table)) - ORDER BY embedding $(_pg_metric_op(metric)) \$1 - LIMIT \$2 -""" +function _pg_search_sql(table::AbstractString, metric::Symbol) + op = _pg_metric_op(metric) + return """ + SELECT id, chunk, embedding $op \$1 AS distance + FROM $(_check_identifier(table)) + ORDER BY embedding $op \$1 + LIMIT \$2 + """ +end """ validate_embeddings_inputs(embeddings, chunks, embedding_dimension) @@ -66,17 +73,7 @@ function validate_embeddings_inputs( chunks::AbstractVector, embedding_dimension::Integer ) - n_rows, n_cols = size(embeddings) - n_rows == embedding_dimension || throw( - DimensionMismatch( - "Embedding height ($n_rows) must match embedding_dimension ($embedding_dimension)." - ) - ) - length(chunks) == n_cols || throw( - DimensionMismatch( - "Number of chunks ($(length(chunks))) must match number of embedding columns ($n_cols)." - ) - ) + check_dims(embeddings, chunks, embedding_dimension) return nothing end @@ -111,13 +108,12 @@ function store_embeddings_pgvector( table::AbstractString="embeddings" ) validate_embeddings_inputs(embeddings, chunks, embedding_dimension) - _check_identifier(table) LibPQ.execute(conn, _pg_create_sql(table, embedding_dimension)) dense_embeddings = Matrix{Float64}(embeddings) chunk_text = String.(chunks) - insert_sql = "INSERT INTO $table (chunk, embedding) VALUES (\$1, \$2)" + insert_sql = "INSERT INTO $(_check_identifier(table)) (chunk, embedding) VALUES (\$1, \$2)" LibPQ.execute(conn, "BEGIN") try @@ -159,8 +155,7 @@ function search_embeddings_pgvector( table::AbstractString="embeddings", metric::Symbol=:cosine ) - k > 0 || throw(ArgumentError("k must be positive, got $k")) - _check_identifier(table) + check_k(k) qvec = _vector_to_pgarray(Vector{Float64}(query)) res = LibPQ.execute(conn, _pg_search_sql(table, metric), (qvec, k)) return [(; id=row.id, chunk=row.chunk, distance=row.distance) for row in res] diff --git a/src/embeddings.jl b/src/embeddings.jl index a25fee8..645a9fd 100644 --- a/src/embeddings.jl +++ b/src/embeddings.jl @@ -23,6 +23,7 @@ module Embeddings using PromptingTools using LinearAlgebra +using ..Utils: get_schema, check_dims export EmbeddingModel, EMBEDDING_MODELS, DEFAULT_EMBEDDING_MODEL, embedding_model, embedding_ref, embedding_dimension, @@ -119,11 +120,6 @@ Return the output dimension of the named embedding model. """ embedding_dimension(name::AbstractString=DEFAULT_EMBEDDING_MODEL) = embedding_model(name).dim -_schema(provider::Symbol) = - provider === :ollama ? PromptingTools.OllamaSchema() : - provider === :huggingface ? PromptingTools.HuggingFaceSchema() : - throw(ArgumentError("provider must be :ollama or :huggingface, got :$provider")) - # Normalise an aiembed result into a dim × n Float32 matrix (columns = chunks). function _as_matrix(content) m = content isa AbstractVector ? reshape(content, :, 1) : content @@ -152,7 +148,7 @@ function embed(texts::AbstractVector{<:AbstractString}, name::AbstractString=DEFAULT_EMBEDDING_MODEL; provider::Symbol=:ollama, kwargs...) isempty(texts) && throw(ArgumentError("`texts` is empty; nothing to embed.")) ref = embedding_ref(name; provider=provider) - res = PromptingTools.aiembed(_schema(provider), texts; model=ref, kwargs...) + res = PromptingTools.aiembed(get_schema(provider), texts; model=ref, kwargs...) return _as_matrix(res.content) end @@ -213,17 +209,10 @@ function validate_embeddings(embeddings::AbstractMatrix; expected_dim::Union{Nothing,Integer}=nothing, chunks::Union{Nothing,AbstractVector}=nothing, atol::Real=1e-4) - dim, n = size(embeddings) - n == 0 && throw(ArgumentError("No embeddings to validate (matrix has zero columns).")) + size(embeddings, 2) == 0 && + throw(ArgumentError("No embeddings to validate (matrix has zero columns).")) + dim, n = check_dims(embeddings, chunks, expected_dim) - if expected_dim !== nothing && dim != expected_dim - throw(DimensionMismatch( - "Embedding dimension ($dim) does not match expected_dim ($expected_dim).")) - end - if chunks !== nothing && length(chunks) != n - throw(DimensionMismatch( - "Number of chunks ($(length(chunks))) does not match embedding columns ($n).")) - end all(isfinite, embeddings) || throw(ArgumentError("Embeddings contain non-finite values (NaN/Inf).")) diff --git a/src/execution.jl b/src/execution.jl new file mode 100644 index 0000000..d67f90c --- /dev/null +++ b/src/execution.jl @@ -0,0 +1,309 @@ +""" + Execution + +Turn a grounded prompt into a FunSQL query and check that the query the model +produced is actually usable. This closes the loop opened by the `Prompt` module: + +``` +build_prompt ─▶ generate_funsql ─▶ FunSQLGeneration ─▶ sanity_check_funsql ─▶ FunSQLCheck + (call the model) (extracted code) (parse / build / render / run) +``` + +Two responsibilities: + +1. **Generation** — [`generate_funsql`](@ref) sends the grounded prompt to a chat + model and extracts the FunSQL code block from the reply + ([`extract_funsql`](@ref)). +2. **Execution-based sanity check** — [`sanity_check_funsql`](@ref) confirms the + generated code *parses* as Julia, *builds* into a `FunSQL.SQLNode`, and + *renders* to SQL — optionally against a schema catalog, which is what catches a + query that references a table or column the OMOP CDM does not have. With a + connection it can also *run* the SQL. + +## Optional dependency + +Building and rendering FunSQL needs the `FunSQL` package loaded in `Main` (it is a +test/driver-side dependency, not a hard dependency of HealthLLM, mirroring how +`FaissVectorStore` treats `Faiss`). The parse stage works without it; the +later stages raise a clear error until `using FunSQL` has run. + +!!! warning "Evaluating model output" + [`sanity_check_funsql`](@ref) `eval`s the generated code to build the query + object — the same approach the FunSQL test suite uses. Evaluated code runs with + full privileges, so only check output from a model you trust, and never point + the `conn` execution path at a production database. +""" +module Execution + +import PromptingTools +using ..Prompt: build_prompt, PromptTemplate, DEFAULT_FUNSQL_TEMPLATE +using ..Utils: require_main_module + +export extract_funsql, generate_funsql, FunSQLGeneration, + sanity_check_funsql, FunSQLCheck + +# --------------------------------------------------------------------------- +# Generation +# --------------------------------------------------------------------------- + +""" + FunSQLGeneration + +Result of [`generate_funsql`](@ref). + +# Fields +- `funsql::String`: The FunSQL code extracted from the model's reply — the code + block only, ready to hand to [`sanity_check_funsql`](@ref). +- `answer::String`: The model's full reply (code block plus any explanation). +""" +struct FunSQLGeneration + funsql::String + answer::String +end + +const _JULIA_FENCE = r"```(?:julia|jl)[ \t]*\r?\n(.*?)```"s +const _ANY_FENCE = r"```[A-Za-z0-9]*[ \t]*\r?\n(.*?)```"s + +""" + extract_funsql(text::AbstractString) -> String + +Pull the FunSQL code out of a model reply. Prefers the first fenced block tagged +` ```julia ` / ` ```jl `, falls back to the first fenced block of any language, +and finally to the whole stripped `text` when no fence is present. The returned +code is stripped of surrounding whitespace. +""" +function extract_funsql(text::AbstractString) + m = match(_JULIA_FENCE, text) + m === nothing && (m = match(_ANY_FENCE, text)) + body = m === nothing ? text : m.captures[1] + return String(strip(body)) +end + +# Turn a build_prompt result (or a plain string) into what aigenerate accepts: +# a system/user conversation when the grounding split is available, else a string. +_to_conversation(p::AbstractString) = String(p) +function _to_conversation(p) + if p isa NamedTuple && haskey(p, :system) && haskey(p, :user) + return [PromptingTools.SystemMessage(p.system), PromptingTools.UserMessage(p.user)] + elseif p isa NamedTuple && haskey(p, :prompt) + return String(p.prompt) + end + return string(p) +end + +_message_content(msg) = hasproperty(msg, :content) ? String(msg.content) : string(msg) + +function _default_llm_generator(prompt; model, schema=nothing, kwargs...) + conv = _to_conversation(prompt) + msg = schema === nothing ? + PromptingTools.aigenerate(conv; model=model, kwargs...) : + PromptingTools.aigenerate(schema, conv; model=model, kwargs...) + return _message_content(msg) +end + +""" + generate_funsql(prompt; model=PromptingTools.MODEL_CHAT, schema=nothing, + generator=, kwargs...) -> FunSQLGeneration + +Generate a FunSQL query from a grounded `prompt` and extract the code block from +the reply. `prompt` is typically the NamedTuple returned by [`build_prompt`](@ref) +— its `system`/`user` split is sent as a proper chat conversation so the grounding +rules land in the system role — but a plain prompt string works too. + +`model` is the chat model name (defaults to the active `PromptingTools.MODEL_CHAT`). +`schema` is an optional `PromptingTools` schema; when given it is passed to +`aigenerate` as the first argument. Extra `kwargs` are forwarded to the generator. + +`generator` is injectable for testing: it is called as +`generator(prompt; model, schema, kwargs...)` and must return the reply text. The +default calls `PromptingTools.aigenerate`. + +# Example + +```julia +hits = retrieve(store, question, 6) +p = build_prompt(question, hits) +gen = generate_funsql(p; model = "llama3.2") +gen.funsql # the extracted FunSQL code +``` +""" +function generate_funsql(prompt; + model::AbstractString=PromptingTools.MODEL_CHAT, + schema=nothing, + generator=_default_llm_generator, + kwargs...) + text = generator(prompt; model=model, schema=schema, kwargs...) + return FunSQLGeneration(extract_funsql(text), String(text)) +end + +# --------------------------------------------------------------------------- +# Execution-based sanity check +# --------------------------------------------------------------------------- + +""" + FunSQLCheck + +Outcome of [`sanity_check_funsql`](@ref). Records how far the generated query got +through the parse → build → render → run pipeline and why it stopped. + +# Fields +- `ok::Bool`: Reached the deepest requested stage with no error. +- `stage::Symbol`: Deepest stage reached — `:parsed`, `:built`, `:rendered`, or + `:executed`. On failure this is the last stage that *succeeded*. +- `parsed::Bool`: The code parsed as Julia. +- `built::Bool`: Evaluating it produced a `FunSQL.SQLNode`. +- `rendered::Bool`: The node rendered to SQL (against `catalog`, when supplied). +- `executed::Bool`: The SQL ran via `executor` (only when one was given). +- `sql::Union{String,Nothing}`: The rendered SQL, when rendering succeeded. +- `error::Union{String,Nothing}`: Message from the failing stage, or `nothing`. +""" +struct FunSQLCheck + ok::Bool + stage::Symbol + parsed::Bool + built::Bool + rendered::Bool + executed::Bool + sql::Union{String,Nothing} + error::Union{String,Nothing} +end + +function Base.show(io::IO, c::FunSQLCheck) + status = c.ok ? "ok" : "failed" + print(io, "FunSQLCheck($status @ :$(c.stage)") + c.error === nothing || print(io, ", error=", repr(c.error)) + print(io, ")") +end + +_funsql_module() = require_main_module(:FunSQL, + "Building or rendering a query needs FunSQL.jl: install it and run " * + "`using FunSQL` before calling sanity_check_funsql with build/render stages.") + +# FunSQL does not export its node constructors (`From`, `Get`, `Agg`, ...), so +# generated code that uses the bare names does not resolve under a plain +# `using FunSQL`. We eval it in a dedicated sandbox module that binds those names +# (and pulls in the `@funsql`/`funsql_*` DSL), so callers need not import anything +# into `Main`. Built lazily on first use because FunSQL is a Main-only dependency. +const _FUNSQL_NODE_NAMES = ( + :From, :Select, :Where, :Join, :LeftJoin, :Group, :Order, :Limit, :Get, :Agg, + :Fun, :As, :Define, :Bind, :Partition, :Append, :With, :Iterate, :Lit, :Var, + :Sort, :Asc, :Desc, :Over, :Highlight, +) +const _EVAL_MOD = Ref{Module}() + +function _funsql_eval_module() + isassigned(_EVAL_MOD) && return _EVAL_MOD[] + FunSQL = _funsql_module() + m = Module(:FunSQLSandbox) + Core.eval(m, :(const FunSQL = $FunSQL)) + Core.eval(m, :(using FunSQL)) # @funsql + funsql_* exports + for name in _FUNSQL_NODE_NAMES # the capitalized node API + isdefined(FunSQL, name) && + Core.eval(m, :(const $name = $(getproperty(FunSQL, name)))) + end + _EVAL_MOD[] = m + return m +end + +""" + sanity_check_funsql(code::AbstractString; catalog=nothing, dialect=:duckdb, + mod=nothing, executor=nothing) -> FunSQLCheck + sanity_check_funsql(gen::FunSQLGeneration; kwargs...) -> FunSQLCheck + +Check that generated FunSQL `code` is usable, without trusting the model's word +for it. Runs as far through the pipeline as the inputs allow and reports the +result as a [`FunSQLCheck`](@ref): + +1. **parse** — `code` parses as Julia. Needs nothing beyond Julia itself. +2. **build** — evaluating it yields a `FunSQL.SQLNode`. Requires `FunSQL` loaded in + `Main` (see the module note). By default the code is evaluated in a built-in + sandbox module that already binds the FunSQL node constructors (`From`, `Get`, + `Agg`, ...), so bare-name generations resolve without importing anything into + `Main`. Pass `mod` to evaluate in your own module instead. +3. **render** — the node serialises to SQL. With `catalog` (a `FunSQL.SQLCatalog`) + the query is resolved **against that schema**, so a reference to a table or + column the schema does not contain fails here — this is the check that catches a + hallucinated OMOP name that slipped past the prompt. Without a catalog the + render is structural, using `dialect` (a dialect name like `:duckdb` or a + `FunSQL.SQLDialect`). +4. **run** — when `executor` is supplied, the rendered SQL string is passed to it + (e.g. `sql -> DuckDB.execute(conn, sql)`), confirming the database accepts it. + +The check never throws for a bad query: any stage failure is captured in the +returned `error` with `ok = false`. It only throws for a genuine misuse (build/ +render requested without FunSQL available). + +# Examples + +```julia +using FunSQL +catalog = FunSQL.SQLCatalog( + FunSQL.SQLTable(:person, columns = [:person_id, :year_of_birth]), + dialect = :duckdb, +) + +sanity_check_funsql("From(:person) |> Group() |> Select(:n => Agg.count())"; + catalog = catalog) +# FunSQLCheck(ok @ :rendered) + +sanity_check_funsql("From(:person) |> Select(Get.made_up_column)"; catalog = catalog) +# FunSQLCheck(failed @ :built) -> error names the unresolved column +``` +""" +function sanity_check_funsql(code::AbstractString; catalog=nothing, dialect=:duckdb, + mod::Union{Module,Nothing}=nothing, executor=nothing) + parsed = built = rendered = executed = false + sql = nothing + + # 1. parse — wrap in a block so multi-line generations parse as one unit. + expr = try + Meta.parse(string("begin\n", code, "\nend")) + catch err + return FunSQLCheck(false, :none, false, false, false, false, nothing, _errmsg(err)) + end + parsed = true + + # 2. build — eval to a FunSQL.SQLNode. + FunSQL = _funsql_module() + eval_mod = mod === nothing ? _funsql_eval_module() : mod + node = try + Core.eval(eval_mod, expr) + catch err + return FunSQLCheck(false, :parsed, parsed, false, false, false, nothing, _errmsg(err)) + end + node isa FunSQL.SQLNode || + return FunSQLCheck(false, :parsed, parsed, false, false, false, nothing, + "generated code evaluated to a $(typeof(node)), not a FunSQL.SQLNode") + built = true + + # 3. render — against the schema catalog when given, else structurally. + sql = try + rendered_sql = catalog === nothing ? + FunSQL.render(node; dialect=dialect) : + FunSQL.render(catalog, node) + string(rendered_sql) + catch err + return FunSQLCheck(false, :built, parsed, built, false, false, nothing, _errmsg(err)) + end + rendered = true + + # 4. run — only when an executor is supplied. + if executor !== nothing + try + executor(sql) + executed = true + catch err + return FunSQLCheck(false, :rendered, parsed, built, rendered, false, sql, _errmsg(err)) + end + end + + stage = executed ? :executed : :rendered + return FunSQLCheck(true, stage, parsed, built, rendered, executed, sql, nothing) +end + +sanity_check_funsql(gen::FunSQLGeneration; kwargs...) = + sanity_check_funsql(gen.funsql; kwargs...) + +_errmsg(err) = first(sprint(showerror, err), 500) + +end diff --git a/src/huggingface.jl b/src/huggingface.jl new file mode 100644 index 0000000..7e777e0 --- /dev/null +++ b/src/huggingface.jl @@ -0,0 +1,357 @@ +""" + HuggingFace + +HuggingFace backend for PromptingTools. + +PromptingTools ships schemas for a long list of OpenAI-compatible providers +(Groq, Together, Fireworks, DeepSeek, Mistral, ...) but — as of v0.94, the latest +release — **none for HuggingFace**. This module supplies the missing one, built +the same way PromptingTools builds its own provider schemas: a marker type under +`AbstractOpenAISchema` plus `create_chat`/`create_embeddings` methods that point +the shared OpenAI transport at HuggingFace's router. + +Because [`HuggingFaceOpenAISchema`](@ref) is an `AbstractOpenAISchema`, everything +PromptingTools already does for OpenAI — `aigenerate`, `aiembed`, message +rendering, streaming, retries — works against HuggingFace unchanged. + +## Credentials + +The token is read from [`huggingface_api_key`](@ref): an explicitly passed +`api_key` wins, then a key set with [`set_huggingface_api_key!`](@ref), then the +first of `HF_API_TOKEN`, `HF_TOKEN`, `HUGGINGFACE_API_KEY`, or +`HUGGING_FACE_HUB_TOKEN` found in the environment. + +## Example + +```julia +using HealthLLM, PromptingTools + +set_huggingface_api_key!(ENV["HF_TOKEN"]) + +msg = PromptingTools.aigenerate(HuggingFaceOpenAISchema(), "Say hi"; + model = "meta-llama/Llama-3.1-8B-Instruct") + +emb = PromptingTools.aiembed(HuggingFaceOpenAISchema(), ["some text"]; + model = "BAAI/bge-m3") + +# or point embeddings at a TEI / Inference Endpoint that speaks OpenAI +emb = PromptingTools.aiembed(HuggingFaceOpenAISchema(), ["some text"]; + model = "BAAI/bge-m3", + api_kwargs = (; url = "https://my-endpoint.hf.space/v1")) +``` + +## Endpoints + +Chat and embeddings use different HuggingFace surfaces, because the router's +OpenAI-compatible API covers chat only: + +| Call | Endpoint | +|-------------|--------------------------------------------------------------| +| `aigenerate`| [`HUGGINGFACE_ROUTER_URL`](@ref) `/chat/completions` | +| `aiembed` | [`HUGGINGFACE_INFERENCE_URL`](@ref) `//pipeline/feature-extraction` | +""" +module HuggingFace + +import PromptingTools +import OpenAI +using HTTP +using JSON3 + +export HuggingFaceOpenAISchema, huggingface_api_key, set_huggingface_api_key!, + huggingface_providers, HUGGINGFACE_ROUTER_URL, HUGGINGFACE_INFERENCE_URL, + HUGGINGFACE_EMBED_TIMEOUT + +""" + HuggingFaceOpenAISchema + +Schema for HuggingFace models served over an OpenAI-compatible API — the +HuggingFace Inference Providers router by default, or any Text Embeddings +Inference / Inference Endpoint deployment you point it at. + +A subtype of `PromptingTools.AbstractOpenAISchema`, so it plugs into `aigenerate` +and `aiembed` exactly like the built-in provider schemas. Override the endpoint +per call with `api_kwargs = (; url = "...")`. + +# Example + +```julia +PromptingTools.aigenerate(HuggingFaceOpenAISchema(), "Hello"; + model = "meta-llama/Llama-3.1-8B-Instruct") +``` +""" +struct HuggingFaceOpenAISchema <: PromptingTools.AbstractOpenAISchema end + +""" + HUGGINGFACE_ROUTER_URL + +Base URL of the HuggingFace Inference Providers router (`https://router.huggingface.co/v1`), +the OpenAI-compatible entry point used when no `url` is supplied. +""" +const HUGGINGFACE_ROUTER_URL = "https://router.huggingface.co/v1" + +# Env vars checked in order. HF_API_TOKEN is what this repository's heavy tests +# already use; HF_TOKEN is the current HuggingFace CLI default; the other two are +# older spellings still common in CI configs. +const _TOKEN_ENV_VARS = ("HF_API_TOKEN", "HF_TOKEN", "HUGGINGFACE_API_KEY", + "HUGGING_FACE_HUB_TOKEN") + +const _API_KEY = Ref{String}("") + +""" + set_huggingface_api_key!(key) -> String + +Set the HuggingFace token for the current session, taking precedence over the +environment. Pass `""` to clear it and fall back to the environment again. +""" +function set_huggingface_api_key!(key::AbstractString) + _API_KEY[] = String(key) + return _API_KEY[] +end + +""" + huggingface_api_key() -> String + +Return the HuggingFace token: the one set by [`set_huggingface_api_key!`](@ref) +if any, else the first of `HF_API_TOKEN`, `HF_TOKEN`, `HUGGINGFACE_API_KEY`, +`HUGGING_FACE_HUB_TOKEN` present in the environment, else `""`. + +An empty result is not an error here — some endpoints are open — but the request +will fail with a 401 if the model requires authentication. +""" +function huggingface_api_key() + isempty(_API_KEY[]) || return _API_KEY[] + for var in _TOKEN_ENV_VARS + key = get(ENV, var, "") + isempty(key) || return String(key) + end + return "" +end + +# PromptingTools defaults `api_key` to OPENAI_API_KEY, so a non-empty value here +# does not mean the caller meant it for HuggingFace. Honour anything that is +# genuinely caller-supplied, otherwise reach for the HuggingFace token. +function _resolve_api_key(api_key::AbstractString) + supplied = !isempty(api_key) && String(api_key) != String(PromptingTools.OPENAI_API_KEY) + supplied && return String(api_key) + hf = huggingface_api_key() + return isempty(hf) ? String(api_key) : hf +end + +""" + hf_model_id(model) -> String + +Strip the `"hf:"` prefix this package uses to mark HuggingFace references +(see `embedding_ref`), leaving the bare repo id the API expects. + +# Example + +```julia +hf_model_id("hf:BAAI/bge-m3") # "BAAI/bge-m3" +hf_model_id("BAAI/bge-m3") # "BAAI/bge-m3" +``` +""" +hf_model_id(model::AbstractString) = + startswith(lowercase(String(model)), "hf:") ? String(model)[4:end] : String(model) + +""" + split_provider(model) -> (repo, provider) + +Split a `"org/repo:provider"` reference into its parts, returning an empty +`provider` when the model is not pinned. HuggingFace repo ids never contain `:`, +so the last colon unambiguously marks a provider pin. + +# Example + +```julia +split_provider("Qwen/Qwen2.5-7B-Instruct:featherless-ai") +# ("Qwen/Qwen2.5-7B-Instruct", "featherless-ai") +``` +""" +function split_provider(model::AbstractString) + s = String(model) + i = findlast(==(':'), s) + i === nothing && return (s, "") + return (s[1:prevind(s, i)], s[nextind(s, i):end]) +end + +""" + huggingface_providers(model; api_key=huggingface_api_key()) -> Vector{String} + +Names of the inference providers currently serving `model`, i.e. those whose +mapping status is `"live"`. Returns an empty vector when nothing serves it. + +The router only auto-routes to providers **enabled on your account**, so a model +can be live somewhere and still be refused. Use this to see the options, then pin +one by appending it to the model name. + +# Example + +```julia +huggingface_providers("Qwen/Qwen2.5-7B-Instruct") # ["featherless-ai"] +# then: model = "hf:Qwen/Qwen2.5-7B-Instruct:featherless-ai" +``` +""" +function huggingface_providers(model::AbstractString; + api_key::AbstractString=huggingface_api_key()) + repo, _ = split_provider(hf_model_id(model)) + url = "https://huggingface.co/api/models/$repo?expand%5B%5D=inferenceProviderMapping" + headers = isempty(api_key) ? Pair{String,String}[] : + ["Authorization" => "Bearer $api_key"] + resp = HTTP.get(url, headers; status_exception=true, readtimeout=30) + mapping = get(JSON3.read(String(resp.body)), :inferenceProviderMapping, nothing) + mapping === nothing && return String[] + return String[String(name) for (name, info) in pairs(mapping) + if String(get(info, :status, "")) == "live"] +end + +# HuggingFace answers an unroutable model with a bare "not supported by any +# provider you have enabled", which does not say that the model *is* served, just +# not by a provider this account has switched on. Name the live providers and the +# pinning syntax so the failure is one edit from fixed. +_is_model_not_supported(err) = occursin("model_not_supported", sprint(showerror, err)) + +function _provider_hint(model::AbstractString, api_key::AbstractString) + repo, pinned = split_provider(model) + isempty(pinned) || return nothing # already pinned; the hint would be noise + providers = try + huggingface_providers(repo; api_key=api_key) + catch + return nothing + end + isempty(providers) && return nothing + served = length(providers) == 1 ? + "$(only(providers)), which is not enabled for your account" : + "$(join(providers, ", ")), none of which is enabled for your account" + return ErrorException( + "HuggingFace refused to route '$repo': it is served by $served. " * + "Pin a provider on the model name, e.g. model = \"hf:$repo:$(first(providers))\", " * + "or enable one at https://huggingface.co/settings/inference-providers.") +end + +function OpenAI.create_chat(schema::HuggingFaceOpenAISchema, + api_key::AbstractString, + model::AbstractString, + conversation; + url::String=HUGGINGFACE_ROUTER_URL, + kwargs...) + key = _resolve_api_key(api_key) + id = hf_model_id(model) + try + return OpenAI.create_chat(PromptingTools.CustomOpenAISchema(), + key, id, conversation; url, kwargs...) + catch err + if _is_model_not_supported(err) + hint = _provider_hint(id, key) + hint === nothing || throw(hint) + end + rethrow() + end +end + +""" + HUGGINGFACE_INFERENCE_URL + +Base URL for HuggingFace pipeline inference +(`https://router.huggingface.co/hf-inference/models`). Embeddings go here rather +than through [`HUGGINGFACE_ROUTER_URL`](@ref): the router's OpenAI-compatible +surface covers chat completions only and answers `/v1/embeddings` with a 404. +""" +const HUGGINGFACE_INFERENCE_URL = "https://router.huggingface.co/hf-inference/models" + +""" + HUGGINGFACE_EMBED_TIMEOUT + +Read timeout in seconds (`300`) used for feature-extraction when the caller did +not choose one. A HuggingFace model that is not already warm loads while holding +the connection open — measured at ~55s for `BAAI/bge-m3` — so PromptingTools' +120s `aiembed` default times out on the first call to a cold large model. +""" +const HUGGINGFACE_EMBED_TIMEOUT = 300 + +# `aiembed`'s exact default. Matching the whole tuple means we only substitute a +# longer timeout when the caller passed no `http_kwargs` at all; any deliberate +# choice, including a 120s one, is left untouched. +const _AIEMBED_DEFAULT_HTTP_KWARGS = (retry_non_idempotent=true, retries=5, readtimeout=120) + +_embed_http_kwargs(http_kwargs::NamedTuple) = + http_kwargs == _AIEMBED_DEFAULT_HTTP_KWARGS ? + merge(http_kwargs, (; readtimeout=HUGGINGFACE_EMBED_TIMEOUT)) : http_kwargs + +""" + OpenAI.create_embeddings(schema::HuggingFaceOpenAISchema, api_key, docs, model; url="", kwargs...) + +Embed `docs` with a HuggingFace model. + +With no `url`, this calls the `feature-extraction` pipeline under +[`HUGGINGFACE_INFERENCE_URL`](@ref) and reshapes the reply into the OpenAI +embeddings response `aiembed` expects, so the HuggingFace backend behaves like +any other from the caller's side. Token-level output (a matrix per input, from +models that do not pool internally) is mean-pooled into one vector per input. + +Pass `api_kwargs = (; url = "https://.../v1")` to target a deployment that *does* +serve an OpenAI-compatible `/embeddings` route — a Text Embeddings Inference +container or a dedicated Inference Endpoint — and the request is forwarded there +unchanged instead. +""" +function OpenAI.create_embeddings(schema::HuggingFaceOpenAISchema, + api_key::AbstractString, + docs, + model::AbstractString; + url::AbstractString="", + http_kwargs::NamedTuple=NamedTuple(), + kwargs...) + key = _resolve_api_key(api_key) + id = hf_model_id(model) + isempty(url) || return OpenAI.create_embeddings(PromptingTools.CustomOpenAISchema(), + key, docs, id; url=String(url), http_kwargs, kwargs...) + return _feature_extraction(key, docs, id; http_kwargs=http_kwargs) +end + +function _feature_extraction(api_key::AbstractString, docs, model::AbstractString; + http_kwargs::NamedTuple=NamedTuple()) + texts = docs isa AbstractString ? [String(docs)] : String[String(d) for d in docs] + isempty(texts) && throw(ArgumentError("`docs` is empty; nothing to embed.")) + + url = string(HUGGINGFACE_INFERENCE_URL, "/", model, "/pipeline/feature-extraction") + headers = ["Authorization" => "Bearer $api_key", "Content-Type" => "application/json"] + # wait_for_model keeps a cold model from failing the first call outright; the + # cost is that the load time is spent on this open connection, hence the + # larger default timeout in `_embed_http_kwargs`. + body = JSON3.write(Dict("inputs" => texts, + "options" => Dict("wait_for_model" => true))) + + resp = try + HTTP.post(url, headers, body; status_exception=true, _embed_http_kwargs(http_kwargs)...) + catch err + err isa HTTP.Exceptions.TimeoutError && throw(ErrorException( + "HuggingFace feature-extraction timed out for '$model'. A cold model can " * + "take minutes to load; retry, or raise the limit with " * + "`api_kwargs = (; http_kwargs = (; readtimeout = 600))`.")) + rethrow() + end + parsed = JSON3.read(String(resp.body)) + vectors = [_pool(item) for item in parsed] + + # Shaped to match OpenAI's embeddings response so `aiembed` needs no special + # case: it reads `.response[:data][i][:embedding]` and `.status`. + return (; + response=Dict( + :data => [Dict(:embedding => v) for v in vectors], + :usage => Dict(:prompt_tokens => 0)), + status=resp.status) +end + +# feature-extraction returns one vector per input for models that pool +# internally, and a token × hidden matrix for those that do not; average the +# token axis so callers always get a single vector per input. +function _pool(item) + isempty(item) && return Float64[] + first(item) isa Number && return Float64[Float64(x) for x in item] + rows = [_pool(row) for row in item] + width = length(first(rows)) + all(r -> length(r) == width, rows) || + throw(ArgumentError("ragged embedding rows returned by feature-extraction")) + return [sum(r[i] for r in rows) / length(rows) for i in 1:width] +end + +end diff --git a/src/ingestion/chunk.jl b/src/ingestion/chunk.jl index 9e3889f..6f9a020 100644 --- a/src/ingestion/chunk.jl +++ b/src/ingestion/chunk.jl @@ -223,16 +223,14 @@ end """ chunk_provenance(c::Chunk) -> String -Render a chunk's grounding as a single provenance string (`" > +Render a chunk's grounding as a single provenance string (`" › "`, truncated to 512 chars) suitable for RAGTools chunk sources. + +Delegates to [`render_provenance`](@ref), which is also what the prompt layer uses +to tag a retrieved chunk — so what is stored alongside a vector and what the model +is shown are the same string. """ -function chunk_provenance(c::Chunk) - base = get(c.metadata, :url, "") - isempty(base) && (base = get(c.metadata, :source, "")) - parent = get(c.metadata, :heading, get(c.metadata, :group, "")) - prov = isempty(parent) ? base : string(base, " > ", parent) - return first(prov, 512) -end +chunk_provenance(c::Chunk) = render_provenance(c.metadata) """ load_funsql_examples(path; source="FunSQL-examples") -> Vector{SourceDocument} diff --git a/src/prompt.jl b/src/prompt.jl new file mode 100644 index 0000000..8b0455a --- /dev/null +++ b/src/prompt.jl @@ -0,0 +1,283 @@ +""" + Prompt + +Prompt construction for retrieval-augmented FunSQL generation. This module sits +between retrieval and the language model: it takes the chunks a vector store +returned for a question — OMOP schema definitions, FunSQL examples, prose docs — +and the user's natural-language *analytical* question, and renders one grounded +prompt asking the model to translate the question into a FunSQL query. + +The two moving parts injected into every prompt are: + +1. **Retrieved context** — the ranked chunks from [`retrieve`](@ref HealthLLM.Storage.retrieve)/[`search`](@ref HealthLLM.Storage.search) + (or raw strings / [`Chunk`](@ref HealthLLM.Ingestion.Chunk)s), formatted as numbered, optionally + provenance-tagged blocks so the model can ground its answer and cite sources. +2. **The analytical question** — the user's natural-language request, verbatim. + +## Interface + +- [`PromptTemplate`](@ref) — the reusable shape (system message, section headers, + per-chunk formatting, and context budgets). [`DEFAULT_FUNSQL_TEMPLATE`](@ref) is + the OMOP/FunSQL default. +- [`format_context`](@ref)`(template, hits)` — render retrieved chunks to a string. +- [`build_prompt`](@ref)`(question, hits; template)` — the entry point; returns + `(; system, user, prompt)`. + +## Example + +```julia +store = LocalVectorStore(embedding_dimension()) +add!(store, embed(chunks), chunks) + +hits = retrieve(store, "How many distinct patients had a diabetes diagnosis in 2020?", 6) +p = build_prompt("How many distinct patients had a diabetes diagnosis in 2020?", hits) + +# split form, for a proper system/user conversation: +msg = PromptingTools.aigenerate(p.system * "\\n\\n" * p.user; model="llama3.2") +# or the single joined string, for the {input_query}-style path: +p.prompt +``` +""" +module Prompt + +using ..Utils: render_provenance + +export FUNSQL_SYSTEM_PROMPT, PromptTemplate, DEFAULT_FUNSQL_TEMPLATE, + format_context, build_prompt + +""" + FUNSQL_SYSTEM_PROMPT + +Default system message for FunSQL generation over the OMOP CDM. It fixes the +model's role and the rules that keep generations grounded: build queries only +from tables/columns that appear in the retrieved context, follow the retrieved +examples' FunSQL style, flag missing schema instead of inventing it, and return +the query as a single fenced `julia` code block. + +The **Grounding** rules are the load-bearing part: they instruct the model to +treat the retrieved context as the *only* admissible source of schema +(table/column) names and FunSQL syntax, and to refuse rather than hallucinate +when the context is insufficient. Weaken them and generations drift toward +plausible-but-nonexistent OMOP columns. +""" +const FUNSQL_SYSTEM_PROMPT = """ +You are an expert data analyst for the OHDSI OMOP Common Data Model (CDM), \ +writing analytical queries with the Julia package FunSQL.jl. + +Your task: translate the user's natural-language analytical question into a \ +correct FunSQL query against the OMOP CDM, using ONLY the retrieved context below. + +Grounding (strict — this overrides your prior knowledge): +- The retrieved context is your ONLY source of truth for table names, column \ + names, and FunSQL syntax. Every table and column you reference MUST appear \ + verbatim in the context. Do not rely on memory of the OMOP CDM or of FunSQL, \ + and do not invent, guess, pluralise, or "correct" any identifier. +- Before using a table or column, confirm it is present in the context. If a name \ + you need is not there, DO NOT substitute a similar-looking one — treat it as \ + unavailable. +- Reproduce FunSQL constructs (`From`, `Where`, `Join`, `Group`, `Select`, \ + `Get`, `Agg`, `Fun`, `|>`) only as they appear in the retrieved examples. Do \ + not use SQL string syntax or FunSQL features you cannot see demonstrated in the \ + context. +- If the context lacks a table, column, or construct required to answer the \ + question, state exactly what is missing, then generate the closest correct \ + query you can from what IS available, labelling any assumption. Never paper over \ + a gap with a fabricated identifier. + +Style: +- Follow the idioms of the retrieved FunSQL examples. +- Join through the standard OMOP keys shown in the context (e.g. `person_id`) and \ + filter on the `*_concept_id` columns present in the context where the question \ + implies a clinical concept. + +Output: +- Return the query as ONE fenced code block tagged `julia`, followed by a one- or \ + two-sentence explanation that names the context tables/columns you used. Do not \ + pad the answer with unrelated commentary. +""" + +""" + PromptTemplate(; kwargs...) + +Reusable shape for a retrieval-augmented prompt. A template is data — the +`system` message, the section headers, how each retrieved chunk is rendered, and +how much context to admit — so the same construction logic serves different +models and tasks by swapping fields. + +# Fields +- `system::String`: System message. Defaults to [`FUNSQL_SYSTEM_PROMPT`](@ref). +- `context_header::String`: Heading printed above the retrieved chunks. +- `question_header::String`: Heading printed above the user's question. +- `answer_cue::String`: Trailing line that cues the model to answer (e.g. a + `# FunSQL query` header). Empty to omit. +- `empty_context_note::String`: Placeholder used when no chunks are supplied, so + the model is told the context is empty rather than seeing a blank section. +- `chunk_label::String`: Per-chunk label prefix; numbered as `"[