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