use llm_router instead of litellm.acompletion, extract helpers, move constants

This commit is contained in:
Ishaan Jaffer 2026-02-18 20:01:29 -08:00
parent 96a02347a8
commit 175a33a4b1
2 changed files with 229 additions and 175 deletions

View file

@ -1474,3 +1474,8 @@ MICROSOFT_USER_FIRST_NAME_ATTRIBUTE = str(
MICROSOFT_USER_LAST_NAME_ATTRIBUTE = str(
os.getenv("MICROSOFT_USER_LAST_NAME_ATTRIBUTE", "surname")
)
# Policy template enrichment
MAX_COMPETITOR_NAMES = 30
COMPETITOR_LLM_TEMPERATURE = 0.3
DEFAULT_COMPETITOR_DISCOVERY_MODEL = "gpt-4o-mini"

View file

@ -9,15 +9,29 @@ All /policy management endpoints
/policy/templates - Get policy templates (GitHub with local fallback)
"""
import copy
import json
import os
from typing import TYPE_CHECKING, List, Literal, Optional, TypedDict, cast
from typing import (
TYPE_CHECKING,
AsyncIterator,
List,
Literal,
Optional,
TypedDict,
cast,
)
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
COMPETITOR_LLM_TEMPERATURE,
DEFAULT_COMPETITOR_DISCOVERY_MODEL,
MAX_COMPETITOR_NAMES,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -516,6 +530,31 @@ class EnrichTemplateRequest(BaseModel):
competitors: Optional[list] = None
def _validate_enrichment_request(data: EnrichTemplateRequest) -> tuple[dict, dict, str]:
"""
Validate enrichment request and return (template, llm_enrichment, brand_name).
Raises HTTPException on validation failure.
"""
templates = _load_policy_templates_from_local_backup()
template = next((t for t in templates if t.get("id") == data.template_id), None)
if template is None:
raise HTTPException(status_code=404, detail=f"Template '{data.template_id}' not found")
llm_enrichment = template.get("llm_enrichment")
if llm_enrichment is None:
raise HTTPException(status_code=400, detail="Template does not support LLM enrichment")
brand_name = data.parameters.get(llm_enrichment["parameter"], "")
if not brand_name:
raise HTTPException(
status_code=400,
detail=f"Parameter '{llm_enrichment['parameter']}' is required",
)
return template, llm_enrichment, brand_name
@router.post(
"/policy/templates/enrich",
tags=["policy management"],
@ -533,38 +572,17 @@ async def enrich_policy_template(
Calls an onboarded LLM to discover competitors for the given brand name,
then returns enriched guardrailDefinitions with the discovered data populated.
"""
templates = _load_policy_templates_from_local_backup()
template = next((t for t in templates if t.get("id") == data.template_id), None)
if template is None:
raise HTTPException(status_code=404, detail=f"Template '{data.template_id}' not found")
llm_enrichment = template.get("llm_enrichment")
if llm_enrichment is None:
raise HTTPException(
status_code=400,
detail="Template does not support LLM enrichment",
)
brand_name = data.parameters.get(llm_enrichment["parameter"], "")
if not brand_name:
raise HTTPException(
status_code=400,
detail=f"Parameter '{llm_enrichment['parameter']}' is required",
)
template, llm_enrichment, brand_name = _validate_enrichment_request(data)
model = data.model or DEFAULT_COMPETITOR_DISCOVERY_MODEL
if data.competitors:
# Free-form mode: use user-provided competitor list directly
competitors = data.competitors
else:
# AI mode: discover competitors via LLM
prompt = llm_enrichment["prompt"].replace(
"{{" + llm_enrichment["parameter"] + "}}", brand_name
)
model = data.model or "gpt-4o-mini"
competitors = await _discover_competitors_via_llm(prompt, model=model)
# Generate common variations/misspellings for each competitor
model = data.model or "gpt-4o-mini"
variations_map = await _generate_competitor_variations(competitors, model=model)
enriched_definitions = _build_competitor_guardrail_definitions(
@ -581,6 +599,66 @@ async def enrich_policy_template(
}
async def _stream_competitor_events(
data: EnrichTemplateRequest,
template: dict,
llm_enrichment: dict,
brand_name: str,
model: str,
) -> AsyncIterator[str]:
"""Stream competitor names as SSE events, then emit a final 'done' event."""
competitors: list[str] = []
if data.competitors:
competitors = data.competitors
for comp in competitors:
yield f"data: {json.dumps({'type': 'competitor', 'name': comp})}\n\n"
else:
prompt = llm_enrichment["prompt"].replace(
"{{" + llm_enrichment["parameter"] + "}}", brand_name
)
try:
from litellm.proxy.proxy_server import llm_router
if llm_router is None:
raise ValueError("LLM router not initialized")
response = await llm_router.acompletion(
model=model,
messages=[{"role": "user", "content": prompt}],
temperature=COMPETITOR_LLM_TEMPERATURE,
stream=True,
)
buffer = ""
async for chunk in response: # type: ignore[union-attr]
delta = chunk.choices[0].delta.content or ""
buffer += delta
while "\n" in buffer:
line, buffer = buffer.split("\n", 1)
name = _clean_competitor_line(line)
if name and len(competitors) < MAX_COMPETITOR_NAMES:
competitors.append(name)
yield f"data: {json.dumps({'type': 'competitor', 'name': name})}\n\n"
# Handle remaining buffer
name = _clean_competitor_line(buffer)
if name and len(competitors) < MAX_COMPETITOR_NAMES:
competitors.append(name)
yield f"data: {json.dumps({'type': 'competitor', 'name': name})}\n\n"
except Exception as e:
verbose_proxy_logger.error("LLM competitor streaming failed: %s", e)
yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n"
return
variations_map = await _generate_competitor_variations(competitors, model=model)
enriched_definitions = _build_competitor_guardrail_definitions(
template.get("guardrailDefinitions", []),
competitors,
brand_name,
variations_map,
)
yield f"data: {json.dumps({'type': 'done', 'competitors': competitors, 'competitor_variations': variations_map, 'guardrailDefinitions': enriched_definitions})}\n\n"
@router.post(
"/policy/templates/enrich/stream",
tags=["policy management"],
@ -598,95 +676,26 @@ async def enrich_policy_template_stream(
- data: {"type": "competitor", "name": "..."} each competitor as discovered
- data: {"type": "done", "competitors": [...], "competitor_variations": {...}, "guardrailDefinitions": [...]}
"""
import litellm
templates = _load_policy_templates_from_local_backup()
template = next((t for t in templates if t.get("id") == data.template_id), None)
if template is None:
raise HTTPException(status_code=404, detail=f"Template '{data.template_id}' not found")
llm_enrichment = template.get("llm_enrichment")
if llm_enrichment is None:
raise HTTPException(status_code=400, detail="Template does not support LLM enrichment")
brand_name = data.parameters.get(llm_enrichment["parameter"], "")
if not brand_name:
raise HTTPException(
status_code=400,
detail=f"Parameter '{llm_enrichment['parameter']}' is required",
)
model = data.model or "gpt-4o-mini"
async def event_generator():
competitors: list[str] = []
if data.competitors:
# Free-form mode — emit all at once
competitors = data.competitors
for comp in competitors:
yield f"data: {json.dumps({'type': 'competitor', 'name': comp})}\n\n"
else:
# AI mode — stream from LLM
prompt = llm_enrichment["prompt"].replace(
"{{" + llm_enrichment["parameter"] + "}}", brand_name
)
try:
response = await litellm.acompletion(
model=model,
messages=[{"role": "user", "content": prompt}],
temperature=0.3,
stream=True,
)
buffer = ""
async for chunk in response: # type: ignore[union-attr]
delta = chunk.choices[0].delta.content or ""
buffer += delta
# Parse complete lines as they come in
while "\n" in buffer:
line, buffer = buffer.split("\n", 1)
name = line.strip().strip(".-) ").strip()
if name and len(name) > 1 and len(competitors) < 30:
competitors.append(name)
yield f"data: {json.dumps({'type': 'competitor', 'name': name})}\n\n"
# Handle remaining buffer
if buffer.strip():
name = buffer.strip().strip(".-) ").strip()
if name and len(name) > 1 and len(competitors) < 30:
competitors.append(name)
yield f"data: {json.dumps({'type': 'competitor', 'name': name})}\n\n"
except Exception as e:
verbose_proxy_logger.error("LLM competitor streaming failed: %s", e)
yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n"
return
# Generate variations
variations_map = await _generate_competitor_variations(competitors, model=model)
# Build enriched definitions
enriched_definitions = _build_competitor_guardrail_definitions(
template.get("guardrailDefinitions", []),
competitors,
brand_name,
variations_map,
)
# Send final event with all data
yield f"data: {json.dumps({'type': 'done', 'competitors': competitors, 'competitor_variations': variations_map, 'guardrailDefinitions': enriched_definitions})}\n\n"
template, llm_enrichment, brand_name = _validate_enrichment_request(data)
model = data.model or DEFAULT_COMPETITOR_DISCOVERY_MODEL
return StreamingResponse(
event_generator(),
_stream_competitor_events(data, template, llm_enrichment, brand_name, model),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)
def _clean_competitor_line(line: str) -> Optional[str]:
"""Strip numbering, bullets, and whitespace from a competitor name line."""
name = line.strip().strip(".-) ").strip()
return name if name and len(name) > 1 else None
async def _generate_competitor_variations(
competitors: list, model: str = "gpt-4o-mini"
competitors: list, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL
) -> dict:
"""Generate common misspellings, abbreviations, and alternate names for each competitor."""
import litellm
if not competitors:
return {}
@ -702,64 +711,81 @@ async def _generate_competitor_variations(
)
try:
response = await litellm.acompletion(
from litellm.proxy.proxy_server import llm_router
if llm_router is None:
raise ValueError("LLM router not initialized")
response = await llm_router.acompletion(
model=model,
messages=[{"role": "user", "content": prompt}],
temperature=0.3,
temperature=COMPETITOR_LLM_TEMPERATURE,
)
raw = response.choices[0].message.content or "" # type: ignore
variations_map: dict[str, list[str]] = {}
for line in raw.strip().split("\n"):
if ":" not in line:
continue
name, _, variations_str = line.partition(":")
name = name.strip()
# Match to original competitor name (case-insensitive)
matched_name = None
for comp in competitors:
if comp.lower() == name.lower():
matched_name = comp
break
if matched_name is None:
continue
variations = [
v.strip() for v in variations_str.split(",") if v.strip()
]
# Filter out variations that are identical to the original
variations = [
v for v in variations if v.lower() != matched_name.lower()
]
variations_map[matched_name] = variations
return variations_map
return _parse_variations_response(raw, competitors)
except Exception as e:
verbose_proxy_logger.error("LLM competitor variation generation failed: %s", e)
return {}
async def _discover_competitors_via_llm(prompt: str, model: str = "gpt-4o-mini") -> list:
"""Call an onboarded LLM to discover competitor names."""
import litellm
def _parse_variations_response(raw: str, competitors: list) -> dict[str, list[str]]:
"""Parse the LLM response for competitor variations into a name -> variations map."""
# Build a lowercase lookup for case-insensitive matching
lower_to_canonical = {comp.lower(): comp for comp in competitors}
variations_map: dict[str, list[str]] = {}
for line in raw.strip().split("\n"):
if ":" not in line:
continue
name, _, variations_str = line.partition(":")
canonical = lower_to_canonical.get(name.strip().lower())
if canonical is None:
continue
variations = [
v.strip()
for v in variations_str.split(",")
if v.strip() and v.strip().lower() != canonical.lower()
]
variations_map[canonical] = variations
return variations_map
async def _discover_competitors_via_llm(
prompt: str, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL
) -> list:
"""Call an onboarded LLM to discover competitor names."""
try:
response = await litellm.acompletion(
from litellm.proxy.proxy_server import llm_router
if llm_router is None:
raise ValueError("LLM router not initialized")
response = await llm_router.acompletion(
model=model,
messages=[{"role": "user", "content": prompt}],
temperature=0.3,
temperature=COMPETITOR_LLM_TEMPERATURE,
)
raw = response.choices[0].message.content or "" # type: ignore
competitors = [
line.strip().strip(".-) ").strip()
name
for line in raw.strip().split("\n")
if line.strip() and len(line.strip()) > 1
if (name := _clean_competitor_line(line)) is not None
]
return competitors[:30]
return competitors[:MAX_COMPETITOR_NAMES]
except Exception as e:
verbose_proxy_logger.error("LLM competitor discovery failed: %s", e)
return []
def _build_all_names_per_competitor(
competitors: list[str], variations_map: dict[str, list[str]]
) -> dict[str, list[str]]:
"""Build canonical + variation name lists for each competitor."""
return {
comp: [comp] + variations_map.get(comp, [])
for comp in competitors
}
def _build_competitor_guardrail_definitions(
definitions: list,
competitors: list,
@ -767,46 +793,13 @@ def _build_competitor_guardrail_definitions(
variations_map: Optional[dict] = None,
) -> list:
"""Build enriched guardrailDefinitions with competitor names and variations populated."""
import copy
variations_map = variations_map or {}
enriched = copy.deepcopy(definitions)
all_names = _build_all_names_per_competitor(competitors, variations_map)
# Build list of all names (canonical + variations) for each competitor
all_names_per_competitor: dict[str, list[str]] = {}
for comp in competitors:
names = [comp]
names.extend(variations_map.get(comp, []))
all_names_per_competitor[comp] = names
output_blocked = []
for comp in competitors:
for name in all_names_per_competitor[comp]:
desc = f"Competitor: {comp}" if name == comp else f"Competitor variation ({comp}): {name}"
output_blocked.append(
{"keyword": name, "action": "BLOCK", "description": desc}
)
recommendation_blocked = []
for comp in competitors:
for name in all_names_per_competitor[comp]:
for prefix in ["try", "use", "switch to", "consider"]:
recommendation_blocked.append(
{"keyword": f"{prefix} {name}", "action": "BLOCK", "description": f"Recommendation to competitor ({comp})"}
)
comparison_blocked = []
for comp in competitors:
for name in all_names_per_competitor[comp]:
comparison_blocked.append(
{"keyword": f"{name} is better", "action": "BLOCK", "description": f"Unfavorable comparison ({comp})"}
)
comparison_blocked.append(
{"keyword": f"better than {brand_name}", "action": "BLOCK", "description": "Unfavorable comparison"}
)
comparison_blocked.append(
{"keyword": f"{brand_name} is worse", "action": "BLOCK", "description": "Unfavorable comparison"}
)
output_blocked = _build_name_blocked_words(competitors, all_names)
recommendation_blocked = _build_recommendation_blocked_words(competitors, all_names)
comparison_blocked = _build_comparison_blocked_words(competitors, all_names, brand_name)
blocked_words_map = {
"competitor-output-blocker": output_blocked,
@ -828,3 +821,59 @@ def _build_competitor_guardrail_definitions(
defn["litellm_params"]["blocked_words"] = blocked_words_map[guardrail_name]
return enriched
def _build_name_blocked_words(
competitors: list[str], all_names: dict[str, list[str]]
) -> list[dict]:
"""Build blocked word entries for direct competitor name mentions."""
result = []
for comp in competitors:
for name in all_names[comp]:
desc = f"Competitor: {comp}" if name == comp else f"Competitor variation ({comp}): {name}"
result.append({"keyword": name, "action": "BLOCK", "description": desc})
return result
def _build_recommendation_blocked_words(
competitors: list[str], all_names: dict[str, list[str]]
) -> list[dict]:
"""Build blocked word entries for competitor recommendations."""
result = []
for comp in competitors:
for name in all_names[comp]:
for prefix in ["try", "use", "switch to", "consider"]:
result.append({
"keyword": f"{prefix} {name}",
"action": "BLOCK",
"description": f"Recommendation to competitor ({comp})",
})
return result
def _build_comparison_blocked_words(
competitors: list[str], all_names: dict[str, list[str]], brand_name: str
) -> list[dict]:
"""Build blocked word entries for unfavorable competitor comparisons."""
result = []
for comp in competitors:
for name in all_names[comp]:
result.append({
"keyword": f"{name} is better",
"action": "BLOCK",
"description": f"Unfavorable comparison ({comp})",
})
# Brand-level comparisons (only need one entry each, not per-competitor)
result.append({
"keyword": f"better than {brand_name}",
"action": "BLOCK",
"description": "Unfavorable comparison",
})
result.append({
"keyword": f"{brand_name} is worse",
"action": "BLOCK",
"description": "Unfavorable comparison",
})
return result