feat(ui): add adaptive routing settings to Auto-Router v2 (#33146)

This commit is contained in:
Krrish Dholakia 2026-07-13 21:29:58 -07:00 • committed by GitHub
parent 6a213de9f4
commit b200d664ee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 733 additions and 239 deletions

View file

@ -0,0 +1,134 @@
import { Card, InputNumber, Radio, Slider, Space, Switch, Typography } from "antd";
import React from "react";
import {
AdaptiveEligible,
ComplexityRouterConfigValue,
DEFAULT_ADAPTIVE_WEIGHTS,
DEFAULT_TIER_DISTANCE_PENALTY,
} from "./ComplexityRouterConfig";
const { Text } = Typography;
interface AdaptiveRoutingConfigProps {
value: ComplexityRouterConfigValue;
onChange: (value: ComplexityRouterConfigValue) => void;
}
const AdaptiveRoutingConfig: React.FC<AdaptiveRoutingConfigProps> = ({ value, onChange }) => {
const adaptiveWeights = value.adaptive_weights ?? DEFAULT_ADAPTIVE_WEIGHTS;
const adaptiveEligible = value.adaptive_eligible ?? "all";
const tierDistancePenalty = value.tier_distance_penalty ?? DEFAULT_TIER_DISTANCE_PENALTY;
const handleAdaptiveToggle = (adaptive: boolean) => {
const nextValue: ComplexityRouterConfigValue = {
...value,
adaptive,
adaptive_weights: adaptiveWeights,
adaptive_eligible: adaptiveEligible,
tier_distance_penalty: tierDistancePenalty,
};
onChange(nextValue);
};
const handleQualityWeightChange = (qualityPercent: number) => {
const quality = qualityPercent / 100;
onChange({ ...value, adaptive_weights: { quality, cost: Math.round((1 - quality) * 100) / 100 } });
};
const handleAdaptiveEligibleChange = (eligible: AdaptiveEligible) => {
onChange({ ...value, adaptive_eligible: eligible });
};
const handleTierDistancePenaltyChange = (penalty: number | null) => {
onChange({ ...value, tier_distance_penalty: penalty ?? DEFAULT_TIER_DISTANCE_PENALTY });
};
return (
<>
<div className="flex items-center gap-2 mb-2">
<Switch checked={value.adaptive ?? false} onChange={handleAdaptiveToggle} />
<Text strong>Enable adaptive bandit selection</Text>
</div>
<Text type="secondary" style={{ display: "block", fontSize: 12 }}>
When disabled, each request always uses the model assigned to its classified tier.
</Text>
<Card className="bg-gray-50 mt-4">
<Text strong style={{ display: "block", marginBottom: 8 }}>
How Adaptive Routing Works
</Text>
<Text type="secondary" style={{ fontSize: 13 }}>
It learns from how each conversation actually goes: does the user have to rephrase or correct the model, does
it get stuck repeating itself, does it run out of tool calls, does the user seem satisfied. Combined with
cost, this live feedback shifts future routing toward the models that are actually working well, and improves
as more conversations come in. Until there&apos;s enough feedback, it defaults to the classified tier&apos;s
model.
</Text>
</Card>
{value.adaptive && (
<div className="mt-4 space-y-4">
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Quality vs. Cost ({Math.round(adaptiveWeights.quality * 100)}% quality /{" "}
{Math.round(adaptiveWeights.cost * 100)}% cost)
</Text>
<Slider
min={0}
max={100}
value={Math.round(adaptiveWeights.quality * 100)}
onChange={handleQualityWeightChange}
tooltip={{ formatter: (v) => `${v}% quality / ${100 - (v ?? 0)}% cost` }}
/>
<Text type="secondary" style={{ fontSize: 12 }}>
Higher quality weight favors more capable (pricier) models; higher cost weight favors cheaper models when
the bandit has feedback to act on. Recommended: 30% quality / 70% cost split.
</Text>
</div>
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Eligible Model Pool
</Text>
<Radio.Group
value={adaptiveEligible}
onChange={(e) => handleAdaptiveEligibleChange(e.target.value)}
className="w-full"
>
<Space direction="vertical" className="w-full">
<Radio value="all">
<Text strong>All tiers (soft floor)</Text>{" "}
<Text type="secondary">— router can pick across tiers, depending on the best fit for the prompt</Text>
</Radio>
<Radio value="classified_tier">
<Text strong>Classified tier only</Text>{" "}
<Text type="secondary">— router can only pick models within tier</Text>
</Radio>
</Space>
</Radio.Group>
</div>
{adaptiveEligible === "all" && (
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Tier Distance Penalty
</Text>
<InputNumber
value={tierDistancePenalty}
onChange={handleTierDistancePenaltyChange}
min={0}
step={0.1}
style={{ width: "100%" }}
/>
<Text type="secondary" style={{ fontSize: 12 }}>
Score penalty applied per tier-step away from the classified tier.
</Text>
</div>
)}
</div>
)}
</>
);
};
export default AdaptiveRoutingConfig;

View file

@ -0,0 +1,171 @@
import { InfoCircleOutlined } from "@ant-design/icons";
import { Select as AntdSelect, Card, InputNumber, Radio, Space, Tooltip, Typography } from "antd";
import React from "react";
import { ClassifierType, ComplexityRouterConfigValue, DEFAULT_CLASSIFIER_TIMEOUT_MS } from "./ComplexityRouterConfig";
const { Text } = Typography;
interface ClassificationMethodConfigProps {
value: ComplexityRouterConfigValue;
onChange: (value: ComplexityRouterConfigValue) => void;
modelOptions: { value: string; label: string }[];
customTechnicalKeywords?: string[];
onCustomTechnicalKeywordsChange?: (keywords: string[]) => void;
showValidationErrors?: boolean;
}
const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
value,
onChange,
modelOptions,
customTechnicalKeywords,
onCustomTechnicalKeywordsChange,
showValidationErrors = false,
}) => {
const classifierModelMissing =
showValidationErrors && value.classifier_type === "llm" && !value.classifier_llm_config?.model;
const handleClassifierTypeChange = (classifierType: ClassifierType) => {
onChange({
...value,
classifier_type: classifierType,
classifier_llm_config:
classifierType === "llm"
? value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS }
: undefined,
});
};
const handleClassifierModelChange = (model: string) => {
onChange({
...value,
classifier_llm_config: {
model,
timeout_ms: value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
});
};
const handleClassifierTimeoutChange = (timeoutMs: number | null) => {
onChange({
...value,
classifier_llm_config: {
model: value.classifier_llm_config?.model ?? "",
timeout_ms: timeoutMs ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
});
};
return (
<>
<Radio.Group
value={value.classifier_type}
onChange={(e) => handleClassifierTypeChange(e.target.value)}
className="w-full"
>
<Space direction="vertical" className="w-full">
<Radio value="heuristic">
<Text strong>Heuristic</Text>{" "}
<Text type="secondary">(default) — rule-based scoring, no API calls, &lt;1ms latency</Text>
</Radio>
<Radio value="llm">
<Text strong>LLM Classifier</Text>{" "}
<Text type="secondary">— use a model to decide the tier (e.g. a small/fast model)</Text>
</Radio>
</Space>
</Radio.Group>
{value.classifier_type === "llm" && (
<div className="mt-4 space-y-3">
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Classifier Model
</Text>
<AntdSelect
value={value.classifier_llm_config?.model || undefined}
onChange={handleClassifierModelChange}
placeholder="Select the model that will classify request complexity"
showSearch
style={{ width: "100%" }}
options={modelOptions}
status={classifierModelMissing ? "error" : undefined}
/>
{classifierModelMissing && (
<Text type="danger" style={{ fontSize: 12 }}>
A classifier model is required
</Text>
)}
</div>
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Timeout (ms)
</Text>
<InputNumber
value={value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS}
onChange={handleClassifierTimeoutChange}
min={1}
style={{ width: "100%" }}
/>
<Text type="secondary" style={{ fontSize: 12 }}>
Falls back to the heuristic scorer if the classifier call errors, times out, or returns an unparseable
response.
</Text>
</div>
</div>
)}
{value.classifier_type === "heuristic" && (
<div className="mt-4">
<div className="flex items-center gap-2 mb-1">
<Text strong>Custom Technical Keywords</Text>
<Tooltip title="Domain-specific terms appended to the built-in technical keyword list. Prompts containing these terms score higher on the technical dimension and route to more capable models.">
<InfoCircleOutlined className="text-gray-400" />
</Tooltip>
</div>
<Text type="secondary" style={{ display: "block", marginBottom: 8, fontSize: 12 }}>
Optional: Add terms to the built-in list to improve classification accuracy on the technical dimension.
(e.g., udp, kafka, terraform).
</Text>
<AntdSelect
mode="tags"
value={customTechnicalKeywords ?? []}
onChange={(keywords: string[]) => onCustomTechnicalKeywordsChange?.(keywords)}
placeholder="Type a keyword and press Enter, or paste a comma-separated list"
tokenSeparators={[","]}
open={false}
suffixIcon={null}
style={{ width: "100%" }}
allowClear
/>
</div>
)}
<Card className="bg-gray-50 mt-4">
<Text strong style={{ display: "block", marginBottom: 8 }}>
How Classification Works
</Text>
<Text type="secondary" style={{ fontSize: 13 }}>
The router scores each request across 7 dimensions: token count, code presence, reasoning markers, technical
terms, simple indicators, multi-step patterns, and question complexity. The weighted score determines the
tier:
</Text>
<ul style={{ marginTop: 8, marginBottom: 0, paddingLeft: 20, fontSize: 13, color: "rgba(0, 0, 0, 0.45)" }}>
<li>
<strong>SIMPLE</strong>: Score &lt; 0.15
</li>
<li>
<strong>MEDIUM</strong>: Score 0.15 - 0.35
</li>
<li>
<strong>COMPLEX</strong>: Score 0.35 - 0.60
</li>
<li>
<strong>REASONING</strong>: Score &gt; 0.60 (or 2+ reasoning markers)
</li>
</ul>
</Card>
</>
);
};
export default ClassificationMethodConfig;

View file

@ -12,10 +12,10 @@ const mockModelInfo = [
const defaultValue: ComplexityRouterConfigValue = {
tiers: {
SIMPLE: "gpt-3.5-turbo",
MEDIUM: "gpt-3.5-turbo",
COMPLEX: "gpt-4",
REASONING: "claude-3-opus",
SIMPLE: ["gpt-3.5-turbo"],
MEDIUM: ["gpt-3.5-turbo"],
COMPLEX: ["gpt-4"],
REASONING: ["claude-3-opus"],
},
classifier_type: "heuristic",
};
@ -58,11 +58,13 @@ describe("ComplexityRouterConfig", () => {
it("should display the how classification works section", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
expect(screen.getByText("How Classification Works")).toBeInTheDocument();
});
it("should show score thresholds in the classification section", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
expect(screen.getByText(/Score < 0.15/)).toBeInTheDocument();
expect(screen.getByText(/Score 0.15 - 0.35/)).toBeInTheDocument();
expect(screen.getByText(/Score 0.35 - 0.60/)).toBeInTheDocument();
@ -107,6 +109,7 @@ describe("ComplexityRouterConfig", () => {
it("should render the custom technical keywords field", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
expect(screen.getByText("Custom Technical Keywords")).toBeInTheDocument();
});
@ -118,6 +121,7 @@ describe("ComplexityRouterConfig", () => {
onCustomTechnicalKeywordsChange={vi.fn()}
/>,
);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
expect(screen.getByText("udp")).toBeInTheDocument();
expect(screen.getByText("kafka")).toBeInTheDocument();
});
@ -132,14 +136,16 @@ describe("ComplexityRouterConfig", () => {
onCustomTechnicalKeywordsChange={onCustomTechnicalKeywordsChange}
/>,
);
const keywordsCard = screen.getByText("Custom Technical Keywords").closest(".ant-card") as HTMLElement;
const input = within(keywordsCard).getByRole("combobox");
fireEvent.click(screen.getByText("Advanced: Classification Method"));
const keywordsSection = screen.getByText("Custom Technical Keywords").closest("div")?.parentElement as HTMLElement;
const input = within(keywordsSection).getByRole("combobox");
await user.type(input, "udp,");
expect(onCustomTechnicalKeywordsChange).toHaveBeenCalledWith(["udp"]);
});
it("should render an empty state when no keyword tier rules exist", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
expect(screen.getByText("Keyword Tier Overrides")).toBeInTheDocument();
expect(screen.getByText("No keyword tier overrides configured")).toBeInTheDocument();
});
@ -158,6 +164,7 @@ describe("ComplexityRouterConfig", () => {
const user = userEvent.setup();
const onKeywordTierRulesChange = vi.fn();
renderWithProviders(<ComplexityRouterConfig {...baseProps} onKeywordTierRulesChange={onKeywordTierRulesChange} />);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
await user.click(screen.getByRole("button", { name: /add keyword rule/i }));
expect(onKeywordTierRulesChange).toHaveBeenCalledTimes(1);
const newRules = onKeywordTierRulesChange.mock.calls[0][0];
@ -175,6 +182,7 @@ describe("ComplexityRouterConfig", () => {
onKeywordTierRulesChange={onKeywordTierRulesChange}
/>,
);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
expect(screen.getByText("invoice")).toBeInTheDocument();
expect(screen.getByText("refund")).toBeInTheDocument();
@ -184,6 +192,7 @@ describe("ComplexityRouterConfig", () => {
it("should not show embedding model or match score fields when semantic matching is disabled", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} semanticMatchingEnabled={false} />);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
expect(screen.getByText("Semantic keyword matching")).toBeInTheDocument();
expect(screen.queryByText("Embedding model")).not.toBeInTheDocument();
expect(screen.queryByText("Minimum match score")).not.toBeInTheDocument();
@ -191,6 +200,7 @@ describe("ComplexityRouterConfig", () => {
it("should show embedding model and match score fields when semantic matching is enabled", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} semanticMatchingEnabled={true} />);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
expect(screen.getByText("Embedding model")).toBeInTheDocument();
expect(screen.getByText("Minimum match score")).toBeInTheDocument();
});
@ -205,6 +215,7 @@ describe("ComplexityRouterConfig", () => {
onSemanticMatchingEnabledChange={onSemanticMatchingEnabledChange}
/>,
);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
await user.click(screen.getByRole("switch"));
expect(onSemanticMatchingEnabledChange).toHaveBeenCalledWith(true, expect.anything());
});
@ -252,7 +263,7 @@ describe("ComplexityRouterConfig", () => {
renderWithProviders(
<ComplexityRouterConfig
{...baseProps}
value={{ ...defaultValue, tiers: { ...defaultValue.tiers, REASONING: "" } }}
value={{ ...defaultValue, tiers: { ...defaultValue.tiers, REASONING: [] } }}
showValidationErrors={true}
/>,
);

View file

@ -1,19 +1,22 @@
import { InfoCircleOutlined } from "@ant-design/icons";
import { Select as AntdSelect, Card, Collapse, Divider, InputNumber, Radio, Space, Tooltip, Typography } from "antd";
import { Select as AntdSelect, Card, Collapse, Divider, Space, Tooltip, Typography } from "antd";
import React from "react";
import { ModelGroup } from "@/components/llm_calls/fetch_models";
import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig";
import ClassificationMethodConfig from "./ClassificationMethodConfig";
import KeywordTierRules, { KeywordTierRule } from "./KeywordTierRules";
import SemanticKeywordMatching from "./SemanticKeywordMatching";
const { Text } = Typography;
export const DEFAULT_CLASSIFIER_TIMEOUT_MS = 3000;
export const DEFAULT_TIER_DISTANCE_PENALTY = 0.5;
export interface ComplexityTiers {
SIMPLE: string;
MEDIUM: string;
COMPLEX: string;
REASONING: string;
SIMPLE: string[];
MEDIUM: string[];
COMPLEX: string[];
REASONING: string[];
}
export interface ClassifierLLMConfig {
@ -23,10 +26,23 @@ export interface ClassifierLLMConfig {
export type ClassifierType = "heuristic" | "llm";
export interface AdaptiveRouterWeights {
quality: number;
cost: number;
}
export const DEFAULT_ADAPTIVE_WEIGHTS: AdaptiveRouterWeights = { quality: 0.3, cost: 0.7 };
export type AdaptiveEligible = "all" | "classified_tier";
export interface ComplexityRouterConfigValue {
tiers: ComplexityTiers;
classifier_type: ClassifierType;
classifier_llm_config?: ClassifierLLMConfig;
adaptive?: boolean;
adaptive_weights?: AdaptiveRouterWeights;
tier_distance_penalty?: number;
adaptive_eligible?: AdaptiveEligible;
}
interface ComplexityRouterConfigProps {
@ -95,44 +111,10 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
label: model.model_group,
}));
const classifierModelMissing =
showValidationErrors && value.classifier_type === "llm" && !value.classifier_llm_config?.model;
const handleTierChange = (tier: keyof ComplexityTiers, model: string) => {
const handleTierChange = (tier: keyof ComplexityTiers, models: string[]) => {
onChange({
...value,
tiers: { ...value.tiers, [tier]: model },
});
};
const handleClassifierTypeChange = (classifierType: ClassifierType) => {
onChange({
...value,
classifier_type: classifierType,
classifier_llm_config:
classifierType === "llm"
? (value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS })
: undefined,
});
};
const handleClassifierModelChange = (model: string) => {
onChange({
...value,
classifier_llm_config: {
model,
timeout_ms: value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
});
};
const handleClassifierTimeoutChange = (timeoutMs: number | null) => {
onChange({
...value,
classifier_llm_config: {
model: value.classifier_llm_config?.model ?? "",
timeout_ms: timeoutMs ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
tiers: { ...value.tiers, [tier]: models },
});
};
@ -142,20 +124,20 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
<Typography.Title level={4} style={{ margin: 0 }}>
Complexity Tier Configuration
</Typography.Title>
<Tooltip title="Map each complexity tier to a model. Simple queries use cheaper/faster models, complex queries use more capable models.">
<Tooltip title="Map each complexity tier to one or more models. Simple queries use cheaper/faster models, complex queries use more capable models.">
<InfoCircleOutlined className="text-gray-400" />
</Tooltip>
</Space>
<Text type="secondary" style={{ display: "block", marginBottom: 24 }}>
The complexity router automatically classifies requests by complexity using rule-based scoring (no API calls,
&lt;1ms latency). Configure which model handles each tier.
&lt;1ms latency). Configure which model(s) handle each tier.
</Text>
<Card>
{(Object.keys(TIER_DESCRIPTIONS) as Array<keyof ComplexityTiers>).map((tier, index) => {
const tierInfo = TIER_DESCRIPTIONS[tier];
const tierMissing = showValidationErrors && !value.tiers[tier];
const tierMissing = showValidationErrors && value.tiers[tier].length === 0;
return (
<div key={tier}>
{index > 0 && <Divider style={{ margin: "16px 0" }} />}
@ -172,14 +154,21 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
Examples: {tierInfo.examples}
</Text>
<AntdSelect
mode="multiple"
value={value.tiers[tier]}
onChange={(model) => handleTierChange(tier, model)}
placeholder={`Select model for ${tierInfo.label.toLowerCase()} queries`}
onChange={(models) => handleTierChange(tier, models)}
placeholder={`Select model(s) for ${tierInfo.label.toLowerCase()} queries`}
showSearch
style={{ width: "100%" }}
options={modelOptions}
status={tierMissing ? "error" : undefined}
/>
{value.tiers[tier].length > 1 && (
<Text type="secondary" style={{ fontSize: 12 }}>
Multiple models selected — the router randomly picks among them per request (or Thompson-samples
within the pool when adaptive routing is on).
</Text>
)}
{tierMissing && (
<Text type="danger" style={{ fontSize: 12 }}>
This tier is required
@ -205,148 +194,61 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
</Text>
),
children: (
<>
<Radio.Group
value={value.classifier_type}
onChange={(e) => handleClassifierTypeChange(e.target.value)}
className="w-full"
>
<Space direction="vertical" className="w-full">
<Radio value="heuristic">
<Text strong>Heuristic</Text>{" "}
<Text type="secondary">(default) — rule-based scoring, no API calls, &lt;1ms latency</Text>
</Radio>
<Radio value="llm">
<Text strong>LLM Classifier</Text>{" "}
<Text type="secondary">— use a model to decide the tier (e.g. a small/fast model)</Text>
</Radio>
</Space>
</Radio.Group>
{value.classifier_type === "llm" && (
<div className="mt-4 space-y-3">
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Classifier Model
</Text>
<AntdSelect
value={value.classifier_llm_config?.model || undefined}
onChange={handleClassifierModelChange}
placeholder="Select the model that will classify request complexity"
showSearch
style={{ width: "100%" }}
options={modelOptions}
status={classifierModelMissing ? "error" : undefined}
/>
{classifierModelMissing && (
<Text type="danger" style={{ fontSize: 12 }}>
A classifier model is required
</Text>
)}
</div>
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Timeout (ms)
</Text>
<InputNumber
value={value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS}
onChange={handleClassifierTimeoutChange}
min={1}
style={{ width: "100%" }}
/>
<Text type="secondary" style={{ fontSize: 12 }}>
Falls back to the heuristic scorer if the classifier call errors, times out, or returns an
unparseable response.
</Text>
</div>
</div>
)}
</>
<ClassificationMethodConfig
value={value}
onChange={onChange}
modelOptions={modelOptions}
customTechnicalKeywords={customTechnicalKeywords}
onCustomTechnicalKeywordsChange={onCustomTechnicalKeywordsChange}
showValidationErrors={showValidationErrors}
/>
),
},
{
key: "adaptive",
label: (
<Text strong style={{ color: "#374151" }}>
Advanced: Adaptive Routing
</Text>
),
children: <AdaptiveRoutingConfig value={value} onChange={onChange} />,
},
...(onKeywordTierRulesChange || onSemanticMatchingEnabledChange
? [
{
key: "keyword-semantic",
label: (
<Text strong style={{ color: "#374151" }}>
Advanced: Keyword/Semantic Matching
</Text>
),
children: (
<>
{onKeywordTierRulesChange && (
<KeywordTierRules rules={keywordTierRules} onChange={onKeywordTierRulesChange} />
)}
{onKeywordTierRulesChange && onSemanticMatchingEnabledChange && (
<Divider style={{ margin: "16px 0" }} />
)}
{onSemanticMatchingEnabledChange && (
<SemanticKeywordMatching
enabled={semanticMatchingEnabled}
onEnabledChange={onSemanticMatchingEnabledChange}
embeddingModel={embeddingModel}
onEmbeddingModelChange={onEmbeddingModelChange}
matchThreshold={matchThreshold}
onMatchThresholdChange={onMatchThresholdChange}
modelInfo={modelInfo}
showValidationErrors={showValidationErrors}
/>
)}
</>
),
},
]
: []),
]}
/>
<Divider />
<Card>
<div className="flex items-center gap-2 mb-2">
<Text strong style={{ fontSize: 16 }}>
Custom Technical Keywords
</Text>
<Tooltip title="Domain-specific terms appended to the built-in technical keyword list. Prompts containing these terms score higher on the technical dimension and route to more capable models.">
<InfoCircleOutlined className="text-gray-400" />
</Tooltip>
</div>
<Text type="secondary" style={{ display: "block", marginBottom: 8, fontSize: 12 }}>
Optional: Add terms to the built-in list to improve classification accuracy on the technical dimension. (e.g.,
udp, kafka, terraform).
</Text>
<AntdSelect
mode="tags"
value={customTechnicalKeywords ?? []}
onChange={(keywords: string[]) => onCustomTechnicalKeywordsChange?.(keywords)}
placeholder="Type a keyword and press Enter, or paste a comma-separated list"
tokenSeparators={[","]}
open={false}
suffixIcon={null}
style={{ width: "100%" }}
allowClear
/>
</Card>
<Divider />
<Card className="bg-gray-50">
<Text strong style={{ display: "block", marginBottom: 8 }}>
How Classification Works
</Text>
<Text type="secondary" style={{ fontSize: 13 }}>
The router scores each request across 7 dimensions: token count, code presence, reasoning markers, technical
terms, simple indicators, multi-step patterns, and question complexity. The weighted score determines the
tier:
</Text>
<ul style={{ marginTop: 8, marginBottom: 0, paddingLeft: 20, fontSize: 13, color: "rgba(0, 0, 0, 0.45)" }}>
<li>
<strong>SIMPLE</strong>: Score &lt; 0.15
</li>
<li>
<strong>MEDIUM</strong>: Score 0.15 - 0.35
</li>
<li>
<strong>COMPLEX</strong>: Score 0.35 - 0.60
</li>
<li>
<strong>REASONING</strong>: Score &gt; 0.60 (or 2+ reasoning markers)
</li>
</ul>
</Card>
{/* Keyword-tier and semantic sections only render when their change handlers are
wired (the add-router flow). The edit-auto-router modal doesn't pass them yet, so
they stay hidden there rather than rendering interactive-but-dead controls. */}
{onKeywordTierRulesChange && (
<>
<Divider />
<KeywordTierRules rules={keywordTierRules} onChange={onKeywordTierRulesChange} />
</>
)}
{onSemanticMatchingEnabledChange && (
<>
<Divider />
<SemanticKeywordMatching
enabled={semanticMatchingEnabled}
onEnabledChange={onSemanticMatchingEnabledChange}
embeddingModel={embeddingModel}
onEmbeddingModelChange={onEmbeddingModelChange}
matchThreshold={matchThreshold}
onMatchThresholdChange={onMatchThresholdChange}
modelInfo={modelInfo}
showValidationErrors={showValidationErrors}
/>
</>
)}
</div>
);
};

View file

@ -1,5 +1,5 @@
import { InfoCircleOutlined } from "@ant-design/icons";
import { Card, InputNumber, Select as AntdSelect, Switch, Tooltip, Typography } from "antd";
import { InputNumber, Select as AntdSelect, Switch, Tooltip, Typography } from "antd";
import React from "react";
import { ModelGroup } from "@/components/llm_calls/fetch_models";
@ -36,7 +36,7 @@ const SemanticKeywordMatching: React.FC<SemanticKeywordMatchingProps> = ({
const embeddingModelMissing = showValidationErrors && !embeddingModel;
return (
<Card className="mb-4">
<div>
<div className="flex items-center justify-between gap-4">
<div>
<div className="flex items-center gap-2">
@ -86,7 +86,7 @@ const SemanticKeywordMatching: React.FC<SemanticKeywordMatchingProps> = ({
</div>
</div>
)}
</Card>
</div>
);
};

View file

@ -8,7 +8,11 @@ import { all_admin_roles } from "@/utils/roles";
import { handleAddAutoRouterSubmit } from "./handle_add_auto_router_submit";
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
import RouterConfigBuilder from "./RouterConfigBuilder";
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
import ComplexityRouterConfig, {
ComplexityRouterConfigValue,
DEFAULT_ADAPTIVE_WEIGHTS,
DEFAULT_TIER_DISTANCE_PENALTY,
} from "./ComplexityRouterConfig";
import { KeywordTierRule } from "./KeywordTierRules";
import { DEFAULT_MATCH_THRESHOLD } from "./SemanticKeywordMatching";
import {
@ -38,7 +42,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
const [routerType, setRouterType] = useState<RouterType>("recommended");
const [complexityRouterConfig, setComplexityRouterConfig] = useState<ComplexityRouterConfigValue>({
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] },
classifier_type: "heuristic",
});
@ -89,6 +93,10 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
tiers,
classifier_type: classifierType,
classifier_llm_config: classifierLlmConfig,
adaptive = false,
adaptive_weights: adaptiveWeights = DEFAULT_ADAPTIVE_WEIGHTS,
tier_distance_penalty: tierDistancePenalty = DEFAULT_TIER_DISTANCE_PENALTY,
adaptive_eligible: adaptiveEligible = "all",
} = complexityRouterConfig;
const missingTiersError = getMissingTiersError(tiers);
@ -111,7 +119,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
return;
}
const defaultModel = tiers.MEDIUM || tiers.SIMPLE || tiers.COMPLEX || tiers.REASONING;
const defaultModel = tiers.MEDIUM[0] || tiers.SIMPLE[0] || tiers.COMPLEX[0] || tiers.REASONING[0];
form.setFieldsValue({
custom_llm_provider: "auto_router",
@ -132,6 +140,10 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
semanticMatchingEnabled,
embeddingModel,
matchThreshold,
adaptive,
adaptiveWeights,
tierDistancePenalty,
adaptiveEligible,
};
const submitValues = {

View file

@ -1,10 +1,10 @@
import { buildAutoRouterTestTargets } from "./build_auto_router_test_targets";
const tiers = {
SIMPLE: "gpt-4o-mini",
MEDIUM: "claude-sonnet-4",
COMPLEX: "claude-sonnet-4",
REASONING: "o3",
SIMPLE: ["gpt-4o-mini"],
MEDIUM: ["claude-sonnet-4"],
COMPLEX: ["claude-sonnet-4"],
REASONING: ["o3"],
};
describe("buildAutoRouterTestTargets", () => {
@ -17,9 +17,21 @@ describe("buildAutoRouterTestTargets", () => {
]);
});
it("emits a target per model when a tier has more than one, and dedups across tiers", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: ["gpt-4o-mini", "claude-sonnet-4"], MEDIUM: ["claude-sonnet-4"], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: false,
embeddingModel: undefined,
});
expect(targets).toEqual([
{ labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" },
{ labels: ["SIMPLE", "MEDIUM"], modelGroup: "claude-sonnet-4", mode: "chat" },
]);
});
it("drops empty/whitespace tiers", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: " ", REASONING: "" },
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: [], COMPLEX: [" "], REASONING: [] },
semanticMatchingEnabled: false,
embeddingModel: undefined,
});
@ -29,7 +41,7 @@ describe("buildAutoRouterTestTargets", () => {
it("returns [] when no tier is configured", () => {
expect(
buildAutoRouterTestTargets({
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: false,
embeddingModel: undefined,
}),
@ -38,7 +50,7 @@ describe("buildAutoRouterTestTargets", () => {
it("appends an embedding target only when semantic matching is on and a model is set", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: [], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: true,
embeddingModel: "voyage-3-5",
});
@ -50,7 +62,7 @@ describe("buildAutoRouterTestTargets", () => {
it("omits the embedding target when semantic matching is on but no model is chosen", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: [], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: true,
embeddingModel: undefined,
});
@ -59,7 +71,7 @@ describe("buildAutoRouterTestTargets", () => {
it("omits the embedding target when a model is set but semantic matching is off", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: [], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: false,
embeddingModel: "voyage-3-5",
});

View file

@ -29,9 +29,11 @@ export const buildAutoRouterTestTargets = ({
embeddingModel,
}: BuildAutoRouterTestTargetsParams): AutoRouterTestTarget[] => {
const groupedByModel = TIER_ORDER.reduce<Record<string, string[]>>((acc, tier) => {
const modelGroup = tiers[tier]?.trim();
if (!modelGroup) return acc;
return { ...acc, [modelGroup]: [...(acc[modelGroup] ?? []), tier] };
return (tiers[tier] ?? []).reduce((tierAcc, rawModel) => {
const modelGroup = rawModel?.trim();
if (!modelGroup) return tierAcc;
return { ...tierAcc, [modelGroup]: [...(tierAcc[modelGroup] ?? []), tier] };
}, acc);
}, {});
const tierTargets: AutoRouterTestTarget[] = Object.entries(groupedByModel).map(([modelGroup, labels]) => ({

View file

@ -6,10 +6,10 @@ import {
} from "./build_complexity_router_config";
const tiers = {
SIMPLE: "gpt-4o-mini",
MEDIUM: "gpt-4o",
COMPLEX: "claude-sonnet-4",
REASONING: "o1-preview",
SIMPLE: ["gpt-4o-mini"],
MEDIUM: ["gpt-4o"],
COMPLEX: ["claude-sonnet-4"],
REASONING: ["o1-preview"],
};
const baseParams: BuildComplexityRouterConfigParams = {
@ -21,6 +21,10 @@ const baseParams: BuildComplexityRouterConfigParams = {
semanticMatchingEnabled: false,
embeddingModel: undefined,
matchThreshold: 0.5,
adaptive: false,
adaptiveWeights: { quality: 0.3, cost: 0.7 },
tierDistancePenalty: 0.5,
adaptiveEligible: "all",
};
describe("buildComplexityRouterConfig", () => {
@ -29,6 +33,14 @@ describe("buildComplexityRouterConfig", () => {
expect(config).toEqual({ tiers, classifier_type: "heuristic" });
});
it("passes through a tier configured with more than one model as a pool", () => {
const config = buildComplexityRouterConfig({
...baseParams,
tiers: { ...tiers, SIMPLE: ["gpt-4o-mini", "gpt-4o", "claude-haiku-4-5"] },
});
expect(config.tiers.SIMPLE).toEqual(["gpt-4o-mini", "gpt-4o", "claude-haiku-4-5"]);
});
it("includes classifier_llm_config only when classifier_type is llm", () => {
const config = buildComplexityRouterConfig({
...baseParams,
@ -123,6 +135,47 @@ describe("buildComplexityRouterConfig", () => {
const config = buildComplexityRouterConfig(params);
expect(config.keyword_tier_rules).toBeUndefined();
});
it("omits adaptive fields when adaptive is disabled even if weights linger in state", () => {
const config = buildComplexityRouterConfig({
...baseParams,
adaptive: false,
adaptiveWeights: { quality: 0.9, cost: 0.1 },
tierDistancePenalty: 2,
adaptiveEligible: "classified_tier",
});
expect(config.adaptive).toBeUndefined();
expect(config.adaptive_weights).toBeUndefined();
expect(config.tier_distance_penalty).toBeUndefined();
expect(config.adaptive_eligible).toBeUndefined();
});
it("includes tier_distance_penalty when adaptive is enabled with eligible='all'", () => {
const config = buildComplexityRouterConfig({
...baseParams,
adaptive: true,
adaptiveWeights: { quality: 0.6, cost: 0.4 },
tierDistancePenalty: 0.75,
adaptiveEligible: "all",
});
expect(config.adaptive).toBe(true);
expect(config.adaptive_weights).toEqual({ quality: 0.6, cost: 0.4 });
expect(config.tier_distance_penalty).toBe(0.75);
expect(config.adaptive_eligible).toBe("all");
});
it("omits tier_distance_penalty when eligible='classified_tier', since the penalty doesn't apply there", () => {
const config = buildComplexityRouterConfig({
...baseParams,
adaptive: true,
adaptiveWeights: { quality: 0.6, cost: 0.4 },
tierDistancePenalty: 0.75,
adaptiveEligible: "classified_tier",
});
expect(config.adaptive).toBe(true);
expect(config.adaptive_eligible).toBe("classified_tier");
expect(config.tier_distance_penalty).toBeUndefined();
});
});
describe("getMissingTiersError", () => {
@ -131,23 +184,27 @@ describe("getMissingTiersError", () => {
});
it("names the specific missing tier when only one is blank", () => {
expect(getMissingTiersError({ ...tiers, REASONING: "" })).toBe(
expect(getMissingTiersError({ ...tiers, REASONING: [] })).toBe(
"Select a model for the following tier(s): REASONING",
);
});
it("names multiple missing tiers in SIMPLE/MEDIUM/COMPLEX/REASONING order", () => {
expect(getMissingTiersError({ ...tiers, SIMPLE: "", REASONING: "" })).toBe(
expect(getMissingTiersError({ ...tiers, SIMPLE: [], REASONING: [] })).toBe(
"Select a model for the following tier(s): SIMPLE, REASONING",
);
});
it("names all four tiers when none are filled", () => {
const noTiers = { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" };
const noTiers = { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] };
expect(getMissingTiersError(noTiers)).toBe(
"Select a model for the following tier(s): SIMPLE, MEDIUM, COMPLEX, REASONING",
);
});
it("treats a tier with more than one model as filled", () => {
expect(getMissingTiersError({ ...tiers, SIMPLE: ["gpt-4o-mini", "gpt-4o"] })).toBeNull();
});
});
describe("getSemanticConfigError", () => {

View file

@ -1,12 +1,11 @@
import { KeywordTierRule } from "./KeywordTierRules";
import { ClassifierLLMConfig, ClassifierType } from "./ComplexityRouterConfig";
export interface ComplexityTiers {
SIMPLE: string;
MEDIUM: string;
COMPLEX: string;
REASONING: string;
}
import {
AdaptiveEligible,
AdaptiveRouterWeights,
ClassifierLLMConfig,
ClassifierType,
ComplexityTiers,
} from "./ComplexityRouterConfig";
export interface BuildComplexityRouterConfigParams {
tiers: ComplexityTiers;
@ -17,6 +16,10 @@ export interface BuildComplexityRouterConfigParams {
semanticMatchingEnabled: boolean;
embeddingModel: string | undefined;
matchThreshold: number;
adaptive: boolean;
adaptiveWeights: AdaptiveRouterWeights;
tierDistancePenalty: number;
adaptiveEligible: AdaptiveEligible;
}
export interface ComplexityRouterConfigPayload {
@ -28,12 +31,16 @@ export interface ComplexityRouterConfigPayload {
semantic_keyword_matching?: boolean;
embedding_model?: string;
match_threshold?: number;
adaptive?: boolean;
adaptive_weights?: AdaptiveRouterWeights;
tier_distance_penalty?: number;
adaptive_eligible?: AdaptiveEligible;
}
const TIER_KEYS: Array<keyof ComplexityTiers> = ["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"];
export const getMissingTiersError = (tiers: ComplexityTiers): string | null => {
const missing = TIER_KEYS.filter((tier) => !tiers[tier]);
const missing = TIER_KEYS.filter((tier) => tiers[tier].length === 0);
if (missing.length === 0) return null;
return `Select a model for the following tier(s): ${missing.join(", ")}`;
};
@ -43,7 +50,8 @@ export const getSemanticConfigError = ({
embeddingModel,
keywordTierRules,
}: Pick<BuildComplexityRouterConfigParams, "semanticMatchingEnabled" | "embeddingModel" | "keywordTierRules">):
string | null => {
| string
| null => {
if (!semanticMatchingEnabled) return null;
if (!embeddingModel) return "Select an embedding model to use semantic keyword matching";
if (keywordTierRules.length === 0) return "Add at least one keyword tier rule to use semantic keyword matching";
@ -61,6 +69,10 @@ export const buildComplexityRouterConfig = ({
semanticMatchingEnabled,
embeddingModel,
matchThreshold,
adaptive,
adaptiveWeights,
tierDistancePenalty,
adaptiveEligible,
}: BuildComplexityRouterConfigParams): ComplexityRouterConfigPayload => {
// Trim keywords and drop empty ones; drop any rule left with no keywords. Clicking
// "Add keyword rule" seeds a rule with an empty keywords list, so without this an
@ -81,5 +93,11 @@ export const buildComplexityRouterConfig = ({
embedding_model: embeddingModel,
match_threshold: matchThreshold,
}),
...(adaptive && {
adaptive: true,
adaptive_weights: adaptiveWeights,
...(adaptiveEligible === "all" && { tier_distance_penalty: tierDistancePenalty }),
adaptive_eligible: adaptiveEligible,
}),
};
};

View file

@ -0,0 +1,104 @@
import { buildUpdatedComplexityRouterConfig } from "./edit_auto_router_modal";
const storedConfigValue = {
tiers: {
SIMPLE: "old-simple",
MEDIUM: "old-medium",
COMPLEX: "old-complex",
REASONING: "old-reasoning",
},
classifier_type: "llm",
classifier_llm_config: { model: "old-classifier", timeout_ms: 1200 },
custom_technical_keywords: ["kafka", "terraform"],
keyword_tier_rules: [{ keywords: ["invoice", "refund"], tier: "MEDIUM" }],
semantic_keyword_matching: true,
embedding_model: "voyage-4-large",
match_threshold: 0.65,
adaptive: true,
adaptive_weights: { quality: 0.3, cost: 0.7 },
tier_distance_penalty: 0.8,
adaptive_eligible: "all",
};
const storedConfig = JSON.stringify(storedConfigValue);
const tiers = {
SIMPLE: ["gpt-4o-mini"],
MEDIUM: ["gpt-4o-mini"],
COMPLEX: ["anthropic-sonnet-4-5"],
REASONING: ["anthropic-sonnet-4-5"],
};
const classifiedTierValue = {
tiers,
classifier_type: "heuristic" as const,
adaptive: true,
adaptive_weights: { quality: 0.4, cost: 0.6 },
tier_distance_penalty: 0.8,
adaptive_eligible: "classified_tier" as const,
};
const expectedClassifiedTierConfig = {
tiers,
classifier_type: "heuristic",
custom_technical_keywords: ["kafka", "terraform"],
keyword_tier_rules: [{ keywords: ["invoice", "refund"], tier: "MEDIUM" }],
semantic_keyword_matching: true,
embedding_model: "voyage-4-large",
match_threshold: 0.65,
adaptive: true,
adaptive_weights: { quality: 0.4, cost: 0.6 },
adaptive_eligible: "classified_tier",
};
const adaptiveDisabledValue = {
tiers,
classifier_type: "heuristic" as const,
adaptive: false,
};
const expectedAdaptiveDisabledConfig = {
tiers,
classifier_type: "heuristic",
custom_technical_keywords: ["kafka", "terraform"],
keyword_tier_rules: [{ keywords: ["invoice", "refund"], tier: "MEDIUM" }],
semantic_keyword_matching: true,
embedding_model: "voyage-4-large",
match_threshold: 0.65,
};
describe("buildUpdatedComplexityRouterConfig", () => {
it("preserves unrelated options and omits the penalty for classified-tier routing", () => {
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, classifiedTierValue);
expect(updatedConfig).toEqual(expectedClassifiedTierConfig);
});
it("removes managed adaptive and classifier fields when they are disabled", () => {
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, adaptiveDisabledValue);
expect(updatedConfig).toEqual(expectedAdaptiveDisabledConfig);
});
it("updates custom technical keywords when they are edited", () => {
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, classifiedTierValue, ["postgres"]);
expect(updatedConfig.custom_technical_keywords).toEqual(["postgres"]);
});
it("removes custom technical keywords when they are cleared", () => {
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, classifiedTierValue, []);
expect(updatedConfig.custom_technical_keywords).toBeUndefined();
});
it("preserves a tier configured with more than one model as a pool", () => {
const multiModelValue = {
...classifiedTierValue,
tiers: { ...tiers, SIMPLE: ["gpt-4o-mini", "claude-haiku-4-5"] },
};
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, multiModelValue);
expect(updatedConfig.tiers).toMatchObject({ SIMPLE: ["gpt-4o-mini", "claude-haiku-4-5"] });
});
});

View file

@ -4,13 +4,23 @@ import { Text, TextInput } from "@tremor/react";
import { modelAvailableCall, modelPatchUpdateCall } from "../networking";
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
import RouterConfigBuilder from "../add_model/RouterConfigBuilder";
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "../add_model/ComplexityRouterConfig";
import ComplexityRouterConfig, {
ComplexityRouterConfigValue,
DEFAULT_ADAPTIVE_WEIGHTS,
DEFAULT_TIER_DISTANCE_PENALTY,
} from "../add_model/ComplexityRouterConfig";
import NotificationsManager from "../molecules/notifications_manager";
const isComplexityRouterModel = (modelData: any): boolean =>
modelData?.litellm_params?.model?.startsWith("auto_router/complexity_router") ||
modelData?.litellm_params?.complexity_router_config != null;
const normalizeTierModels = (value: unknown): string[] => {
if (Array.isArray(value)) return value;
if (typeof value === "string" && value) return [value];
return [];
};
interface EditAutoRouterModalProps {
isVisible: boolean;
onCancel: () => void;
@ -20,6 +30,57 @@ interface EditAutoRouterModalProps {
userRole: string;
}
const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([
"tiers",
"classifier_type",
"classifier_llm_config",
"adaptive",
"adaptive_weights",
"tier_distance_penalty",
"adaptive_eligible",
]);
const toRecord = (value: unknown): Record<string, unknown> => {
const parsed: unknown = typeof value === "string" ? JSON.parse(value) : value;
return typeof parsed === "object" && parsed !== null && !Array.isArray(parsed)
? (parsed as Record<string, unknown>)
: {};
};
export const buildUpdatedComplexityRouterConfig = (
storedConfig: unknown,
value: ComplexityRouterConfigValue,
customTechnicalKeywords?: string[],
): Record<string, unknown> => {
const preservedConfig = Object.fromEntries(
Object.entries(toRecord(storedConfig)).filter(
([key]) =>
!MANAGED_COMPLEXITY_ROUTER_KEYS.has(key) &&
(customTechnicalKeywords === undefined || key !== "custom_technical_keywords"),
),
);
const adaptiveEligible = value.adaptive_eligible ?? "all";
return {
...preservedConfig,
tiers: value.tiers,
classifier_type: value.classifier_type,
...(value.classifier_type === "llm" ? { classifier_llm_config: value.classifier_llm_config } : {}),
...(customTechnicalKeywords &&
customTechnicalKeywords.length > 0 && {
custom_technical_keywords: customTechnicalKeywords,
}),
...(value.adaptive && {
adaptive: true,
adaptive_weights: value.adaptive_weights ?? DEFAULT_ADAPTIVE_WEIGHTS,
...(adaptiveEligible === "all" && {
tier_distance_penalty: value.tier_distance_penalty ?? DEFAULT_TIER_DISTANCE_PENALTY,
}),
adaptive_eligible: adaptiveEligible,
}),
};
};
const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
isVisible,
onCancel,
@ -35,8 +96,9 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
const [showCustomDefaultModel, setShowCustomDefaultModel] = useState<boolean>(false);
const [showCustomEmbeddingModel, setShowCustomEmbeddingModel] = useState<boolean>(false);
const [routerConfig, setRouterConfig] = useState<any>(null);
const [customTechnicalKeywords, setCustomTechnicalKeywords] = useState<string[]>([]);
const [complexityRouterConfig, setComplexityRouterConfig] = useState<ComplexityRouterConfigValue>({
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] },
classifier_type: "heuristic",
});
const isComplexityRouter = isComplexityRouterModel(modelData);
@ -85,14 +147,21 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
setComplexityRouterConfig({
tiers: {
SIMPLE: parsedConfig.tiers?.SIMPLE || "",
MEDIUM: parsedConfig.tiers?.MEDIUM || "",
COMPLEX: parsedConfig.tiers?.COMPLEX || "",
REASONING: parsedConfig.tiers?.REASONING || "",
SIMPLE: normalizeTierModels(parsedConfig.tiers?.SIMPLE),
MEDIUM: normalizeTierModels(parsedConfig.tiers?.MEDIUM),
COMPLEX: normalizeTierModels(parsedConfig.tiers?.COMPLEX),
REASONING: normalizeTierModels(parsedConfig.tiers?.REASONING),
},
classifier_type: parsedConfig.classifier_type || "heuristic",
classifier_llm_config: parsedConfig.classifier_llm_config,
adaptive: parsedConfig.adaptive || false,
adaptive_weights: parsedConfig.adaptive_weights,
tier_distance_penalty: parsedConfig.tier_distance_penalty,
adaptive_eligible: parsedConfig.adaptive_eligible || "all",
});
setCustomTechnicalKeywords(
Array.isArray(parsedConfig.custom_technical_keywords) ? parsedConfig.custom_technical_keywords : [],
);
form.setFieldsValue({
auto_router_name: modelData.model_name,
@ -138,7 +207,7 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
if (isComplexityRouter) {
const { tiers, classifier_type, classifier_llm_config } = complexityRouterConfig;
if (Object.values(tiers).filter(Boolean).length === 0) {
if (Object.values(tiers).every((models) => models.length === 0)) {
NotificationsManager.fromBackend("Please select at least one model for a complexity tier");
return;
}
@ -147,14 +216,14 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
return;
}
const defaultModel = tiers.MEDIUM || tiers.SIMPLE || tiers.COMPLEX || tiers.REASONING;
const defaultModel = tiers.MEDIUM[0] || tiers.SIMPLE[0] || tiers.COMPLEX[0] || tiers.REASONING[0];
const updatedLitellmParams = {
...modelData.litellm_params,
complexity_router_config: {
tiers,
classifier_type,
...(classifier_type === "llm" ? { classifier_llm_config } : {}),
},
complexity_router_config: buildUpdatedComplexityRouterConfig(
modelData.litellm_params?.complexity_router_config,
complexityRouterConfig,
customTechnicalKeywords,
),
complexity_router_default_model: defaultModel,
};
const updatedModelInfo = {
@ -264,6 +333,8 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
onChange={(config) => {
setComplexityRouterConfig(config);
}}
customTechnicalKeywords={customTechnicalKeywords}
onCustomTechnicalKeywordsChange={setCustomTechnicalKeywords}
/>
</div>
) : (