diff --git a/litellm/constants.py b/litellm/constants.py index 17ad742e419..472fd27983c 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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" diff --git a/litellm/proxy/management_endpoints/policy_endpoints.py b/litellm/proxy/management_endpoints/policy_endpoints.py index 4c45cd4193b..8699f41227a 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints.py +++ b/litellm/proxy/management_endpoints/policy_endpoints.py @@ -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