From ba86c546464303aba583e325469ee999f793a53e Mon Sep 17 00:00:00 2001 From: Yash-Chindam Date: Sat, 3 Oct 2026 19:36:40 +0530 Subject: [PATCH] feat: route on a calibrated classifier, live load, and measured quality Section 7.2 asks for a lightweight classifier or calibrated model for task and complexity, and lists structured-output requirement, current queue delay, GPU capacity, and historical quality by task and model as routing features. Task was keyword matching, complexity did not exist, queue delay was a static catalog estimate, and quality was one figure per model. Task and complexity now come from a multinomial naive Bayes model trained at start-up from a committed dataset, with posteriors temperature-scaled against a held-out split. A prediction under 0.5 abstains to the general task, and a declared task is never overridden. Keyword rules remain the fallback when no dataset is deployed. Capability stays deterministic: a structured request never reaches a model that cannot produce structured output. Scoring uses benchmarked quality for the task, an exponentially weighted average of observed queue delay in place of the estimate, and engine saturation from the last scrape. Predicted complexity pulls a request toward higher measured quality. Saturation can tip an eligible request to the external model but never overrides privacy. The response and the trace report how the task was established. Co-Authored-By: Claude Opus 5.5 --- README.md | 38 ++++ benchmarks/datasets/routing-tasks-v1.jsonl | 42 ++++ config/registry.yaml | 1 + config/routing/task-classifier-v1.jsonl | 185 +++++++++++++++ src/llm_router/app.py | 20 ++ src/llm_router/classifier.py | 253 +++++++++++++++++++++ src/llm_router/config.py | 3 + src/llm_router/load.py | 67 ++++++ src/llm_router/models.py | 12 + src/llm_router/registry.py | 29 ++- src/llm_router/routing.py | 87 ++++++- src/llm_router/tracing.py | 3 + tests/integration/test_api.py | 49 ++++ tests/unit/test_classifier.py | 142 ++++++++++++ tests/unit/test_routing_features.py | 223 ++++++++++++++++++ 15 files changed, 1147 insertions(+), 7 deletions(-) create mode 100644 benchmarks/datasets/routing-tasks-v1.jsonl create mode 100644 config/routing/task-classifier-v1.jsonl create mode 100644 src/llm_router/classifier.py create mode 100644 src/llm_router/load.py create mode 100644 tests/unit/test_classifier.py create mode 100644 tests/unit/test_routing_features.py diff --git a/README.md b/README.md index 84bb619..6eba090 100644 --- a/README.md +++ b/README.md @@ -73,6 +73,44 @@ workflow tags `main` (`vMAJOR.MINOR.PATCH`) and publishes a GitHub Release with auto-generated notes. Tags are never created by hand, and nothing is ever tagged off a branch other than `main`. +## Routing + +Privacy, tenant entitlement, context size, and capability are hard filters: deterministic, and +applied before any score is computed. Among the models that remain, the router scores measured +quality for the task against observed queue delay, cost, engine saturation, and how complex the +request is predicted to be. + +**Task and complexity** come from a calibrated classifier — multinomial naive Bayes over word +unigrams and bigrams, trained at start-up from +[`config/routing/task-classifier-v1.jsonl`](config/routing/task-classifier-v1.jsonl). It needs no +accelerator and no extra dependency, so a routing decision never waits on the models it is +choosing between. Its posteriors are temperature-scaled against a held-out split, so a confidence +reads as a probability, and a prediction under 0.5 abstains to the `general` task rather than +being trusted. A task the caller declares in `routing.task` is never overridden. Without the +dataset the router falls back to keyword rules. + +On the held-out prompts in +[`benchmarks/datasets/routing-tasks-v1.jsonl`](benchmarks/datasets/routing-tasks-v1.jsonl) the +classifier gets every task right, 88% of complexity labels, and an expected calibration error of +0.035. Both datasets are small and were written by hand in one voice, so treat that as a floor +check on the mechanism, not as evidence of accuracy on real traffic: replace them with labelled +production prompts before relying on the numbers. + +| Routing feature | Source | +|---|---| +| Task and complexity | The classifier; `low` leaves work on the cheapest capable model, `high` outweighs the specialization and cost terms. | +| Structured-output requirement | `routing.structured` excludes any model whose card sets `supports_structured_output: false`. | +| Quality by task and model | Mean benchmarked quality for the task from the catalog; the card's headline `quality` only for a task never measured. | +| Current queue delay | An exponentially weighted average of what requests for that model actually waited, replacing the catalog estimate after the first observation. | +| GPU capacity | Engine saturation — the worse of KV-cache occupancy and the share of admitted work not yet started — from the last metrics scrape, discarded after 30 seconds. | + +One gateway faces one engine, so saturation costs every local model equally: it can tip an +eligible request to the approved external model, and it never reorders local models or overrides +privacy. Required modality is not a routing feature yet; message content is text only. + +Every response reports `task_source` (`declared`, `classifier`, `abstained`, `keyword`, or +`cached`), `task_confidence`, and `complexity` beside the route reason. + ## Serving backends `ROUTER_BACKEND=mock` (the default) keeps CI deterministic and GPU-free. diff --git a/benchmarks/datasets/routing-tasks-v1.jsonl b/benchmarks/datasets/routing-tasks-v1.jsonl new file mode 100644 index 0000000..d3ef70f --- /dev/null +++ b/benchmarks/datasets/routing-tasks-v1.jsonl @@ -0,0 +1,42 @@ +{"prompt": "Extract the order number and delivery date from this confirmation.", "task": "extraction", "complexity": "low"} +{"prompt": "Return the patient name and appointment time as JSON.", "task": "extraction", "complexity": "low"} +{"prompt": "Pull the VAT number out of the supplier letter.", "task": "extraction", "complexity": "low"} +{"prompt": "Extract the amount due from this bill.", "task": "extraction", "complexity": "low"} +{"prompt": "Parse the following address into its fields.", "task": "extraction", "complexity": "low"} +{"prompt": "Extract the serial number and model from the warranty card.", "task": "extraction", "complexity": "low"} +{"prompt": "Classify this message as complaint, question or praise.", "task": "classification", "complexity": "low"} +{"prompt": "Choose one label for this support email.", "task": "classification", "complexity": "low"} +{"prompt": "Is this ticket a bug or a feature request?", "task": "classification", "complexity": "low"} +{"prompt": "Classify the sentiment of this review.", "task": "classification", "complexity": "low"} +{"prompt": "Which category does this expense belong to?", "task": "classification", "complexity": "low"} +{"prompt": "Label this transaction as personal or business.", "task": "classification", "complexity": "low"} +{"prompt": "According to the documents, what is the cancellation fee?", "task": "rag", "complexity": "medium"} +{"prompt": "Using the provided context, answer how refunds are issued.", "task": "rag", "complexity": "medium"} +{"prompt": "Answer only from the passages below: what is the SLA?", "task": "rag", "complexity": "low"} +{"prompt": "Based on the retrieved context, who signs off the release?", "task": "rag", "complexity": "medium"} +{"prompt": "Use the provided context to answer the question about holidays.", "task": "rag", "complexity": "medium"} +{"prompt": "From the documents provided, what is the retention period?", "task": "rag", "complexity": "low"} +{"prompt": "Summarize the annual review.", "task": "summarization", "complexity": "medium"} +{"prompt": "Give me a summary of this long email.", "task": "summarization", "complexity": "medium"} +{"prompt": "Summarise the transcript in five bullet points.", "task": "summarization", "complexity": "medium"} +{"prompt": "Write a brief summary of the audit findings.", "task": "summarization", "complexity": "medium"} +{"prompt": "Condense the article into its key points.", "task": "summarization", "complexity": "medium"} +{"prompt": "Provide a short overview of this policy document.", "task": "summarization", "complexity": "medium"} +{"prompt": "Reason step by step about why the test fails.", "task": "reasoning", "complexity": "high"} +{"prompt": "Prove that the algorithm always halts.", "task": "reasoning", "complexity": "high"} +{"prompt": "Analyze deeply the trade-offs between these two designs.", "task": "reasoning", "complexity": "high"} +{"prompt": "Work through the puzzle and explain each step.", "task": "reasoning", "complexity": "high"} +{"prompt": "Deduce the correct ordering from the constraints.", "task": "reasoning", "complexity": "high"} +{"prompt": "Reason about whether the claim follows from the premises.", "task": "reasoning", "complexity": "high"} +{"prompt": "Critique this system design.", "task": "critique", "complexity": "high"} +{"prompt": "Find flaws in this proposal.", "task": "critique", "complexity": "high"} +{"prompt": "Review this essay and point out its weaknesses.", "task": "critique", "complexity": "medium"} +{"prompt": "Give a critical review of this architecture.", "task": "critique", "complexity": "high"} +{"prompt": "What is wrong with this argument?", "task": "critique", "complexity": "high"} +{"prompt": "Point out the weaknesses in this security plan.", "task": "critique", "complexity": "high"} +{"prompt": "Hello", "task": "general", "complexity": "low"} +{"prompt": "Tell me a joke about computers.", "task": "general", "complexity": "low"} +{"prompt": "What is the capital of Japan?", "task": "general", "complexity": "low"} +{"prompt": "Write a short poem about the sea.", "task": "general", "complexity": "medium"} +{"prompt": "Thanks for your help.", "task": "general", "complexity": "low"} +{"prompt": "Draft an email inviting the team to lunch.", "task": "general", "complexity": "medium"} diff --git a/config/registry.yaml b/config/registry.yaml index 03fea7c..5f0e6ca 100644 --- a/config/registry.yaml +++ b/config/registry.yaml @@ -144,6 +144,7 @@ benchmarks: container_digest: sha256:mock-container engine_revision: mock-engine-0.1.0 model_revision: mock-small@sha256:dev + task: extraction concurrency: 32 prompt_tokens_p50: 420 prompt_tokens_p95: 1100 diff --git a/config/routing/task-classifier-v1.jsonl b/config/routing/task-classifier-v1.jsonl new file mode 100644 index 0000000..1a8c459 --- /dev/null +++ b/config/routing/task-classifier-v1.jsonl @@ -0,0 +1,185 @@ +{"prompt": "Extract the invoice number and total from this invoice.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Extract customer fields", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Extract the claim id from the note below and reply as JSON.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Pull out every date mentioned in the contract and return them as a JSON array.", "task": "extraction", "complexity": "medium", "split": "train"} +{"prompt": "Extract the fields from this purchase order: vendor, amount, due date.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Return the sender, recipient and subject of this email as JSON.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Parse the address into street, city, postcode and country.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Extract the claim fields", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "From the lab report, extract each test name with its value and unit.", "task": "extraction", "complexity": "medium", "split": "train"} +{"prompt": "Get the order id, sku and quantity out of this shipping notice.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Extract the invoice fields", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Fill this JSON schema using the values found in the document.", "task": "extraction", "complexity": "medium", "split": "train"} +{"prompt": "List the named people and organisations that appear in the article as JSON.", "task": "extraction", "complexity": "medium", "split": "train"} +{"prompt": "Extract the policy number and the effective date from the letter.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Pull the phone numbers and email addresses out of the signature block.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Extract the fields", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Read the receipt and output merchant, date, tax and total as key value pairs.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Extract all line items with description, unit price and quantity from the quote.", "task": "extraction", "complexity": "medium", "split": "train"} +{"prompt": "Identify the medication names and dosages in the discharge summary and return JSON.", "task": "extraction", "complexity": "medium", "split": "train"} +{"prompt": "Extract the json schema fields from the form text.", "task": "extraction", "complexity": "low", "split": "train"} +{"prompt": "Extract parties, governing law, term and termination clauses from this agreement into a structured record, noting any clause that is missing.", "task": "extraction", "complexity": "high", "split": "train"} +{"prompt": "Classify this ticket as billing, technical, or other.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Classify this refund request", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Choose one label for this review: positive, negative or neutral.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Which category does this email belong to: sales, support or spam?", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Label the sentiment of the following tweet.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Classify this ticket", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Is this message spam or not spam? Answer with one word.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Assign a priority label of low, medium or high to the incident.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Tag this document with exactly one topic from the list.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Decide whether the transaction is fraudulent or legitimate.", "task": "classification", "complexity": "medium", "split": "train"} +{"prompt": "Categorise the support request into one of the queues below.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Classify the intent of this utterance: book, cancel, or reschedule.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "classify this private matter", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Pick the single best category for this product listing.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Label each sentence as claim, evidence or neither.", "task": "classification", "complexity": "medium", "split": "train"} +{"prompt": "Choose one label for this item", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Is the tone of this message formal or informal?", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Route this request: classify it as hardware, software or account.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Determine the language of the text and answer with its name only.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Classify the severity of this bug report as critical, major or minor.", "task": "classification", "complexity": "low", "split": "train"} +{"prompt": "Classify each of these forty clauses by risk category and flag any that fit more than one category with a short justification.", "task": "classification", "complexity": "high", "split": "train"} +{"prompt": "Using the provided context, answer the question.", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "According to the documents, what is the refund window?", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Answer only from the passages below and cite the passage you used.", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Based on the retrieved context, when does the warranty expire?", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Context: the handbook excerpt follows. Question: how many leave days do staff get?", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Given the provided context, what did the committee decide?", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Answer the user question using the knowledge base articles supplied.", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "According to the documents provided, who approves expenses over the limit?", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Use the sources below to answer, and say you do not know if they do not cover it.", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "From the retrieved passages, what is the maximum file size allowed?", "task": "rag", "complexity": "low", "split": "train"} +{"prompt": "With reference to the attached context, explain the escalation path.", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Answer from the provided context only: which regions are supported?", "task": "rag", "complexity": "low", "split": "train"} +{"prompt": "Here are search results. Use them to answer what the default timeout is.", "task": "rag", "complexity": "low", "split": "train"} +{"prompt": "Grounded in the documents, what are the eligibility requirements?", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Question answering over the provided context: what changed in version two?", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Refer to the context and tell me the name of the data controller.", "task": "rag", "complexity": "low", "split": "train"} +{"prompt": "According to the documents, which clause covers liability?", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Use the retrieved chunks to answer and quote the supporting sentence.", "task": "rag", "complexity": "medium", "split": "train"} +{"prompt": "Consult the provided context and state the shipping cost for Europe.", "task": "rag", "complexity": "low", "split": "train"} +{"prompt": "Reconcile the three provided context documents, which disagree, and answer which policy is currently in force, citing each source.", "task": "rag", "complexity": "high", "split": "train"} +{"prompt": "Summarize the quarterly report", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Summarize this report", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Give me a summary of the meeting notes.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Write a short summary of the article in three sentences.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Condense this email thread into the key points.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Summarize the report", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Provide a brief overview of the document below.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "TL;DR of this long post please.", "task": "summarization", "complexity": "low", "split": "train"} +{"prompt": "Summarise the customer call transcript for the account manager.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Shorten this text to one paragraph while keeping the main points.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Produce an executive summary of the incident review.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Summarize the patient intake notes", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Give the main takeaways from this research abstract.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Boil the changelog down to a two line summary.", "task": "summarization", "complexity": "low", "split": "train"} +{"prompt": "Summarize the ticket backlog", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Write a headline and a one sentence summary for this story.", "task": "summarization", "complexity": "low", "split": "train"} +{"prompt": "Recap what was agreed in the discussion above.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Summarise the chapter in bullet points.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Brief summary of the release notes for a non technical reader.", "task": "summarization", "complexity": "medium", "split": "train"} +{"prompt": "Summarize these twelve board papers into a single briefing that keeps every decision, owner and deadline and notes where papers conflict.", "task": "summarization", "complexity": "high", "split": "train"} +{"prompt": "Analyze deeply", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Reason step by step about this proof", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Prove that the sum of two even numbers is even.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Work through this logic puzzle and explain each step.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Reason step by step to find which schedule satisfies all constraints.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Analyze deeply why the system deadlocks under this interleaving.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Derive the time complexity of this algorithm and justify it.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Think through the trade-offs and decide which design is correct.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Solve this word problem, showing your reasoning.", "task": "reasoning", "complexity": "medium", "split": "train"} +{"prompt": "Given these premises, what follows logically? Explain why.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Prove or disprove the following claim about prime numbers.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Plan the migration in ordered steps and reason about the risks of each.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Determine the root cause from these symptoms by elimination.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Reason about whether this invariant holds after every iteration.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Compute the probability and explain the reasoning behind each step.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Which of the two hypotheses better explains the data, and why?", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Analyze deeply the second order effects of this pricing change.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Deduce the missing value in the sequence and justify the rule.", "task": "reasoning", "complexity": "medium", "split": "train"} +{"prompt": "Carefully reason through the edge cases of this state machine.", "task": "reasoning", "complexity": "high", "split": "train"} +{"prompt": "Step by step, work out how many moves the puzzle needs.", "task": "reasoning", "complexity": "medium", "split": "train"} +{"prompt": "Critique this architecture", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Find flaws in this argument.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Review this design document and point out its weaknesses.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Critique the essay and suggest how to improve it.", "task": "critique", "complexity": "medium", "split": "train"} +{"prompt": "What is wrong with this plan? Be critical.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Find flaws in the proposed database schema.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Give a critical review of this pull request description.", "task": "critique", "complexity": "medium", "split": "train"} +{"prompt": "Critique this business case and list the risks it ignores.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Evaluate the strengths and weaknesses of this proposal.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Identify the logical fallacies in the following paragraph.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Play devil's advocate against this strategy.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Critique my cover letter honestly.", "task": "critique", "complexity": "medium", "split": "train"} +{"prompt": "Review the code below and point out bugs and bad practices.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Find the weaknesses in this security model.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Assess this experiment design and say what would invalidate its results.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Give harsh but fair feedback on this pitch.", "task": "critique", "complexity": "medium", "split": "train"} +{"prompt": "Critique the methodology of this study.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Point out the flaws and gaps in this test plan.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Red team this policy and find how it could be abused.", "task": "critique", "complexity": "high", "split": "train"} +{"prompt": "Peer review this abstract and list its shortcomings.", "task": "critique", "complexity": "medium", "split": "train"} +{"prompt": "Hello there", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "hello", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "hello again", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Hi, how are you today?", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Write a short poem about autumn.", "task": "general", "complexity": "medium", "split": "train"} +{"prompt": "What is the capital of France?", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Tell me a joke.", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Draft a polite email asking to move our meeting.", "task": "general", "complexity": "medium", "split": "train"} +{"prompt": "Translate good morning into Spanish.", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Thanks, that was helpful.", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Can you help me with something?", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Write a function that reverses a string in Python.", "task": "general", "complexity": "medium", "split": "train"} +{"prompt": "Give me three ideas for a team offsite.", "task": "general", "complexity": "medium", "split": "train"} +{"prompt": "What does HTTP stand for?", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Rewrite this sentence to sound friendlier.", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Explain what a hash table is.", "task": "general", "complexity": "medium", "split": "train"} +{"prompt": "Suggest a name for my bakery.", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "shared state probe", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Good morning", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "How do I boil an egg?", "task": "general", "complexity": "low", "split": "train"} +{"prompt": "Write a product description for a steel water bottle.", "task": "general", "complexity": "medium", "split": "train"} +{"prompt": "Extract the invoice date and the supplier name.", "task": "extraction", "complexity": "low", "split": "calibration"} +{"prompt": "Return the contract value and currency as JSON.", "task": "extraction", "complexity": "low", "split": "calibration"} +{"prompt": "Pull out the tracking number from this message.", "task": "extraction", "complexity": "low", "split": "calibration"} +{"prompt": "Extract every figure with its label from the financial table.", "task": "extraction", "complexity": "medium", "split": "calibration"} +{"prompt": "Parse the name and date of birth from the form.", "task": "extraction", "complexity": "low", "split": "calibration"} +{"prompt": "Extract the customer id", "task": "extraction", "complexity": "low", "split": "calibration"} +{"prompt": "Classify this refund request please", "task": "classification", "complexity": "low", "split": "calibration"} +{"prompt": "Which label fits this complaint best?", "task": "classification", "complexity": "low", "split": "calibration"} +{"prompt": "Is this review positive or negative?", "task": "classification", "complexity": "low", "split": "calibration"} +{"prompt": "Categorise the email as urgent or not urgent.", "task": "classification", "complexity": "low", "split": "calibration"} +{"prompt": "Classify this document by department.", "task": "classification", "complexity": "low", "split": "calibration"} +{"prompt": "Choose one label for the feedback below.", "task": "classification", "complexity": "low", "split": "calibration"} +{"prompt": "According to the documents, what is the notice period?", "task": "rag", "complexity": "medium", "split": "calibration"} +{"prompt": "Using the provided context, list the supported currencies.", "task": "rag", "complexity": "low", "split": "calibration"} +{"prompt": "Answer from the passages below: who is the account owner?", "task": "rag", "complexity": "low", "split": "calibration"} +{"prompt": "Based on the provided context, is overtime paid?", "task": "rag", "complexity": "medium", "split": "calibration"} +{"prompt": "Use the retrieved documents to answer the customer.", "task": "rag", "complexity": "medium", "split": "calibration"} +{"prompt": "From the context supplied, what is the return address?", "task": "rag", "complexity": "low", "split": "calibration"} +{"prompt": "Summarize the incident report", "task": "summarization", "complexity": "medium", "split": "calibration"} +{"prompt": "Give a short summary of this contract.", "task": "summarization", "complexity": "medium", "split": "calibration"} +{"prompt": "Summarise the feedback in two sentences.", "task": "summarization", "complexity": "medium", "split": "calibration"} +{"prompt": "Condense these notes into key points.", "task": "summarization", "complexity": "medium", "split": "calibration"} +{"prompt": "Write a brief summary of the proposal.", "task": "summarization", "complexity": "medium", "split": "calibration"} +{"prompt": "TL;DR this thread.", "task": "summarization", "complexity": "low", "split": "calibration"} +{"prompt": "Reason step by step about which option is cheaper overall.", "task": "reasoning", "complexity": "high", "split": "calibration"} +{"prompt": "Prove that this recursive function terminates.", "task": "reasoning", "complexity": "high", "split": "calibration"} +{"prompt": "Analyze deeply the cause of the regression.", "task": "reasoning", "complexity": "high", "split": "calibration"} +{"prompt": "Work out the answer and explain each step of your reasoning.", "task": "reasoning", "complexity": "medium", "split": "calibration"} +{"prompt": "Deduce who owns the fish from these clues.", "task": "reasoning", "complexity": "high", "split": "calibration"} +{"prompt": "Reason through the constraints and find a valid assignment.", "task": "reasoning", "complexity": "high", "split": "calibration"} +{"prompt": "Critique this API design.", "task": "critique", "complexity": "high", "split": "calibration"} +{"prompt": "Find flaws in my reasoning above.", "task": "critique", "complexity": "high", "split": "calibration"} +{"prompt": "Review this plan and list its weaknesses.", "task": "critique", "complexity": "high", "split": "calibration"} +{"prompt": "Give critical feedback on this draft.", "task": "critique", "complexity": "medium", "split": "calibration"} +{"prompt": "What are the shortcomings of this approach?", "task": "critique", "complexity": "high", "split": "calibration"} +{"prompt": "Point out the flaws in this contract clause.", "task": "critique", "complexity": "high", "split": "calibration"} +{"prompt": "Hey", "task": "general", "complexity": "low", "split": "calibration"} +{"prompt": "What time zone is Tokyo in?", "task": "general", "complexity": "low", "split": "calibration"} +{"prompt": "Write a limerick about a cat.", "task": "general", "complexity": "medium", "split": "calibration"} +{"prompt": "Help me word a thank you note.", "task": "general", "complexity": "medium", "split": "calibration"} +{"prompt": "What is two plus two?", "task": "general", "complexity": "low", "split": "calibration"} +{"prompt": "Recommend a book about history.", "task": "general", "complexity": "low", "split": "calibration"} diff --git a/src/llm_router/app.py b/src/llm_router/app.py index 9d749b6..5ff9687 100644 --- a/src/llm_router/app.py +++ b/src/llm_router/app.py @@ -39,9 +39,11 @@ exact_cache_eligible, semantic_cache_eligible, ) +from llm_router.classifier import TaskClassifier, load_classifier from llm_router.config import Settings, get_settings from llm_router.engine_stats import ColdStartTracker, EngineStatsCollector from llm_router.evaluation import structured_output_valid +from llm_router.load import LoadTracker from llm_router.models import ( ChatCompletionChoice, ChatCompletionRequest, @@ -95,6 +97,14 @@ def _load_catalog(path: str) -> Registry | None: return load_registry(path) +def _load_task_classifier(path: str) -> TaskClassifier | None: + """Train the routing classifier, or fall back to keyword rules when absent.""" + + if not Path(path).exists(): + return None + return load_classifier(path) + + def create_app( settings: Settings | None = None, *, @@ -112,10 +122,13 @@ def create_app( policy_version = ( catalog.policy.version if catalog is not None else runtime_settings.routing_policy_version ) + load = LoadTracker() router = Router( profiles=profiles, external_fallback_enabled=runtime_settings.external_fallback_enabled, registry=catalog, + classifier=_load_task_classifier(runtime_settings.task_classifier_path), + load=load, ) admission = AdmissionController( runtime_settings.max_concurrency, @@ -278,6 +291,7 @@ async def prometheus_metrics() -> Response: stats = await engine_telemetry.sample() if stats is not None: telemetry.record_engine_stats(stats, engine=engine_label) + load.observe_engine(stats) payload, content_type = telemetry.render() return Response(content=payload, media_type=content_type) @@ -647,6 +661,9 @@ async def _complete( async with admission.slot(): telemetry.queued_requests.dec() queue_seconds = time.perf_counter() - started + # What this request actually waited becomes the next + # request's estimate for the same model. + load.observe_queue(decision.profile.id, queue_seconds * 1000) telemetry.inflight_requests.inc() try: if payload.stream: @@ -713,6 +730,9 @@ async def _complete( "adapter_id": decision.adapter_id, "adapter_revision": decision.adapter_revision, "task": decision.task.value, + "task_source": decision.task_source, + "task_confidence": decision.task_confidence, + "complexity": decision.complexity, "reason": decision.reason, "score": decision.score, "candidate_count": decision.candidate_count, diff --git a/src/llm_router/classifier.py b/src/llm_router/classifier.py new file mode 100644 index 0000000..41037a3 --- /dev/null +++ b/src/llm_router/classifier.py @@ -0,0 +1,253 @@ +"""Calibrated task and complexity prediction for routing (section 7.2). + +Section 7.2 keeps privacy and hard capability restrictions deterministic and +asks for a lightweight classifier or calibrated model for task and complexity. +This is a multinomial naive Bayes model over word unigrams and bigrams, trained +at start-up from a committed dataset. It needs no accelerator and no extra +dependency, and trains in milliseconds, so the routing decision never waits on +the models it is choosing between. + +Naive Bayes posteriors are overconfident, so they are temperature-scaled +against a held-out calibration split. A confidence can then be read as a +probability, and a prediction below the threshold falls back to the general +task rather than being trusted. +""" + +import json +import math +import re +from collections import Counter, defaultdict +from collections.abc import Iterable, Sequence +from dataclasses import dataclass +from enum import StrEnum +from functools import lru_cache +from itertools import pairwise +from pathlib import Path +from typing import Generic, TypeVar + +from llm_router.models import TaskClass + +DEFAULT_DATASET = "config/routing/task-classifier-v1.jsonl" +# The instruction is nearly always at the start or the end of a prompt. Long +# retrieved context in between says what the prompt is about, not what is being +# asked, and would otherwise drown the instruction out. +HEAD_TOKENS = 48 +TAIL_TOKENS = 48 +SMOOTHING = 0.5 +TEMPERATURES = (0.5, 0.75, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, 8.0, 12.0, 16.0) +DEFAULT_CONFIDENCE_THRESHOLD = 0.5 + +_TOKEN = re.compile(r"[a-z0-9]+") +Label = TypeVar("Label") + + +class Complexity(StrEnum): + LOW = "low" + MEDIUM = "medium" + HIGH = "high" + + +class ClassifierError(RuntimeError): + """Raised when the training data cannot produce a usable model.""" + + +def featurize(text: str) -> list[str]: + """Turn a prompt into unigram, bigram, and coarse length features.""" + + tokens = _TOKEN.findall(text.lower()) + if len(tokens) > HEAD_TOKENS + TAIL_TOKENS: + tokens = tokens[:HEAD_TOKENS] + tokens[-TAIL_TOKENS:] + features = list(tokens) + features.extend(f"{left}_{right}" for left, right in pairwise(tokens)) + # Length carries signal for complexity that word identity alone does not. + words = len(text.split()) + features.append("__short__" if words <= 12 else "__medium__" if words <= 60 else "__long__") + return features + + +@dataclass(frozen=True) +class Prediction(Generic[Label]): + label: Label + confidence: float + distribution: dict[Label, float] + + +class NaiveBayes(Generic[Label]): + """Multinomial naive Bayes with additive smoothing and temperature scaling.""" + + def __init__(self, examples: Sequence[tuple[list[str], Label]]) -> None: + if not examples: + raise ClassifierError("cannot train on an empty dataset") + self._label_counts: Counter[Label] = Counter(label for _, label in examples) + self._feature_counts: dict[Label, Counter[str]] = defaultdict(Counter) + for features, label in examples: + self._feature_counts[label].update(features) + self._vocabulary = { + feature for counts in self._feature_counts.values() for feature in counts + } + self._totals = { + label: sum(counts.values()) for label, counts in self._feature_counts.items() + } + self._examples = len(examples) + self.temperature = 1.0 + + @property + def labels(self) -> tuple[Label, ...]: + return tuple(self._label_counts) + + def _log_joint(self, features: Iterable[str]) -> dict[Label, float]: + known = [feature for feature in features if feature in self._vocabulary] + vocabulary = len(self._vocabulary) + scores: dict[Label, float] = {} + for label, count in self._label_counts.items(): + score = math.log(count / self._examples) + denominator = self._totals[label] + SMOOTHING * vocabulary + counts = self._feature_counts[label] + for feature in known: + score += math.log((counts[feature] + SMOOTHING) / denominator) + scores[label] = score + return scores + + def predict(self, features: Iterable[str]) -> Prediction[Label]: + scores = self._log_joint(features) + scaled = {label: score / self.temperature for label, score in scores.items()} + peak = max(scaled.values()) + exponentials = {label: math.exp(score - peak) for label, score in scaled.items()} + total = sum(exponentials.values()) + distribution = {label: value / total for label, value in exponentials.items()} + label = max(distribution, key=lambda candidate: distribution[candidate]) + return Prediction(label=label, confidence=distribution[label], distribution=distribution) + + def calibrate(self, examples: Sequence[tuple[list[str], Label]]) -> float: + """Pick the temperature that minimizes held-out negative log-likelihood.""" + + if not examples: + raise ClassifierError("cannot calibrate without a held-out split") + + def loss(temperature: float) -> float: + self.temperature = temperature + return -sum( + math.log(max(self.predict(features).distribution.get(label, 0.0), 1e-12)) + for features, label in examples + ) + + self.temperature = min(TEMPERATURES, key=loss) + return self.temperature + + +def expected_calibration_error( + model: NaiveBayes[Label], examples: Sequence[tuple[list[str], Label]], *, bins: int = 5 +) -> float: + """Mean gap between stated confidence and observed accuracy, weighted by bin size.""" + + if not examples: + raise ClassifierError("cannot measure calibration on an empty sample") + buckets: dict[int, list[tuple[float, bool]]] = defaultdict(list) + for features, label in examples: + prediction = model.predict(features) + index = min(int(prediction.confidence * bins), bins - 1) + buckets[index].append((prediction.confidence, prediction.label == label)) + error = 0.0 + for members in buckets.values(): + confidence = sum(value for value, _ in members) / len(members) + accuracy = sum(correct for _, correct in members) / len(members) + error += abs(confidence - accuracy) * len(members) / len(examples) + return error + + +@dataclass(frozen=True) +class LabelledPrompt: + prompt: str + task: TaskClass + complexity: Complexity + split: str = "train" + + +def load_labelled_prompts(path: str | Path) -> tuple[LabelledPrompt, ...]: + """Read a JSON Lines file of prompts labelled with task and complexity.""" + + rows: list[LabelledPrompt] = [] + for number, line in enumerate(Path(path).read_text(encoding="utf-8").splitlines(), start=1): + if not line.strip(): + continue + try: + document = json.loads(line) + rows.append( + LabelledPrompt( + prompt=str(document["prompt"]), + task=TaskClass(document["task"]), + complexity=Complexity(document["complexity"]), + split=str(document.get("split", "train")), + ) + ) + except (ValueError, KeyError) as error: + raise ClassifierError(f"{path}:{number} is not a usable example: {error}") from error + if not rows: + raise ClassifierError(f"{path} contains no examples") + return tuple(rows) + + +@dataclass(frozen=True) +class RoutingPrediction: + task: TaskClass + task_confidence: float + complexity: Complexity + complexity_confidence: float + # False when the task fell back to general because confidence was too low. + trusted: bool + + +class TaskClassifier: + """Predicts task and complexity, abstaining to the general task when unsure.""" + + def __init__( + self, + task_model: NaiveBayes[TaskClass], + complexity_model: NaiveBayes[Complexity], + *, + confidence_threshold: float = DEFAULT_CONFIDENCE_THRESHOLD, + ) -> None: + self.task_model = task_model + self.complexity_model = complexity_model + self.confidence_threshold = confidence_threshold + + def predict(self, prompt: str) -> RoutingPrediction: + features = featurize(prompt) + task = self.task_model.predict(features) + complexity = self.complexity_model.predict(features) + trusted = task.confidence >= self.confidence_threshold + return RoutingPrediction( + # The general task is served by the broadest models, so abstaining + # to it is the conservative choice when the model is unsure. + task=task.label if trusted else TaskClass.GENERAL, + task_confidence=task.confidence, + complexity=complexity.label, + complexity_confidence=complexity.confidence, + trusted=trusted, + ) + + +def train_classifier( + rows: Sequence[LabelledPrompt], + *, + confidence_threshold: float = DEFAULT_CONFIDENCE_THRESHOLD, +) -> TaskClassifier: + """Fit both heads on the training split and calibrate on the held-out one.""" + + train = [row for row in rows if row.split == "train"] + held_out = [row for row in rows if row.split == "calibration"] + if not train or not held_out: + raise ClassifierError("the dataset needs both a train and a calibration split") + + task_model = NaiveBayes([(featurize(row.prompt), row.task) for row in train]) + task_model.calibrate([(featurize(row.prompt), row.task) for row in held_out]) + complexity_model = NaiveBayes([(featurize(row.prompt), row.complexity) for row in train]) + complexity_model.calibrate([(featurize(row.prompt), row.complexity) for row in held_out]) + return TaskClassifier(task_model, complexity_model, confidence_threshold=confidence_threshold) + + +@lru_cache(maxsize=4) +def load_classifier(path: str = DEFAULT_DATASET) -> TaskClassifier: + """Train once per dataset path; the model is immutable after calibration.""" + + return train_classifier(load_labelled_prompts(path)) diff --git a/src/llm_router/config.py b/src/llm_router/config.py index ba45041..00d49c4 100644 --- a/src/llm_router/config.py +++ b/src/llm_router/config.py @@ -26,6 +26,9 @@ class Settings(BaseSettings): backend_timeout_seconds: float = Field(default=60.0, gt=0) registry_path: str = "config/registry.yaml" routing_policy_version: str = "v1" + # Labelled prompts the task and complexity classifier is trained on at + # start-up; keyword rules decide the task when the file is absent. + task_classifier_path: str = "config/routing/task-classifier-v1.jsonl" redis_url: str = "" # Traces are exported only when a collector endpoint is set. Prompt content # is never recorded unless an operator opts in, and then only for public diff --git a/src/llm_router/load.py b/src/llm_router/load.py new file mode 100644 index 0000000..65620e3 --- /dev/null +++ b/src/llm_router/load.py @@ -0,0 +1,67 @@ +"""Live load observed by the gateway, fed back into routing (section 7.2). + +Section 7.2 lists current queue delay and GPU capacity as routing features. A +catalog can only carry an estimate of queue delay; this tracker replaces the +estimate with what requests actually waited, and remembers how saturated the +engine last reported itself to be. +""" + +import time +from dataclasses import dataclass, field + +from llm_router.engine_stats import EngineStats + +DEFAULT_SMOOTHING = 0.2 +DEFAULT_ENGINE_TTL_SECONDS = 30.0 + + +@dataclass +class LoadTracker: + """Exponentially weighted queue delay per model, plus engine saturation.""" + + smoothing: float = DEFAULT_SMOOTHING + engine_ttl_seconds: float = DEFAULT_ENGINE_TTL_SECONDS + _queue_ms: dict[str, float] = field(default_factory=dict) + _saturation: float | None = None + _saturation_at: float = 0.0 + + def observe_queue(self, model: str, queue_ms: float) -> None: + previous = self._queue_ms.get(model) + self._queue_ms[model] = ( + queue_ms if previous is None else previous + self.smoothing * (queue_ms - previous) + ) + + def queue_ms(self, model: str, *, default: float) -> float: + """Observed queue delay, or the catalog estimate before any observation.""" + + return self._queue_ms.get(model, default) + + def observe_engine(self, stats: EngineStats, *, now: float | None = None) -> None: + """Remember how full the engine is, from whichever signals it published.""" + + signals: list[float] = [] + if stats.kv_cache_usage_ratio is not None: + signals.append(stats.kv_cache_usage_ratio) + if stats.running_requests is not None and stats.waiting_requests is not None: + total = stats.running_requests + stats.waiting_requests + if total > 0: + # The share of admitted work the engine has not started on. + signals.append(stats.waiting_requests / total) + if not signals: + return + self._saturation = min(1.0, max(signals)) + self._saturation_at = time.monotonic() if now is None else now + + def engine_saturation(self, *, now: float | None = None) -> float: + """Last reported saturation, or zero once the reading has gone stale. + + Engine state is only refreshed when metrics are scraped, so an old + reading is discarded rather than allowed to steer routing indefinitely. + """ + + if self._saturation is None: + return 0.0 + current = time.monotonic() if now is None else now + if current - self._saturation_at > self.engine_ttl_seconds: + return 0.0 + return self._saturation diff --git a/src/llm_router/models.py b/src/llm_router/models.py index 8e00377..effa86e 100644 --- a/src/llm_router/models.py +++ b/src/llm_router/models.py @@ -60,6 +60,13 @@ class ModelProfile(BaseModel): quality: float = Field(ge=0.0, le=1.0) estimated_queue_ms: int = Field(default=0, ge=0) cost_weight: float = Field(default=0.0, ge=0.0) + # Measured quality per task from benchmark history; the single + # `quality` figure is only the fallback for a task never measured. + quality_by_task: dict[TaskClass, float] = Field(default_factory=dict) + supports_structured_output: bool = True + + def quality_for(self, task: TaskClass) -> float: + return self.quality_by_task.get(task, self.quality) class RouteDecision(BaseModel): @@ -70,6 +77,11 @@ class RouteDecision(BaseModel): candidate_count: int adapter_id: str | None = None adapter_revision: str | None = None + # How the task was established: declared by the caller, predicted by + # the classifier, or defaulted to general when the classifier abstained. + task_source: str = "declared" + task_confidence: float | None = None + complexity: str | None = None class ChatCompletionChoice(BaseModel): diff --git a/src/llm_router/registry.py b/src/llm_router/registry.py index f027c4b..d7ce1aa 100644 --- a/src/llm_router/registry.py +++ b/src/llm_router/registry.py @@ -76,6 +76,7 @@ class ModelCard(BaseModel): cost_weight: float = Field(default=0.0, ge=0.0) stage: LifecycleStage = LifecycleStage.DEVELOPMENT healthy: bool = True + supports_structured_output: bool = True intended_tasks: str limitations: str evaluation_references: tuple[str, ...] = () @@ -86,8 +87,10 @@ def require_evidence_for_production(self) -> "ModelCard": raise ValueError(f"model {self.id} cannot reach production without evaluation evidence") return self - def to_profile(self) -> ModelProfile: + def to_profile(self, quality_by_task: dict[TaskClass, float] | None = None) -> ModelProfile: return ModelProfile( + quality_by_task=quality_by_task or {}, + supports_structured_output=self.supports_structured_output, id=self.id, revision=self.revision, local=self.local, @@ -138,6 +141,9 @@ class BenchmarkRun(BaseModel): engine_revision: str model_revision: str adapter_revision: str | None = None + # The task the dataset measures, so quality history can be kept per + # task and model rather than as one figure per model. + task: TaskClass | None = None concurrency: int = Field(ge=1) prompt_tokens_p50: int = Field(ge=1) prompt_tokens_p95: int = Field(ge=1) @@ -274,7 +280,26 @@ def servable_models(self) -> tuple[ModelCard, ...]: ) def profiles(self) -> tuple[ModelProfile, ...]: - return tuple(card.to_profile() for card in self.servable_models()) + return tuple( + card.to_profile(self.quality_history(card.revision)) for card in self.servable_models() + ) + + def quality_history(self, model_revision: str) -> dict[TaskClass, float]: + """Mean benchmarked quality per task for one base-model revision. + + Adapter runs are excluded: they measure the adapter, and are already + accounted for through the adapter's own quality delta. + """ + + scores: dict[TaskClass, list[float]] = {} + for run in self.benchmarks: + if ( + run.model_revision == model_revision + and run.adapter_revision is None + and run.task is not None + ): + scores.setdefault(run.task, []).append(run.quality_score) + return {task: sum(values) / len(values) for task, values in scores.items()} def model_card(self, model_id: str) -> ModelCard: for card in self.models: diff --git a/src/llm_router/routing.py b/src/llm_router/routing.py index ce60d0b..d3abc7e 100644 --- a/src/llm_router/routing.py +++ b/src/llm_router/routing.py @@ -1,6 +1,8 @@ from dataclasses import dataclass from typing import TYPE_CHECKING +from llm_router.classifier import Complexity, TaskClassifier +from llm_router.load import LoadTracker from llm_router.models import ( ChatCompletionRequest, ModelProfile, @@ -13,6 +15,20 @@ from llm_router.registry import Registry +# How strongly predicted complexity pulls a request toward higher measured +# quality. Low-complexity work is left to the cheapest capable model; a +# high-complexity request outweighs the specialization and cost terms. +COMPLEXITY_QUALITY_WEIGHT: dict[Complexity, float] = { + Complexity.LOW: 0.0, + Complexity.MEDIUM: 0.25, + Complexity.HIGH: 1.5, +} +COMPLEXITY_QUALITY_BASELINE = 0.8 +# Points a fully saturated engine costs every model it serves, which only +# matters against a candidate served elsewhere. +ENGINE_SATURATION_PENALTY = 20.0 + + class NoEligibleModelError(RuntimeError): """Raised when policy removes every model candidate.""" @@ -75,11 +91,20 @@ class Router: profiles: tuple[ModelProfile, ...] external_fallback_enabled: bool = False registry: "Registry | None" = None + # Without a classifier the keyword rules below decide the task, which + # keeps the router usable where no training data is deployed. + classifier: TaskClassifier | None = None + load: LoadTracker | None = None def classify_task(self, request: ChatCompletionRequest) -> TaskClass: if request.routing.task is not None: return request.routing.task + if self.classifier is not None: + return self.classifier.predict(request.prompt).task + return self._keyword_task(request) + @staticmethod + def _keyword_task(request: ChatCompletionRequest) -> TaskClass: prompt = request.prompt.lower() keywords = ( (TaskClass.EXTRACTION, ("extract", "json schema", "fields from")), @@ -104,7 +129,21 @@ def select( tenant_allows_external: bool = True, privacy_raised_from: PrivacyClass | None = None, ) -> RouteDecision: - task = task if task is not None else self.classify_task(request) + prediction = ( + self.classifier.predict(request.prompt) if self.classifier is not None else None + ) + complexity = prediction.complexity if prediction is not None else None + task_confidence: float | None = None + if request.routing.task is not None: + # A declared task is a fact about the request, not a prediction. + task, task_source = request.routing.task, "declared" + elif task is not None: + task_source = "cached" + elif prediction is not None: + task, task_confidence = prediction.task, prediction.task_confidence + task_source = "classifier" if prediction.trusted else "abstained" + else: + task, task_source = self._keyword_task(request), "keyword" estimated_tokens = max(1, len(request.prompt) // 4) + request.max_tokens # The request arrives carrying its effective privacy class: the gateway # raises it to the tenant floor before anything reads it, including the @@ -118,7 +157,11 @@ def select( if profile.healthy and task in profile.supported_tasks and estimated_tokens <= profile.context_limit - and profile.quality >= floor + and profile.quality_for(task) >= floor + # Capability restrictions stay deterministic: a model that + # cannot produce structured output is never a candidate for a + # request that requires it, whatever it would have scored. + and (not request.routing.structured or profile.supports_structured_output) and self._privacy_allows(profile, effective_privacy) and self._external_allows(profile, request, tenant_allows_external) # A tenant entitlement is a hard restriction, like privacy: it is @@ -131,21 +174,41 @@ def select( if not candidates: raise NoEligibleModelError( - "no healthy model satisfies capability, context, quality, " + "no healthy model satisfies capability, structured output, context, quality, " "privacy, tenant entitlement, and fallback policy" ) + saturation = self.load.engine_saturation() if self.load is not None else 0.0 + def score(profile: ModelProfile) -> float: - latency_penalty = profile.estimated_queue_ms / 100 + # Observed queue delay replaces the catalog estimate as soon as + # any request for the model has actually waited. + queue_ms = ( + self.load.queue_ms(profile.id, default=profile.estimated_queue_ms) + if self.load is not None + else profile.estimated_queue_ms + ) + latency_penalty = queue_ms / 100 if request.routing.latency_tier == "interactive": latency_penalty *= 2 elif request.routing.latency_tier == "batch": latency_penalty *= 0.5 specialization_bonus = 10 if len(profile.supported_tasks) <= 2 else 0 + quality = profile.quality_for(task) + complexity_bonus = ( + (quality - COMPLEXITY_QUALITY_BASELINE) + * 100 + * COMPLEXITY_QUALITY_WEIGHT[complexity] + if complexity is not None + else 0.0 + ) + capacity_penalty = saturation * ENGINE_SATURATION_PENALTY if profile.local else 0.0 return ( - (profile.quality * 100) + (quality * 100) + specialization_bonus + + complexity_bonus - latency_penalty + - capacity_penalty - (profile.cost_weight * 10) ) @@ -165,6 +228,17 @@ def score(profile: ModelProfile) -> float: f"task={task.value}, privacy={effective_privacy.value}, " f"latency_tier={request.routing.latency_tier}" ) + if task_source == "classifier": + reason += f"; task predicted with confidence {task_confidence:.2f}" + elif task_source == "abstained": + reason += ( + f"; classifier abstained at confidence {task_confidence:.2f}, " + "so the general task was used" + ) + if complexity is not None: + reason += f"; complexity={complexity.value}" + if saturation > 0: + reason += f"; engine saturation {saturation:.2f}" if privacy_raised_from is not None: # Raising the class is a policy decision and is attributed rather # than applied silently. @@ -184,6 +258,9 @@ def score(profile: ModelProfile) -> float: candidate_count=len(candidates), adapter_id=None if adapter is None else adapter.id, adapter_revision=None if adapter is None else adapter.adapter_revision, + task_source=task_source, + task_confidence=task_confidence, + complexity=None if complexity is None else complexity.value, ) @staticmethod diff --git a/src/llm_router/tracing.py b/src/llm_router/tracing.py index 9a1f178..ae4b356 100644 --- a/src/llm_router/tracing.py +++ b/src/llm_router/tracing.py @@ -103,6 +103,9 @@ def set_route(self, decision: RouteDecision) -> None: "router.adapter": decision.adapter_id, "router.adapter.revision": decision.adapter_revision, "router.task": decision.task.value, + "router.task.source": decision.task_source, + "router.task.confidence": decision.task_confidence, + "router.complexity": decision.complexity, "router.route.reason": decision.reason, "router.route.score": decision.score, "router.route.candidates": decision.candidate_count, diff --git a/tests/integration/test_api.py b/tests/integration/test_api.py index 0c9538f..92d53ec 100644 --- a/tests/integration/test_api.py +++ b/tests/integration/test_api.py @@ -142,3 +142,52 @@ def test_metrics_endpoint_counts_policy_rejections() -> None: metrics = client.get("/metrics").text assert 'router_rejections_total{type="no_eligible_model"} 1.0' in metrics + + +def test_response_attributes_how_the_task_was_established() -> None: + with build_client() as client: + predicted = client.post( + "/v1/chat/completions", + headers={"Authorization": "Bearer integration-key"}, + json={ + "model": "auto", + "messages": [{"role": "user", "content": "Which label fits this complaint?"}], + }, + ).json()["routing"] + declared = client.post( + "/v1/chat/completions", + headers={"Authorization": "Bearer integration-key"}, + json={ + "model": "auto", + "messages": [{"role": "user", "content": "Which label fits this complaint?"}], + "routing": {"task": "summarization", "privacy": "restricted"}, + }, + ).json()["routing"] + + assert predicted["task"] == "classification" + assert predicted["task_source"] == "classifier" + assert predicted["task_confidence"] >= 0.5 + assert predicted["complexity"] == "low" + assert declared["task"] == "summarization" + assert declared["task_source"] == "declared" + assert declared["task_confidence"] is None + + +def test_observed_queue_delay_feeds_back_into_the_gateway_router() -> None: + with build_client() as client: + for _ in range(2): + response = client.post( + "/v1/chat/completions", + headers={"Authorization": "Bearer integration-key"}, + json={ + "model": "auto", + "messages": [{"role": "user", "content": "Classify this ticket"}], + "routing": {"privacy": "restricted"}, + }, + ) + assert response.status_code == 200 + + # The mock engine admits instantly, so the observed delay stays near zero + # and the specialist keeps winning; the point is that routing still works + # once the estimate has been replaced by an observation. + assert response.json()["model"] == "small-specialist" diff --git a/tests/unit/test_classifier.py b/tests/unit/test_classifier.py new file mode 100644 index 0000000..094e336 --- /dev/null +++ b/tests/unit/test_classifier.py @@ -0,0 +1,142 @@ +from pathlib import Path + +import pytest + +from llm_router.classifier import ( + DEFAULT_DATASET, + TEMPERATURES, + ClassifierError, + Complexity, + LabelledPrompt, + NaiveBayes, + expected_calibration_error, + featurize, + load_classifier, + load_labelled_prompts, + train_classifier, +) +from llm_router.models import TaskClass + +# Prompts the classifier never trained or calibrated on. +EVALUATION = load_labelled_prompts("benchmarks/datasets/routing-tasks-v1.jsonl") +CLASSIFIER = load_classifier(DEFAULT_DATASET) + + +def test_features_include_bigrams_and_a_length_bucket() -> None: + features = featurize("Extract the invoice total") + + assert {"extract", "invoice", "extract_the", "invoice_total", "__short__"} <= set(features) + + +def test_long_prompts_keep_the_instruction_at_either_end() -> None: + context = " ".join(f"filler{index}" for index in range(400)) + + features = set(featurize(f"Summarize the following. {context} Answer in one line.")) + + # The instruction survives at both ends; the middle of the context does not. + assert {"summarize", "answer", "__long__"} <= features + assert "filler200" not in features + + +def test_committed_dataset_has_both_splits_and_every_task() -> None: + rows = load_labelled_prompts(DEFAULT_DATASET) + + assert {row.split for row in rows} == {"train", "calibration"} + assert {row.task for row in rows} == set(TaskClass) + assert {row.complexity for row in rows} == set(Complexity) + + +def test_evaluation_prompts_are_disjoint_from_the_training_data() -> None: + seen = {row.prompt for row in load_labelled_prompts(DEFAULT_DATASET)} + + assert not seen & {row.prompt for row in EVALUATION} + + +def test_task_accuracy_on_prompts_never_seen_in_training() -> None: + correct = sum(CLASSIFIER.predict(row.prompt).task == row.task for row in EVALUATION) + + assert correct / len(EVALUATION) >= 0.9 + + +def test_complexity_accuracy_on_prompts_never_seen_in_training() -> None: + correct = sum(CLASSIFIER.predict(row.prompt).complexity == row.complexity for row in EVALUATION) + + assert correct / len(EVALUATION) >= 0.8 + + +def test_stated_confidence_matches_observed_accuracy() -> None: + pairs = [(featurize(row.prompt), row.task) for row in EVALUATION] + + assert expected_calibration_error(CLASSIFIER.task_model, pairs) <= 0.1 + + +def test_calibration_chooses_a_temperature_from_the_grid() -> None: + assert CLASSIFIER.task_model.temperature in TEMPERATURES + assert CLASSIFIER.complexity_model.temperature in TEMPERATURES + + +def test_an_unrecognisable_prompt_abstains_to_the_general_task() -> None: + prediction = CLASSIFIER.predict("zxqv flurble quonk") + + assert prediction.trusted is False + assert prediction.task is TaskClass.GENERAL + assert prediction.task_confidence < CLASSIFIER.confidence_threshold + + +def test_a_recognisable_prompt_is_trusted_with_a_probability() -> None: + prediction = CLASSIFIER.predict("Classify this ticket as billing or technical") + + assert prediction.trusted is True + assert prediction.task is TaskClass.CLASSIFICATION + assert 0.5 <= prediction.task_confidence <= 1.0 + assert prediction.complexity is Complexity.LOW + + +def test_posterior_is_a_probability_distribution() -> None: + prediction = CLASSIFIER.task_model.predict(featurize("Summarize the report")) + + assert sum(prediction.distribution.values()) == pytest.approx(1.0) + assert prediction.confidence == max(prediction.distribution.values()) + + +def test_a_higher_temperature_lowers_confidence() -> None: + examples = [ + (featurize("extract the fields"), "a"), + (featurize("extract the totals"), "a"), + (featurize("summarize the notes"), "b"), + (featurize("summarize the report"), "b"), + ] + model = NaiveBayes(examples) + features = featurize("extract the report") + + model.temperature = 1.0 + sharp = model.predict(features).confidence + model.temperature = 8.0 + soft = model.predict(features).confidence + + assert soft < sharp + + +def test_training_needs_data_and_both_splits() -> None: + row = LabelledPrompt("Extract it", TaskClass.EXTRACTION, Complexity.LOW, "train") + + with pytest.raises(ClassifierError, match="empty dataset"): + NaiveBayes([]) + with pytest.raises(ClassifierError, match="train and a calibration split"): + train_classifier([row]) + with pytest.raises(ClassifierError, match="held-out split"): + NaiveBayes([(featurize("x"), "a")]).calibrate([]) + with pytest.raises(ClassifierError, match="empty sample"): + expected_calibration_error(NaiveBayes([(featurize("x"), "a")]), []) + + +def test_loader_reports_unusable_and_empty_files(tmp_path: Path) -> None: + unusable = tmp_path / "bad.jsonl" + unusable.write_text('{"prompt": "x", "task": "not-a-task", "complexity": "low"}\n', "utf-8") + empty = tmp_path / "empty.jsonl" + empty.write_text("\n\n", "utf-8") + + with pytest.raises(ClassifierError, match="not a usable example"): + load_labelled_prompts(unusable) + with pytest.raises(ClassifierError, match="no examples"): + load_labelled_prompts(empty) diff --git a/tests/unit/test_routing_features.py b/tests/unit/test_routing_features.py new file mode 100644 index 0000000..2cda0d3 --- /dev/null +++ b/tests/unit/test_routing_features.py @@ -0,0 +1,223 @@ +import pytest + +from llm_router.classifier import DEFAULT_DATASET, load_classifier +from llm_router.engine_stats import EngineStats +from llm_router.load import LoadTracker +from llm_router.models import ChatCompletionRequest, ModelProfile, TaskClass +from llm_router.registry import load_registry +from llm_router.routing import NoEligibleModelError, Router, default_model_profiles + +CLASSIFIER = load_classifier(DEFAULT_DATASET) + + +def request_for(prompt: str, **routing: object) -> ChatCompletionRequest: + return ChatCompletionRequest.model_validate( + {"model": "auto", "messages": [{"role": "user", "content": prompt}], "routing": routing} + ) + + +def profile(model_id: str, **overrides: object) -> ModelProfile: + values: dict[str, object] = { + "id": model_id, + "revision": f"{model_id}@rev", + "local": True, + "context_limit": 8192, + "supported_tasks": frozenset(TaskClass), + "quality": 0.9, + } + values.update(overrides) + return ModelProfile.model_validate(values) + + +def test_the_classifier_replaces_keyword_rules_and_reports_its_confidence() -> None: + router = Router(default_model_profiles(), classifier=CLASSIFIER) + + decision = router.select(request_for("Which category does this complaint belong to?")) + + assert decision.task is TaskClass.CLASSIFICATION + assert decision.task_source == "classifier" + assert decision.task_confidence is not None and decision.task_confidence >= 0.5 + assert "task predicted with confidence" in decision.reason + + +def test_keyword_rules_still_decide_without_a_classifier() -> None: + decision = Router(default_model_profiles()).select(request_for("Classify this ticket")) + + assert decision.task is TaskClass.CLASSIFICATION + assert decision.task_source == "keyword" + assert decision.task_confidence is None + assert decision.complexity is None + + +def test_a_declared_task_is_never_overridden_by_a_prediction() -> None: + router = Router(default_model_profiles(), classifier=CLASSIFIER) + + decision = router.select(request_for("Summarize the report", task="extraction")) + + assert decision.task is TaskClass.EXTRACTION + assert decision.task_source == "declared" + assert decision.task_confidence is None + + +def test_an_abstention_routes_as_general_and_says_so() -> None: + router = Router(default_model_profiles(), classifier=CLASSIFIER) + + decision = router.select(request_for("zxqv flurble quonk")) + + assert decision.task is TaskClass.GENERAL + assert decision.task_source == "abstained" + assert "classifier abstained" in decision.reason + + +def test_low_complexity_work_stays_on_the_cheapest_capable_model() -> None: + router = Router(default_model_profiles(), classifier=CLASSIFIER) + + decision = router.select(request_for("Classify this ticket as billing or technical")) + + assert decision.complexity == "low" + assert decision.profile.id == "small-specialist" + + +def test_high_complexity_pulls_the_same_task_to_a_stronger_model() -> None: + router = Router(default_model_profiles(), classifier=CLASSIFIER) + # Not in the training data: the complexity is predicted, not memorised. + prompt = ( + "Label every one of these thirty contract provisions with a risk category, justify " + "each label, and flag the provisions that plausibly belong to two categories." + ) + + decision = router.select(request_for(prompt, task="classification")) + + assert decision.complexity == "high" + assert decision.profile.id == "high-capability" + assert "complexity=high" in decision.reason + + +def test_a_structured_request_never_reaches_a_model_that_cannot_produce_it() -> None: + profiles = ( + profile("cheap", quality=0.99, supports_structured_output=False), + profile("capable", quality=0.80), + ) + router = Router(profiles) + + plain = router.select(request_for("Extract the fields")) + structured = router.select(request_for("Extract the fields", structured=True)) + + # The incapable model scores far higher, and still is not a candidate. + assert plain.profile.id == "cheap" + assert structured.profile.id == "capable" + + +def test_a_structured_request_with_no_capable_model_is_refused() -> None: + router = Router((profile("cheap", supports_structured_output=False),)) + + with pytest.raises(NoEligibleModelError, match="structured output"): + router.select(request_for("Extract the fields", structured=True)) + + +def test_measured_quality_for_the_task_outranks_the_headline_figure() -> None: + profiles = ( + profile("headline", quality=0.95, quality_by_task={TaskClass.EXTRACTION: 0.70}), + profile("measured", quality=0.85, quality_by_task={TaskClass.EXTRACTION: 0.93}), + ) + router = Router(profiles) + + extraction = router.select(request_for("Extract the fields")) + summary = router.select(request_for("Summarize the report")) + + assert extraction.profile.id == "measured" + # A task with no history falls back to the headline figure. + assert summary.profile.id == "headline" + + +def test_the_quality_floor_is_checked_against_quality_for_the_task() -> None: + router = Router((profile("m", quality=0.95, quality_by_task={TaskClass.EXTRACTION: 0.70}),)) + + with pytest.raises(NoEligibleModelError): + router.select(request_for("Extract the fields", quality_floor=0.9)) + + +def test_observed_queue_delay_replaces_the_catalog_estimate() -> None: + profiles = ( + profile("near", quality=0.90, estimated_queue_ms=10), + profile("far", quality=0.90, estimated_queue_ms=40), + ) + load = LoadTracker() + router = Router(profiles, load=load) + + assert router.select(request_for("Summarize the report")).profile.id == "near" + + # The model the catalog called fast has in fact been queueing for seconds. + load.observe_queue("near", 4000) + + assert router.select(request_for("Summarize the report")).profile.id == "far" + + +def test_queue_delay_is_smoothed_rather_than_taken_from_one_request() -> None: + load = LoadTracker(smoothing=0.5) + + load.observe_queue("m", 100) + load.observe_queue("m", 300) + + assert load.queue_ms("m", default=0) == pytest.approx(200) + assert load.queue_ms("unseen", default=35) == 35 + + +def test_engine_saturation_reads_the_worst_signal_and_expires() -> None: + load = LoadTracker(engine_ttl_seconds=30) + + load.observe_engine( + EngineStats(running_requests=6, waiting_requests=2, kv_cache_usage_ratio=0.9), now=100.0 + ) + + assert load.engine_saturation(now=110.0) == pytest.approx(0.9) + # Engine state is only refreshed on scrape, so a stale reading is dropped. + assert load.engine_saturation(now=200.0) == 0.0 + + +def test_an_engine_that_published_nothing_usable_leaves_saturation_unset() -> None: + load = LoadTracker() + + load.observe_engine(EngineStats(running_requests=0, waiting_requests=0), now=1.0) + load.observe_engine(EngineStats(preemptions_total=3), now=1.0) + + assert load.engine_saturation(now=2.0) == 0.0 + + +def test_a_saturated_engine_tips_an_eligible_request_to_the_external_model() -> None: + load = LoadTracker() + router = Router(default_model_profiles(), external_fallback_enabled=True, load=load) + body = request_for("Analyze deeply", privacy="public", allow_external_fallback=True) + + assert router.select(body).profile.id == "high-capability" + + load.observe_engine(EngineStats(kv_cache_usage_ratio=1.0)) + saturated = router.select(body) + + assert saturated.profile.id == "approved-external-fallback" + assert "engine saturation 1.00" in saturated.reason + + +def test_saturation_never_overrides_privacy() -> None: + load = LoadTracker() + load.observe_engine(EngineStats(kv_cache_usage_ratio=1.0)) + router = Router(default_model_profiles(), external_fallback_enabled=True, load=load) + + decision = router.select( + request_for("Analyze deeply", privacy="private", allow_external_fallback=True) + ) + + assert decision.profile.local is True + + +def test_quality_history_is_built_from_benchmarks_per_task() -> None: + registry = load_registry("config/registry.yaml") + small = registry.model_card("small-specialist") + + history = registry.quality_history(small.revision) + profiles = {item.id: item for item in registry.profiles()} + + assert history == {TaskClass.EXTRACTION: pytest.approx(0.82)} + assert profiles["small-specialist"].quality_for(TaskClass.EXTRACTION) == pytest.approx(0.82) + # No benchmark covers this task, so the card's headline figure applies. + assert profiles["small-specialist"].quality_for(TaskClass.CLASSIFICATION) == small.quality