diff --git a/.circleci/config.yml b/.circleci/config.yml index 46d2ea2c6e..540ebbbc32 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1388,6 +1388,7 @@ jobs: - run: python ./tests/documentation_tests/test_circular_imports.py - run: python ./tests/code_coverage_tests/prevent_key_leaks_in_exceptions.py - run: python ./tests/code_coverage_tests/check_unsafe_enterprise_import.py + - run: python ./tests/code_coverage_tests/ban_copy_deepcopy_kwargs.py - run: helm lint ./deploy/charts/litellm-helm db_migration_disable_update_check: diff --git a/deploy/charts/litellm-helm/README.md b/deploy/charts/litellm-helm/README.md index 31bda3f7d7..cef2b8d162 100644 --- a/deploy/charts/litellm-helm/README.md +++ b/deploy/charts/litellm-helm/README.md @@ -110,6 +110,22 @@ data: Source: [GitHub Gist from troyharvey](https://gist.github.com/troyharvey/4506472732157221e04c6b15e3b3f094) +### Migration Job Settings + +The migration job supports both ArgoCD and Helm hooks to ensure database migrations run at the appropriate time during deployments. + +| Name | Description | Value | +| ---------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | ----- | +| `migrationJob.enabled` | Enable or disable the schema migration Job | `true` | +| `migrationJob.backoffLimit` | Backoff limit for Job restarts | `4` | +| `migrationJob.ttlSecondsAfterFinished` | TTL for completed migration jobs | `120` | +| `migrationJob.annotations` | Additional annotations for the migration job pod | `{}` | +| `migrationJob.extraContainers` | Additional containers to run alongside the migration job | `[]` | +| `migrationJob.hooks.argocd.enabled` | Enable ArgoCD hooks for the migration job (uses PreSync hook with BeforeHookCreation delete policy) | `true` | +| `migrationJob.hooks.helm.enabled` | Enable Helm hooks for the migration job (uses pre-install,pre-upgrade hooks with before-hook-creation delete policy) | `false` | +| `migrationJob.hooks.helm.weight` | Helm hook execution order (lower weights executed first). Optional - defaults to "1" if not specified. | N/A | + + ## Accessing the Admin UI When browsing to the URL published per the settings in `ingress.*`, you will be prompted for **Admin Configuration**. The **Proxy Endpoint** is the internal diff --git a/deploy/charts/litellm-helm/templates/migrations-job.yaml b/deploy/charts/litellm-helm/templates/migrations-job.yaml index 4917678831..143e62fceb 100644 --- a/deploy/charts/litellm-helm/templates/migrations-job.yaml +++ b/deploy/charts/litellm-helm/templates/migrations-job.yaml @@ -5,8 +5,15 @@ kind: Job metadata: name: {{ include "litellm.fullname" . }}-migrations annotations: + {{- if .Values.migrationJob.hooks.argocd.enabled }} argocd.argoproj.io/hook: PreSync - argocd.argoproj.io/hook-delete-policy: BeforeHookCreation # delete old migration on a new deploy in case the migration needs to make updates + argocd.argoproj.io/hook-delete-policy: BeforeHookCreation + {{- end }} + {{- if .Values.migrationJob.hooks.helm.enabled }} + helm.sh/hook: "pre-install,pre-upgrade" + helm.sh/hook-delete-policy: "before-hook-creation" + helm.sh/hook-weight: {{ .Values.migrationJob.hooks.helm.weight | default "1" | quote }} + {{- end }} checksum/config: {{ toYaml .Values | sha256sum }} spec: template: diff --git a/deploy/charts/litellm-helm/values.yaml b/deploy/charts/litellm-helm/values.yaml index 0c00d2325a..5324f6de76 100644 --- a/deploy/charts/litellm-helm/values.yaml +++ b/deploy/charts/litellm-helm/values.yaml @@ -201,6 +201,13 @@ migrationJob: annotations: {} ttlSecondsAfterFinished: 120 extraContainers: [] + + # Hook configuration + hooks: + argocd: + enabled: true + helm: + enabled: false # Additional environment variables to be added to the deployment as a map of key-value pairs envVars: { diff --git a/docker-compose.yml b/docker-compose.yml index 189dde5e3e..366fbe51b5 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,67 +1,66 @@ -services: - litellm: - build: - context: . - args: - target: runtime - image: ghcr.io/berriai/litellm:main-stable - ######################################### - ## Uncomment these lines to start proxy with a config.yaml file ## - # volumes: - # - ./config.yaml:/app/config.yaml <<- this is missing in the docker-compose file currently - # command: - # - "--config=/app/config.yaml" - ############################################## - ports: - - "4000:4000" # Map the container port to the host, change the host port if necessary - environment: - DATABASE_URL: "postgresql://llmproxy:dbpassword9090@db:5432/litellm" - STORE_MODEL_IN_DB: "True" # allows adding models to proxy via UI - env_file: - - .env # Load local .env file - depends_on: - - db # Indicates that this service depends on the 'db' service, ensuring 'db' starts first - healthcheck: # Defines the health check configuration for the container - test: [ "CMD-SHELL", "wget --no-verbose --tries=1 http://localhost:4000/health/liveliness || exit 1" ] # Command to execute for health check - interval: 30s # Perform health check every 30 seconds - timeout: 10s # Health check command times out after 10 seconds - retries: 3 # Retry up to 3 times if health check fails - start_period: 40s # Wait 40 seconds after container start before beginning health checks - - db: - image: postgres:16 - restart: always - container_name: litellm_db - environment: - POSTGRES_DB: litellm - POSTGRES_USER: llmproxy - POSTGRES_PASSWORD: dbpassword9090 - ports: - - "5432:5432" - volumes: - - postgres_data:/var/lib/postgresql/data # Persists Postgres data across container restarts - healthcheck: - test: ["CMD-SHELL", "pg_isready -d litellm -U llmproxy"] - interval: 1s - timeout: 5s - retries: 10 - - prometheus: - image: prom/prometheus - volumes: - - prometheus_data:/prometheus - - ./prometheus.yml:/etc/prometheus/prometheus.yml - ports: - - "9090:9090" - command: - - "--config.file=/etc/prometheus/prometheus.yml" - - "--storage.tsdb.path=/prometheus" - - "--storage.tsdb.retention.time=15d" - restart: always - -volumes: - prometheus_data: - driver: local - postgres_data: - name: litellm_postgres_data # Named volume for Postgres data persistence - +services: + litellm: + build: + context: . + args: + target: runtime + image: ghcr.io/berriai/litellm:main-stable + ######################################### + ## Uncomment these lines to start proxy with a config.yaml file ## + # volumes: + # - ./config.yaml:/app/config.yaml <<- this is missing in the docker-compose file currently + # command: + # - "--config=/app/config.yaml" + ############################################## + ports: + - "4000:4000" # Map the container port to the host, change the host port if necessary + environment: + DATABASE_URL: "postgresql://llmproxy:dbpassword9090@db:5432/litellm" + STORE_MODEL_IN_DB: "True" # allows adding models to proxy via UI + env_file: + - .env # Load local .env file + depends_on: + - db # Indicates that this service depends on the 'db' service, ensuring 'db' starts first + healthcheck: # Defines the health check configuration for the container + test: [ "CMD-SHELL", "wget --no-verbose --tries=1 http://localhost:4000/health/liveliness || exit 1" ] # Command to execute for health check + interval: 30s # Perform health check every 30 seconds + timeout: 10s # Health check command times out after 10 seconds + retries: 3 # Retry up to 3 times if health check fails + start_period: 40s # Wait 40 seconds after container start before beginning health checks + + db: + image: postgres:16 + restart: always + container_name: litellm_db + environment: + POSTGRES_DB: litellm + POSTGRES_USER: llmproxy + POSTGRES_PASSWORD: dbpassword9090 + ports: + - "5432:5432" + volumes: + - postgres_data:/var/lib/postgresql/data # Persists Postgres data across container restarts + healthcheck: + test: ["CMD-SHELL", "pg_isready -d litellm -U llmproxy"] + interval: 1s + timeout: 5s + retries: 10 + + prometheus: + image: prom/prometheus + volumes: + - prometheus_data:/prometheus + - ./prometheus.yml:/etc/prometheus/prometheus.yml + ports: + - "9090:9090" + command: + - "--config.file=/etc/prometheus/prometheus.yml" + - "--storage.tsdb.path=/prometheus" + - "--storage.tsdb.retention.time=15d" + restart: always + +volumes: + prometheus_data: + driver: local + postgres_data: + name: litellm_postgres_data # Named volume for Postgres data persistence diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index d4e672251e..3f204908f3 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -33,7 +33,7 @@ WORKDIR /app # Install runtime dependencies USER root RUN apk upgrade --no-cache && \ - apk add --no-cache bash + apk add --no-cache bash libstdc++ ca-certificates openssl # Copy only necessary artifacts from builder stage for runtime COPY --from=builder /app/docker/entrypoint.sh /app/docker/prod_entrypoint.sh /app/docker/ @@ -86,4 +86,4 @@ ENTRYPOINT ["/app/docker/prod_entrypoint.sh"] # Append "--detailed_debug" to the end of CMD to view detailed debug logs # CMD ["--port", "4000", "--detailed_debug"] -CMD ["--port", "4000"] +CMD ["--port", "4000"] \ No newline at end of file diff --git a/docker/README.md b/docker/README.md index 8dbc59d01b..1c3c208988 100644 --- a/docker/README.md +++ b/docker/README.md @@ -1,3 +1,65 @@ -# LiteLLM Docker +# Docker Development Guide -This is a minimal Docker Compose setup for self-hosting LiteLLM. \ No newline at end of file +This guide provides instructions for building and running the LiteLLM application using Docker and Docker Compose. + +## Prerequisites + +- Docker +- Docker Compose + +## Building and Running the Application + +To build and run the application, you will use the `docker-compose.yml` file located in the root of the project. This file is configured to use the `Dockerfile.non_root` for a secure, non-root container environment. + +### 1. Set the Master Key + +The application requires a `MASTER_KEY` for signing and validating tokens. You must set this key as an environment variable before running the application. + +Create a `.env` file in the root of the project and add the following line: + +``` +MASTER_KEY=your-secret-key +``` + +Replace `your-secret-key` with a strong, randomly generated secret. + +### 2. Build and Run the Containers + +Once you have set the `MASTER_KEY`, you can build and run the containers using the following command: + +```bash +docker-compose up -d --build +``` + +This command will: + +- Build the Docker image using `Dockerfile.non_root`. +- Start the `litellm`, `litellm_db`, and `prometheus` services in detached mode (`-d`). +- The `--build` flag ensures that the image is rebuilt if there are any changes to the Dockerfile or the application code. + +### 3. Verifying the Application is Running + +You can check the status of the running containers with the following command: + +```bash +docker-compose ps +``` + +To view the logs of the `litellm` container, run: + +```bash +docker-compose logs -f litellm +``` + +### 4. Stopping the Application + +To stop the running containers, use the following command: + +```bash +docker-compose down +``` + +## Troubleshooting + +- **`build_admin_ui.sh: not found`**: This error can occur if the Docker build context is not set correctly. Ensure that you are running the `docker-compose` command from the root of the project. +- **`Master key is not initialized`**: This error means the `MASTER_key` environment variable is not set. Make sure you have created a `.env` file in the project root with the `MASTER_KEY` defined. diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 1243d7232a..24771ff63c 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -327,6 +327,7 @@ router_settings: | ATHINA_BASE_URL | Base URL for Athina service (defaults to `https://log.athina.ai`) | AUTH_STRATEGY | Strategy used for authentication (e.g., OAuth, API key) | ANTHROPIC_API_KEY | API key for Anthropic service +| ANTHROPIC_API_BASE | Base URL for Anthropic API. Default is https://api.anthropic.com | AWS_ACCESS_KEY_ID | Access Key ID for AWS services | AWS_PROFILE_NAME | AWS CLI profile name to be used | AWS_REGION_NAME | Default AWS region for service interactions @@ -372,6 +373,7 @@ router_settings: | CONFIDENT_API_KEY | API key for DeepEval integration | CUSTOM_TIKTOKEN_CACHE_DIR | Custom directory for Tiktoken cache | CONFIDENT_API_KEY | API key for Confident AI (Deepeval) Logging service +| COHERE_API_BASE | Base URL for Cohere API. Default is https://api.cohere.com | DATABASE_HOST | Hostname for the database server | DATABASE_NAME | Name of the database | DATABASE_PASSWORD | Password for the database user @@ -482,6 +484,7 @@ router_settings: | GENERIC_USER_PROVIDER_ATTRIBUTE | Attribute specifying the user's provider | GENERIC_USER_ROLE_ATTRIBUTE | Attribute specifying the user's role | GENERIC_USERINFO_ENDPOINT | Endpoint to fetch user information in generic OAuth +| GEMINI_API_BASE | Base URL for Gemini API. Default is https://generativelanguage.googleapis.com | GALILEO_BASE_URL | Base URL for Galileo platform | GALILEO_PASSWORD | Password for Galileo authentication | GALILEO_PROJECT_ID | Project ID for Galileo usage @@ -581,7 +584,7 @@ router_settings: | MAX_LANGFUSE_INITIALIZED_CLIENTS | Maximum number of Langfuse clients to initialize on proxy. Default is 20. This is set since langfuse initializes 1 thread everytime a client is initialized. We've had an incident in the past where we reached 100% cpu utilization because Langfuse was initialized several times. | MIN_NON_ZERO_TEMPERATURE | Minimum non-zero temperature value. Default is 0.0001 | MINIMUM_PROMPT_CACHE_TOKEN_COUNT | Minimum token count for caching a prompt. Default is 1024 -| MISTRAL_API_BASE | Base URL for Mistral API +| MISTRAL_API_BASE | Base URL for Mistral API. Default is https://api.mistral.ai | MISTRAL_API_KEY | API key for Mistral API | MICROSOFT_CLIENT_ID | Client ID for Microsoft services | MICROSOFT_CLIENT_SECRET | Client secret for Microsoft services @@ -593,7 +596,7 @@ router_settings: | NON_LLM_CONNECTION_TIMEOUT | Timeout in seconds for non-LLM service connections. Default is 15 | OAUTH_TOKEN_INFO_ENDPOINT | Endpoint for OAuth token info retrieval | OPENAI_BASE_URL | Base URL for OpenAI API -| OPENAI_API_BASE | Base URL for OpenAI API +| OPENAI_API_BASE | Base URL for OpenAI API. Default is https://api.openai.com/ | OPENAI_API_KEY | API key for OpenAI services | OPENAI_FILE_SEARCH_COST_PER_1K_CALLS | Cost per 1000 calls for OpenAI file search. Default is 0.0025 | OPENAI_ORGANIZATION | Organization identifier for OpenAI diff --git a/litellm/__init__.py b/litellm/__init__.py index 8eaf2e1d4a..8400668bed 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -61,6 +61,11 @@ from litellm.constants import ( DEFAULT_SOFT_BUDGET, DEFAULT_ALLOWED_FAILS, ) +from litellm.integrations.dotprompt import ( + global_prompt_manager, + global_prompt_directory, + set_global_prompt_directory, +) from litellm.types.guardrails import GuardrailItem from litellm.types.secret_managers.main import ( KeyManagementSystem, @@ -83,7 +88,6 @@ if litellm_mode == "DEV": # Register async client cleanup to prevent resource leaks register_async_client_cleanup() - #################################################### if set_verbose == True: _turn_on_debug() @@ -130,6 +134,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "s3_v2", "aws_sqs", "vector_store_pre_call_hook", + "dotprompt", ] logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None _known_custom_logger_compatible_callbacks: List = list( @@ -145,22 +150,22 @@ prometheus_initialize_budget_metrics: Optional[bool] = False require_auth_for_metrics_endpoint: Optional[bool] = False argilla_batch_size: Optional[int] = None datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload. -gcs_pub_sub_use_v1: Optional[ - bool -] = False # if you want to use v1 gcs pubsub logged payload -generic_api_use_v1: Optional[ - bool -] = False # if you want to use v1 generic api logged payload +gcs_pub_sub_use_v1: Optional[bool] = ( + False # if you want to use v1 gcs pubsub logged payload +) +generic_api_use_v1: Optional[bool] = ( + False # if you want to use v1 generic api logged payload +) argilla_transformation_object: Optional[Dict[str, Any]] = None -_async_input_callback: List[ - Union[str, Callable, CustomLogger] -] = [] # internal variable - async custom callbacks are routed here. -_async_success_callback: List[ - Union[str, Callable, CustomLogger] -] = [] # internal variable - async custom callbacks are routed here. -_async_failure_callback: List[ - Union[str, Callable, CustomLogger] -] = [] # internal variable - async custom callbacks are routed here. +_async_input_callback: List[Union[str, Callable, CustomLogger]] = ( + [] +) # internal variable - async custom callbacks are routed here. +_async_success_callback: List[Union[str, Callable, CustomLogger]] = ( + [] +) # internal variable - async custom callbacks are routed here. +_async_failure_callback: List[Union[str, Callable, CustomLogger]] = ( + [] +) # internal variable - async custom callbacks are routed here. pre_call_rules: List[Callable] = [] post_call_rules: List[Callable] = [] turn_off_message_logging: Optional[bool] = False @@ -168,18 +173,18 @@ log_raw_request_response: bool = False redact_messages_in_exceptions: Optional[bool] = False redact_user_api_key_info: Optional[bool] = False filter_invalid_headers: Optional[bool] = False -add_user_information_to_llm_headers: Optional[ - bool -] = None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers +add_user_information_to_llm_headers: Optional[bool] = ( + None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers +) store_audit_logs = False # Enterprise feature, allow users to see audit logs ### end of callbacks ############# -email: Optional[ - str -] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -token: Optional[ - str -] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +email: Optional[str] = ( + None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) +token: Optional[str] = ( + None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) telemetry = True max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False)) @@ -268,11 +273,15 @@ enable_loadbalancing_on_batch_endpoints: Optional[bool] = None enable_caching_on_provider_specific_optional_params: bool = ( False # feature-flag for caching on optional params - e.g. 'top_k' ) -caching: bool = False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -caching_with_models: bool = False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -cache: Optional[ - Cache -] = None # cache object <- use this - https://docs.litellm.ai/docs/caching +caching: bool = ( + False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) +caching_with_models: bool = ( + False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) +cache: Optional[Cache] = ( + None # cache object <- use this - https://docs.litellm.ai/docs/caching +) default_in_memory_ttl: Optional[float] = None default_redis_ttl: Optional[float] = None default_redis_batch_cache_expiry: Optional[float] = None @@ -280,9 +289,9 @@ model_alias_map: Dict[str, str] = {} model_group_alias_map: Dict[str, str] = {} model_group_settings: Optional["ModelGroupSettings"] = None max_budget: float = 0.0 # set the max budget across all providers -budget_duration: Optional[ - str -] = None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). +budget_duration: Optional[str] = ( + None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). +) default_soft_budget: float = ( DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0 ) @@ -291,11 +300,15 @@ forward_traceparent_to_llm_provider: bool = False _current_cost = 0.0 # private variable, used if max budget is set error_logs: Dict = {} -add_function_to_prompt: bool = False # if function calling not supported by api, append function call details to system prompt +add_function_to_prompt: bool = ( + False # if function calling not supported by api, append function call details to system prompt +) client_session: Optional[httpx.Client] = None aclient_session: Optional[httpx.AsyncClient] = None model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks' -model_cost_map_url: str = "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json" +model_cost_map_url: str = ( + "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json" +) suppress_debug_info = False dynamodb_table_name: Optional[str] = None s3_callback_params: Optional[Dict] = None @@ -324,7 +337,9 @@ prometheus_metrics_config: Optional[List] = None disable_add_prefix_to_prompt: bool = ( False # used by anthropic, to disable adding prefix to prompt ) -disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. +disable_copilot_system_to_assistant: bool = ( + False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. +) public_model_groups: Optional[List[str]] = None public_model_groups_links: Dict[str, str] = {} #### REQUEST PRIORITIZATION ##### @@ -332,13 +347,17 @@ priority_reservation: Optional[Dict[str, float]] = None ######## Networking Settings ######## -use_aiohttp_transport: bool = True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead. +use_aiohttp_transport: bool = ( + True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead. +) aiohttp_trust_env: bool = False # set to true to use HTTP_ Proxy settings disable_aiohttp_transport: bool = False # Set this to true to use httpx instead disable_aiohttp_trust_env: bool = ( False # When False, aiohttp will respect HTTP(S)_PROXY env vars ) -force_ipv4: bool = False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6. +force_ipv4: bool = ( + False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6. +) module_level_aclient = AsyncHTTPHandler( timeout=request_timeout, client_alias="module level aclient" ) @@ -352,13 +371,13 @@ fallbacks: Optional[List] = None context_window_fallbacks: Optional[List] = None content_policy_fallbacks: Optional[List] = None allowed_fails: int = 3 -num_retries_per_request: Optional[ - int -] = None # for the request overall (incl. fallbacks + model retries) +num_retries_per_request: Optional[int] = ( + None # for the request overall (incl. fallbacks + model retries) +) ####### SECRET MANAGERS ##################### -secret_manager_client: Optional[ - Any -] = None # list of instantiated key management clients - e.g. azure kv, infisical, etc. +secret_manager_client: Optional[Any] = ( + None # list of instantiated key management clients - e.g. azure kv, infisical, etc. +) _google_kms_resource_name: Optional[str] = None _key_management_system: Optional[KeyManagementSystem] = None _key_management_settings: KeyManagementSettings = KeyManagementSettings() @@ -498,6 +517,7 @@ lambda_ai_models: List = [] hyperbolic_models: List = [] recraft_models: List = [] + def is_bedrock_pricing_only_model(key: str) -> bool: """ Excludes keys with the pattern 'bedrock//'. These are in the model_prices_and_context_window.json file for pricing purposes only. @@ -1232,12 +1252,12 @@ from .types.llms.custom_llm import CustomLLMItem from .types.utils import GenericStreamingChunk custom_provider_map: List[CustomLLMItem] = [] -_custom_providers: List[ - str -] = [] # internal helper util, used to track names of custom providers -disable_hf_tokenizer_download: Optional[ - bool -] = None # disable huggingface tokenizer download. Defaults to openai clk100 +_custom_providers: List[str] = ( + [] +) # internal helper util, used to track names of custom providers +disable_hf_tokenizer_download: Optional[bool] = ( + None # disable huggingface tokenizer download. Defaults to openai clk100 +) global_disable_no_log_param: bool = False ### PASSTHROUGH ### diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index c8892cd26a..6fecc7fa97 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -789,6 +789,7 @@ def completion_cost( # noqa: PLR0915 from litellm.llms.recraft.cost_calculator import ( cost_calculator as recraft_image_cost_calculator, ) + return recraft_image_cost_calculator( model=model, image_response=completion_response, @@ -797,6 +798,7 @@ def completion_cost( # noqa: PLR0915 from litellm.llms.gemini.image_generation.cost_calculator import ( cost_calculator as gemini_image_cost_calculator, ) + return gemini_image_cost_calculator( model=model, image_response=completion_response, @@ -867,7 +869,10 @@ def completion_cost( # noqa: PLR0915 from litellm.proxy._experimental.mcp_server.cost_calculator import ( MCPCostCalculator, ) - return MCPCostCalculator.calculate_mcp_tool_call_cost(litellm_logging_obj=litellm_logging_obj) + + return MCPCostCalculator.calculate_mcp_tool_call_cost( + litellm_logging_obj=litellm_logging_obj + ) # Calculate cost based on prompt_tokens, completion_tokens if ( "togethercomputer" in model @@ -1318,7 +1323,7 @@ class BaseTokenUsageProcessor: combined.completion_tokens_details = CompletionTokensDetails() # Check what keys exist in the model's completion_tokens_details - for attr in dir(usage.completion_tokens_details): + for attr in usage.completion_tokens_details.model_fields: if not attr.startswith("_") and not callable( getattr(usage.completion_tokens_details, attr) ): @@ -1326,7 +1331,8 @@ class BaseTokenUsageProcessor: combined.completion_tokens_details, attr, 0 ) new_val = getattr(usage.completion_tokens_details, attr, 0) - if new_val is not None: + + if new_val is not None and current_val is not None: setattr( combined.completion_tokens_details, attr, diff --git a/litellm/exceptions.py b/litellm/exceptions.py index 9f3411143a..153230518c 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -829,3 +829,65 @@ class BlockedPiiEntityError(Exception): self.guardrail_name = guardrail_name self.message = f"Blocked entity detected: {entity_type} by Guardrail: {guardrail_name}. This entity is not allowed to be used in this request." super().__init__(self.message) + + +class MidStreamFallbackError(ServiceUnavailableError): # type: ignore + def __init__( + self, + message: str, + model: str, + llm_provider: str, + original_exception: Optional[Exception] = None, + response: Optional[httpx.Response] = None, + litellm_debug_info: Optional[str] = None, + max_retries: Optional[int] = None, + num_retries: Optional[int] = None, + generated_content: str = "", + is_pre_first_chunk: bool = False, + ): + self.status_code = 503 # Service Unavailable + self.message = f"litellm.MidStreamFallbackError: {message}" + self.model = model + self.llm_provider = llm_provider + self.original_exception = original_exception + self.litellm_debug_info = litellm_debug_info + self.max_retries = max_retries + self.num_retries = num_retries + self.generated_content = generated_content + self.is_pre_first_chunk = is_pre_first_chunk + + # Create a response if one wasn't provided + if response is None: + self.response = httpx.Response( + status_code=self.status_code, + request=httpx.Request( + method="POST", + url=f"https://{llm_provider}.com/v1/", + ), + ) + else: + self.response = response + + # Call the parent constructor + super().__init__( + message=self.message, + llm_provider=llm_provider, + model=model, + response=self.response, + litellm_debug_info=self.litellm_debug_info, + max_retries=self.max_retries, + num_retries=self.num_retries, + ) + + def __str__(self): + _message = self.message + if self.num_retries: + _message += f" LiteLLM Retried: {self.num_retries} times" + if self.max_retries: + _message += f", LiteLLM Max Retries: {self.max_retries}" + if self.original_exception: + _message += f" Original exception: {type(self.original_exception).__name__}: {str(self.original_exception)}" + return _message + + def __repr__(self): + return self.__str__() diff --git a/litellm/integrations/dotprompt/README.md b/litellm/integrations/dotprompt/README.md new file mode 100644 index 0000000000..7eaeca9a33 --- /dev/null +++ b/litellm/integrations/dotprompt/README.md @@ -0,0 +1,316 @@ +# LiteLLM Dotprompt Manager + +A powerful prompt management system for LiteLLM that supports [Google's Dotprompt specification](https://google.github.io/dotprompt/getting-started/). This allows you to manage your AI prompts in organized `.prompt` files with YAML frontmatter, Handlebars templating, and full integration with LiteLLM's completion API. + +## Features + +- **📁 File-based prompt management**: Organize prompts in `.prompt` files +- **🎯 YAML frontmatter**: Define model, parameters, and schemas in file headers +- **🔧 Handlebars templating**: Use `{{variable}}` syntax with Jinja2 backend +- **✅ Input validation**: Automatic validation against defined schemas +- **🔗 LiteLLM integration**: Works seamlessly with `litellm.completion()` +- **💬 Smart message parsing**: Converts prompts to proper chat messages +- **⚙️ Parameter extraction**: Automatically applies model settings from prompts + +## Quick Start + +### 1. Create a `.prompt` file + +Create a file called `chat_assistant.prompt`: + +```yaml +--- +model: gpt-4 +temperature: 0.7 +max_tokens: 150 +input: + schema: + user_message: string + system_context?: string +--- + +{% if system_context %}System: {{system_context}} + +{% endif %}User: {{user_message}} +``` + +### 2. Use with LiteLLM + +```python +import litellm + +litellm.set_global_prompt_directory("path/to/your/prompts") + +# Use with completion - the model prefix 'dotprompt/' tells LiteLLM to use prompt management +response = litellm.completion( + model="dotprompt/gpt-4", # The actual model comes from the .prompt file + prompt_id="chat_assistant", + prompt_variables={ + "user_message": "What is machine learning?", + "system_context": "You are a helpful AI tutor." + }, + # Any additional messages will be appended after the prompt + messages=[{"role": "user", "content": "Please explain it simply."}] +) + +print(response.choices[0].message.content) +``` + +## Prompt File Format + +### Basic Structure + +```yaml +--- +# Model configuration +model: gpt-4 +temperature: 0.7 +max_tokens: 500 + +# Input schema (optional) +input: + schema: + name: string + age: integer + preferences?: array +--- + +# Template content using Handlebars syntax +Hello {{name}}! + +{% if age >= 18 %} +You're an adult, so here are some mature recommendations: +{% else %} +Here are some age-appropriate suggestions: +{% endif %} + +{% for pref in preferences %} +- Based on your interest in {{pref}}, I recommend... +{% endfor %} +``` + +### Supported Frontmatter Fields + +- **`model`**: The LLM model to use (e.g., `gpt-4`, `claude-3-sonnet`) +- **`input.schema`**: Define expected input variables and their types +- **`output.format`**: Expected output format (`json`, `text`, etc.) +- **`output.schema`**: Structure of expected output + +### Additional Parameters + +- **`temperature`**: Model temperature (0.0 to 1.0) +- **`max_tokens`**: Maximum tokens to generate +- **`top_p`**: Nucleus sampling parameter (0.0 to 1.0) +- **`frequency_penalty`**: Frequency penalty (0.0 to 1.0) +- **`presence_penalty`**: Presence penalty (0.0 to 1.0) +- any other parameters that are not model or schema-related will be treated as optional parameters to the model. + +### Input Schema Types + +- `string` or `str`: Text values +- `integer` or `int`: Whole numbers +- `float`: Decimal numbers +- `boolean` or `bool`: True/false values +- `array` or `list`: Lists of values +- `object` or `dict`: Key-value objects + +Use `?` suffix for optional fields: `name?: string` + +## Message Format Conversion + +The dotprompt manager intelligently converts your rendered prompts into proper chat messages: + +### Simple Text → User Message +```yaml +--- +model: gpt-4 +--- +Tell me about {{topic}}. +``` +Becomes: `[{"role": "user", "content": "Tell me about AI."}]` + +### Role-Based Format → Multiple Messages +```yaml +--- +model: gpt-4 +--- +System: You are a {{role}}. + +User: {{question}} +``` + +Becomes: +```python +[ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "What is AI?"} +] +``` + + +## Example Prompts + +### Data Extraction +```yaml +# extract_info.prompt +--- +model: gemini/gemini-1.5-pro +input: + schema: + text: string +output: + format: json + schema: + title?: string + summary: string + tags: array +--- + +Extract the requested information from the given text. Return JSON format. + +Text: {{text}} +``` + +### Code Assistant +```yaml +# code_helper.prompt +--- +model: claude-3-5-sonnet-20241022 +temperature: 0.2 +max_tokens: 2000 +input: + schema: + language: string + task: string + code?: string +--- + +You are an expert {{language}} programmer. + +Task: {{task}} + +{% if code %} +Current code: +```{{language}} +{{code}} +``` +{% endif %} + +Please provide a complete, well-documented solution. +``` + +### Multi-turn Conversation +```yaml +# conversation.prompt +--- +model: gpt-4 +temperature: 0.8 +input: + schema: + personality: string + context: string +--- + +System: You are a {{personality}}. {{context}} + +User: Let's start our conversation. +``` + +## API Reference + +### PromptManager + +The core class for managing `.prompt` files. + +#### Methods + +- **`__init__(prompt_directory: str)`**: Initialize with directory path +- **`render(prompt_id: str, variables: dict) -> str`**: Render prompt with variables +- **`list_prompts() -> List[str]`**: Get all available prompt IDs +- **`get_prompt(prompt_id: str) -> PromptTemplate`**: Get prompt template object +- **`get_prompt_metadata(prompt_id: str) -> dict`**: Get prompt metadata +- **`reload_prompts() -> None`**: Reload all prompts from directory +- **`add_prompt(prompt_id: str, content: str, metadata: dict)`**: Add prompt programmatically + +### DotpromptManager + +LiteLLM integration class extending `PromptManagementBase`. + +#### Methods + +- **`__init__(prompt_directory: str)`**: Initialize with directory path +- **`should_run_prompt_management(prompt_id: str, params: dict) -> bool`**: Check if prompt exists +- **`set_prompt_directory(directory: str)`**: Change prompt directory +- **`reload_prompts()`**: Reload prompts from directory + +### PromptTemplate + +Represents a single prompt with metadata. + +#### Properties + +- **`content: str`**: The prompt template content +- **`metadata: dict`**: Full metadata from frontmatter +- **`model: str`**: Specified model name +- **`temperature: float`**: Model temperature +- **`max_tokens: int`**: Token limit +- **`input_schema: dict`**: Input validation schema +- **`output_format: str`**: Expected output format +- **`output_schema: dict`**: Output structure schema + +## Best Practices + +1. **Organize by purpose**: Group related prompts in subdirectories +2. **Use descriptive names**: `extract_user_info.prompt` vs `prompt1.prompt` +3. **Define schemas**: Always specify input schemas for validation +4. **Version control**: Store `.prompt` files in git for change tracking +5. **Test prompts**: Use the test framework to validate prompt behavior +6. **Keep templates focused**: One prompt should do one thing well +7. **Use includes**: Break complex prompts into reusable components + +## Troubleshooting + +### Common Issues + +**Prompt not found**: Ensure the `.prompt` file exists and has correct extension +```python +# Check available prompts +from litellm.prompts import get_dotprompt_manager +manager = get_dotprompt_manager() +print(manager.prompt_manager.list_prompts()) +``` + +**Template errors**: Verify Handlebars syntax and variable names +```python +# Test rendering directly +manager.prompt_manager.render("my_prompt", {"test": "value"}) +``` + +**Model not working**: Check that model name in frontmatter is correct +```python +# Check prompt metadata +metadata = manager.prompt_manager.get_prompt_metadata("my_prompt") +print(metadata) +``` + +### Validation Errors + +Input validation failures show helpful error messages: +``` +ValueError: Invalid type for field 'age': expected int, got str +``` + +Make sure your variables match the defined schema types. + +## Contributing + +The LiteLLM Dotprompt manager follows the [Dotprompt specification](https://google.github.io/dotprompt/) for maximum compatibility. When contributing: + +1. Ensure compatibility with existing `.prompt` files +2. Add tests for new features +3. Update documentation +4. Follow the existing code style + +## License + +This prompt management system is part of LiteLLM and follows the same license terms. \ No newline at end of file diff --git a/litellm/integrations/dotprompt/__init__.py b/litellm/integrations/dotprompt/__init__.py new file mode 100644 index 0000000000..bbd8be8025 --- /dev/null +++ b/litellm/integrations/dotprompt/__init__.py @@ -0,0 +1,33 @@ +from typing import TYPE_CHECKING, Optional + +if TYPE_CHECKING: + from .prompt_manager import PromptManager, PromptTemplate + +from .dotprompt_manager import DotpromptManager + +# Global instances +global_prompt_directory: Optional[str] = None +global_prompt_manager: Optional["PromptManager"] = None + + +def set_global_prompt_directory(directory: str) -> None: + """ + Set the global prompt directory for dotprompt files. + + Args: + directory: Path to directory containing .prompt files + """ + import litellm + + litellm.global_prompt_directory = directory # type: ignore + + +# Export public API +__all__ = [ + "PromptManager", + "DotpromptManager", + "PromptTemplate", + "set_global_prompt_directory", + "global_prompt_directory", + "global_prompt_manager", +] diff --git a/litellm/integrations/dotprompt/dotprompt_manager.py b/litellm/integrations/dotprompt/dotprompt_manager.py new file mode 100644 index 0000000000..830e950832 --- /dev/null +++ b/litellm/integrations/dotprompt/dotprompt_manager.py @@ -0,0 +1,225 @@ +""" +Dotprompt manager that integrates with LiteLLM's prompt management system. +Builds on top of PromptManagementBase to provide .prompt file support. +""" + +from typing import List, Optional + +from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.prompt_management_base import ( + PromptManagementBase, + PromptManagementClient, +) +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import StandardCallbackDynamicParams + +from .prompt_manager import PromptManager, PromptTemplate + + +class DotpromptManager(PromptManagementBase, CustomLogger): + """ + Dotprompt manager that integrates with LiteLLM's prompt management system. + + This class enables using .prompt files with the litellm completion() function + by implementing the PromptManagementBase interface. + + Usage: + # Set global prompt directory + litellm.prompt_directory = "path/to/prompts" + + # Use with completion + response = litellm.completion( + model="dotprompt/gpt-4", + prompt_id="my_prompt", + prompt_variables={"variable": "value"}, + messages=[{"role": "user", "content": "This will be combined with the prompt"}] + ) + """ + + def __init__(self, prompt_directory: Optional[str] = None): + import litellm + + self.prompt_directory = prompt_directory or litellm.global_prompt_directory + + self._prompt_manager: Optional[PromptManager] = None + + @property + def integration_name(self) -> str: + """Integration name used in model names like 'dotprompt/gpt-4'.""" + return "dotprompt" + + @property + def prompt_manager(self) -> PromptManager: + """Lazy-load the prompt manager.""" + if self._prompt_manager is None: + if self.prompt_directory is None: + raise ValueError( + "prompt_directory must be set before using dotprompt manager. " + "Set litellm.global_prompt_directory or initialize with prompt_directory parameter." + ) + self._prompt_manager = PromptManager(self.prompt_directory) + return self._prompt_manager + + def should_run_prompt_management( + self, + prompt_id: str, + dynamic_callback_params: StandardCallbackDynamicParams, + ) -> bool: + """ + Determine if prompt management should run based on the prompt_id. + + Returns True if the prompt_id exists in our prompt manager. + """ + try: + return prompt_id in self.prompt_manager.list_prompts() + except Exception: + # If there's any error accessing prompts, don't run prompt management + return False + + def _compile_prompt_helper( + self, + prompt_id: str, + prompt_variables: Optional[dict], + dynamic_callback_params: StandardCallbackDynamicParams, + prompt_label: Optional[str] = None, + prompt_version: Optional[int] = None, + ) -> PromptManagementClient: + """ + Compile a .prompt file into a PromptManagementClient structure. + + This method: + 1. Loads the prompt template from the .prompt file + 2. Renders it with the provided variables + 3. Converts the rendered text into chat messages + 4. Extracts model and optional parameters from metadata + """ + try: + # Get the prompt template + template = self.prompt_manager.get_prompt(prompt_id) + if template is None: + raise ValueError(f"Prompt '{prompt_id}' not found in prompt directory") + + # Render the template with variables + rendered_content = self.prompt_manager.render(prompt_id, prompt_variables) + + # Convert rendered content to chat messages + messages = self._convert_to_messages(rendered_content) + + # Extract model from metadata (if specified) + template_model = template.model + + # Extract optional parameters from metadata + optional_params = self._extract_optional_params(template) + + return PromptManagementClient( + prompt_id=prompt_id, + prompt_template=messages, + prompt_template_model=template_model, + prompt_template_optional_params=optional_params, + completed_messages=None, + ) + + except Exception as e: + raise ValueError(f"Error compiling prompt '{prompt_id}': {e}") + + def _convert_to_messages(self, rendered_content: str) -> List[AllMessageValues]: + """ + Convert rendered prompt content to chat messages. + + This method supports multiple formats: + 1. Simple text -> converted to user message + 2. Text with role prefixes (System:, User:, Assistant:) -> parsed into separate messages + 3. Already formatted as a single message + """ + # Clean up the content + content = rendered_content.strip() + + # Try to parse role-based format (System: ..., User: ..., etc.) + messages = [] + current_role = None + current_content = [] + + lines = content.split("\n") + + for line in lines: + line = line.strip() + + # Check for role prefixes + if line.startswith("System:"): + if current_role and current_content: + messages.append( + self._create_message( + current_role, "\n".join(current_content).strip() + ) + ) + current_role = "system" + current_content = [line[7:].strip()] # Remove "System:" prefix + elif line.startswith("User:"): + if current_role and current_content: + messages.append( + self._create_message( + current_role, "\n".join(current_content).strip() + ) + ) + current_role = "user" + current_content = [line[5:].strip()] # Remove "User:" prefix + elif line.startswith("Assistant:"): + if current_role and current_content: + messages.append( + self._create_message( + current_role, "\n".join(current_content).strip() + ) + ) + current_role = "assistant" + current_content = [line[10:].strip()] # Remove "Assistant:" prefix + else: + # Continue current message content + if current_role: + current_content.append(line) + else: + # No role prefix found, treat as user message + current_role = "user" + current_content = [line] + + # Add the last message + if current_role and current_content: + content_text = "\n".join(current_content).strip() + if content_text: # Only add if there's actual content + messages.append(self._create_message(current_role, content_text)) + + # If no messages were created, treat the entire content as a user message + if not messages and content: + messages.append(self._create_message("user", content)) + + return messages + + def _create_message(self, role: str, content: str) -> AllMessageValues: + """Create a message with the specified role and content.""" + return { + "role": role, # type: ignore + "content": content, + } + + def _extract_optional_params(self, template: PromptTemplate) -> dict: + """ + Extract optional parameters from the prompt template metadata. + + Includes parameters like temperature, max_tokens, etc. + """ + optional_params = {} + + # Extract common parameters from metadata + if template.optional_params is not None: + optional_params.update(template.optional_params) + + return optional_params + + def set_prompt_directory(self, prompt_directory: str) -> None: + """Set the prompt directory and reload prompts.""" + self.prompt_directory = prompt_directory + self._prompt_manager = None # Reset to force reload + + def reload_prompts(self) -> None: + """Reload all prompts from the directory.""" + if self._prompt_manager: + self._prompt_manager.reload_prompts() diff --git a/litellm/integrations/dotprompt/prompt_manager.py b/litellm/integrations/dotprompt/prompt_manager.py new file mode 100644 index 0000000000..c8bfd6e68b --- /dev/null +++ b/litellm/integrations/dotprompt/prompt_manager.py @@ -0,0 +1,220 @@ +""" +Based on Google's GenAI Kit dotprompt implementation: https://google.github.io/dotprompt/reference/frontmatter/ +""" + +import re +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple, Union + +import yaml +from jinja2 import DictLoader, Environment, select_autoescape + + +class PromptTemplate: + """Represents a single prompt template with metadata and content.""" + + def __init__( + self, + content: str, + metadata: Optional[Dict[str, Any]] = None, + template_id: Optional[str] = None, + ): + self.content = content + self.metadata = metadata or {} + self.template_id = template_id + + # Extract common metadata fields + restricted_keys = ["model", "input", "output"] + self.model = self.metadata.get("model") + self.input_schema = self.metadata.get("input", {}).get("schema", {}) + self.output_format = self.metadata.get("output", {}).get("format") + self.output_schema = self.metadata.get("output", {}).get("schema", {}) + self.optional_params = {} + for key in self.metadata.keys(): + if key not in restricted_keys: + self.optional_params[key] = self.metadata[key] + + def __repr__(self): + return f"PromptTemplate(id='{self.template_id}', model='{self.model}')" + + +class PromptManager: + """ + Manager for loading and rendering .prompt files following the Dotprompt specification. + + Supports: + - YAML frontmatter for metadata + - Handlebars-style templating (using Jinja2) + - Input/output schema validation + - Model configuration + """ + + def __init__(self, prompt_directory: str): + self.prompt_directory = Path(prompt_directory) + self.prompts: Dict[str, PromptTemplate] = {} + self.jinja_env = Environment( + loader=DictLoader({}), + autoescape=select_autoescape(["html", "xml"]), + # Use Handlebars-style delimiters to match Dotprompt spec + variable_start_string="{{", + variable_end_string="}}", + block_start_string="{%", + block_end_string="%}", + comment_start_string="{#", + comment_end_string="#}", + ) + + # Load all prompts in the directory + self._load_prompts() + + def _load_prompts(self) -> None: + """Load all .prompt files from the prompt directory.""" + if not self.prompt_directory.exists(): + raise ValueError( + f"Prompt directory does not exist: {self.prompt_directory}" + ) + + prompt_files = list(self.prompt_directory.glob("*.prompt")) + + for prompt_file in prompt_files: + try: + prompt_id = prompt_file.stem # filename without extension + template = self._load_prompt_file(prompt_file, prompt_id) + self.prompts[prompt_id] = template + # Optional: print(f"Loaded prompt: {prompt_id}") + except Exception: + # Optional: print(f"Error loading prompt file {prompt_file}") + pass + + def _load_prompt_file(self, file_path: Path, prompt_id: str) -> PromptTemplate: + """Load and parse a single .prompt file.""" + content = file_path.read_text(encoding="utf-8") + + # Split frontmatter and content + frontmatter, template_content = self._parse_frontmatter(content) + + return PromptTemplate( + content=template_content.strip(), + metadata=frontmatter, + template_id=prompt_id, + ) + + def _parse_frontmatter(self, content: str) -> Tuple[Dict[str, Any], str]: + """Parse YAML frontmatter from prompt content.""" + # Match YAML frontmatter between --- delimiters + frontmatter_pattern = r"^---\s*\n(.*?)\n---\s*\n(.*)$" + match = re.match(frontmatter_pattern, content, re.DOTALL) + + if match: + frontmatter_yaml = match.group(1) + template_content = match.group(2) + + try: + frontmatter = yaml.safe_load(frontmatter_yaml) or {} + except yaml.YAMLError as e: + raise ValueError(f"Invalid YAML frontmatter: {e}") + else: + # No frontmatter found, treat entire content as template + frontmatter = {} + template_content = content + + return frontmatter, template_content + + def render( + self, prompt_id: str, prompt_variables: Optional[Dict[str, Any]] = None + ) -> str: + """ + Render a prompt template with the given variables. + + Args: + prompt_id: The ID of the prompt template to render + prompt_variables: Variables to substitute in the template + + Returns: + The rendered prompt string + + Raises: + KeyError: If prompt_id is not found + ValueError: If template rendering fails + """ + if prompt_id not in self.prompts: + available_prompts = list(self.prompts.keys()) + raise KeyError( + f"Prompt '{prompt_id}' not found. Available prompts: {available_prompts}" + ) + + template = self.prompts[prompt_id] + variables = prompt_variables or {} + + # Validate input variables against schema if defined + if template.input_schema: + self._validate_input(variables, template.input_schema) + + try: + # Create Jinja2 template and render + jinja_template = self.jinja_env.from_string(template.content) + rendered = jinja_template.render(**variables) + return rendered + except Exception as e: + raise ValueError(f"Error rendering template '{prompt_id}': {e}") + + def _validate_input( + self, variables: Dict[str, Any], schema: Dict[str, Any] + ) -> None: + """Basic validation of input variables against schema.""" + for field_name, field_type in schema.items(): + if field_name in variables: + value = variables[field_name] + expected_type = self._get_python_type(field_type) + + if not isinstance(value, expected_type): + raise ValueError( + f"Invalid type for field '{field_name}': " + f"expected {getattr(expected_type, '__name__', str(expected_type))}, got {type(value).__name__}" + ) + + def _get_python_type(self, schema_type: str) -> Union[type, tuple]: + """Convert schema type string to Python type.""" + type_mapping: Dict[str, Union[type, tuple]] = { + "string": str, + "str": str, + "number": (int, float), + "integer": int, + "int": int, + "float": float, + "boolean": bool, + "bool": bool, + "array": list, + "list": list, + "object": dict, + "dict": dict, + } + + return type_mapping.get(schema_type.lower(), str) # type: ignore + + def get_prompt(self, prompt_id: str) -> Optional[PromptTemplate]: + """Get a prompt template by ID.""" + return self.prompts.get(prompt_id) + + def list_prompts(self) -> List[str]: + """Get a list of all available prompt IDs.""" + return list(self.prompts.keys()) + + def get_prompt_metadata(self, prompt_id: str) -> Optional[Dict[str, Any]]: + """Get metadata for a specific prompt.""" + template = self.prompts.get(prompt_id) + return template.metadata if template else None + + def reload_prompts(self) -> None: + """Reload all prompts from the directory.""" + self.prompts.clear() + self._load_prompts() + + def add_prompt( + self, prompt_id: str, content: str, metadata: Optional[Dict[str, Any]] = None + ) -> None: + """Add a prompt template programmatically.""" + template = PromptTemplate( + content=content, metadata=metadata or {}, template_id=prompt_id + ) + self.prompts[prompt_id] = template diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 86e7eb89a2..13a2e554f1 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -18,24 +18,22 @@ else: def safe_divide_seconds( - seconds: float, - denominator: float, - default: Optional[float] = None + seconds: float, denominator: float, default: Optional[float] = None ) -> Optional[float]: """ Safely divide seconds by denominator, handling zero division. - + Args: seconds: Time duration in seconds denominator: The divisor (e.g., number of tokens) default: Value to return if division by zero (defaults to None) - + Returns: The result of the division as a float (seconds per unit), or default if denominator is zero """ if denominator <= 0: return default - + return float(seconds / denominator) @@ -203,3 +201,50 @@ def preserve_upstream_non_openai_attributes( for key, value in original_chunk.model_dump().items(): if key not in expected_keys: setattr(model_response, key, value) + + +def safe_deep_copy(data): + """ + Safe Deep Copy + + The LiteLLM Request has some object that can-not be pickled / deep copied + + Use this function to safely deep copy the LiteLLM Request + """ + import copy + + import litellm + + if litellm.safe_memory_mode is True: + return data + + litellm_parent_otel_span: Optional[Any] = None + # Step 1: Remove the litellm_parent_otel_span + litellm_parent_otel_span = None + if isinstance(data, dict): + # remove litellm_parent_otel_span since this is not picklable + if "metadata" in data and "litellm_parent_otel_span" in data["metadata"]: + litellm_parent_otel_span = data["metadata"].pop("litellm_parent_otel_span") + data["metadata"]["litellm_parent_otel_span"] = "placeholder" + if ( + "litellm_metadata" in data + and "litellm_parent_otel_span" in data["litellm_metadata"] + ): + litellm_parent_otel_span = data["litellm_metadata"].pop( + "litellm_parent_otel_span" + ) + data["litellm_metadata"]["litellm_parent_otel_span"] = "placeholder" + new_data = copy.deepcopy(data) + + # Step 2: re-add the litellm_parent_otel_span after doing a deep copy + if isinstance(data, dict) and litellm_parent_otel_span is not None: + if "metadata" in data and "litellm_parent_otel_span" in data["metadata"]: + data["metadata"]["litellm_parent_otel_span"] = litellm_parent_otel_span + if ( + "litellm_metadata" in data + and "litellm_parent_otel_span" in data["litellm_metadata"] + ): + data["litellm_metadata"][ + "litellm_parent_otel_span" + ] = litellm_parent_otel_span + return new_data diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index 252fb29eb3..9606b47b9b 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -7,6 +7,7 @@ Example: "datadog" -> DataDogLogger "prometheus" -> PrometheusLogger """ + from typing import Union from litellm.integrations.agentops import AgentOps @@ -31,10 +32,12 @@ from litellm.integrations.mlflow import MlflowLogger from litellm.integrations.openmeter import OpenMeterLogger from litellm.integrations.opentelemetry import OpenTelemetry from litellm.integrations.opik.opik import OpikLogger + try: from litellm_enterprise.integrations.prometheus import PrometheusLogger except Exception: PrometheusLogger = None +from litellm.integrations.dotprompt import DotpromptManager from litellm.integrations.s3_v2 import S3Logger from litellm.integrations.sqs import SQSLogger from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( @@ -47,6 +50,7 @@ class CustomLoggerRegistry: """ Registry mapping the callback class string to the class type. """ + CALLBACK_CLASS_STR_TO_CLASS_TYPE = { "lago": LagoLogger, "openmeter": OpenMeterLogger, @@ -80,6 +84,7 @@ class CustomLoggerRegistry: "aws_sqs": SQSLogger, "dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler, "vector_store_pre_call_hook": VectorStorePreCallHook, + "dotprompt": DotpromptManager, } try: @@ -110,14 +115,17 @@ class CustomLoggerRegistry: def get_callback_str_from_class_type(cls, class_type: type) -> Union[str, None]: """ Get the callback string from the class type. - + Args: class_type: The class type to find the string for - + Returns: str: The callback string, or None if not found """ - for callback_str, callback_class in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items(): + for ( + callback_str, + callback_class, + ) in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items(): if callback_class == class_type: return callback_str return None @@ -127,15 +135,18 @@ class CustomLoggerRegistry: """ Get all callback strings that map to the same class type. Some class types (like OpenTelemetry) have multiple string mappings. - + Args: class_type: The class type to find all strings for - + Returns: list: List of callback strings that map to the class type """ callback_strs: list[str] = [] - for callback_str, callback_class in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items(): + for ( + callback_str, + callback_class, + ) in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items(): if callback_class == class_type: callback_strs.append(callback_str) - return callback_strs \ No newline at end of file + return callback_strs diff --git a/litellm/litellm_core_utils/fallback_utils.py b/litellm/litellm_core_utils/fallback_utils.py index d5610d5fdd..a5b0c85c81 100644 --- a/litellm/litellm_core_utils/fallback_utils.py +++ b/litellm/litellm_core_utils/fallback_utils.py @@ -1,9 +1,9 @@ import uuid -from copy import deepcopy from typing import Optional import litellm from litellm._logging import verbose_logger +from litellm.litellm_core_utils.core_helpers import safe_deep_copy from .asyncify import run_async_function @@ -41,7 +41,7 @@ async def async_completion_with_fallbacks(**kwargs): most_recent_exception_str: Optional[str] = None for fallback in fallbacks: try: - completion_kwargs = deepcopy(base_kwargs) + completion_kwargs = safe_deep_copy(base_kwargs) # Handle dictionary fallback configurations if isinstance(fallback, dict): model = fallback.pop("model", original_model) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 12af18804d..029a829f2b 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -120,6 +120,7 @@ from ..integrations.azure_storage.azure_storage import AzureBlobStorageLogger from ..integrations.custom_prompt_management import CustomPromptManagement from ..integrations.datadog.datadog import DataDogLogger from ..integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger +from ..integrations.dotprompt import DotpromptManager from ..integrations.dynamodb import DyanmoDBLogger from ..integrations.galileo import GalileoObserve from ..integrations.gcs_bucket.gcs_bucket import GCSBucketLogger @@ -172,7 +173,6 @@ try: StandardLoggingPayloadSetup as EnterpriseStandardLoggingPayloadSetup, ) - EnterpriseStandardLoggingPayloadSetupVAR: Optional[ Type[EnterpriseStandardLoggingPayloadSetup] ] = EnterpriseStandardLoggingPayloadSetup @@ -599,9 +599,7 @@ class Logging(LiteLLMLoggingBaseClass): custom_logger = ( prompt_management_logger or self.get_custom_logger_for_prompt_management( - model=model, - tools=tools, - non_default_params=non_default_params + model=model, tools=tools, non_default_params=non_default_params ) ) @@ -673,16 +671,16 @@ class Logging(LiteLLMLoggingBaseClass): # Vector Store / Knowledge Base hooks ######################################################### if litellm.vector_store_registry is not None: - - vector_store_custom_logger = _init_custom_logger_compatible_class( - logging_integration="vector_store_pre_call_hook", - internal_usage_cache=None, - llm_router=None, - ) - self.model_call_details["prompt_integration"] = ( - vector_store_custom_logger.__class__.__name__ - ) - return vector_store_custom_logger + + vector_store_custom_logger = _init_custom_logger_compatible_class( + logging_integration="vector_store_pre_call_hook", + internal_usage_cache=None, + llm_router=None, + ) + self.model_call_details["prompt_integration"] = ( + vector_store_custom_logger.__class__.__name__ + ) + return vector_store_custom_logger return None @@ -1315,9 +1313,9 @@ class Logging(LiteLLMLoggingBaseClass): if ( EnterpriseCallbackControls is not None and EnterpriseCallbackControls.is_callback_disabled_dynamically( - callback=callback, + callback=callback, litellm_params=litellm_params, - standard_callback_dynamic_params = self.standard_callback_dynamic_params + standard_callback_dynamic_params=self.standard_callback_dynamic_params, ) ): verbose_logger.debug( @@ -2266,7 +2264,7 @@ class Logging(LiteLLMLoggingBaseClass): start_time=start_time, end_time=end_time, ) - + if isinstance(callback, CustomLogger): # custom logger class model_call_details: Dict = self.model_call_details ################################## @@ -2276,10 +2274,7 @@ class Logging(LiteLLMLoggingBaseClass): ) ################################## if self.stream is True: - if ( - "async_complete_streaming_response" - in model_call_details - ): + if "async_complete_streaming_response" in model_call_details: await callback.async_log_success_event( kwargs=model_call_details, response_obj=model_call_details[ @@ -3217,11 +3212,10 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 _literalai_logger = LiteralAILogger() _in_memory_loggers.append(_literalai_logger) return _literalai_logger # type: ignore - elif logging_integration == "prometheus": - if PrometheusLogger is not None: - for callback in _in_memory_loggers: - if isinstance(callback, PrometheusLogger): - return callback # type: ignore + elif logging_integration == "prometheus" and PrometheusLogger is not None: + for callback in _in_memory_loggers: + if isinstance(callback, PrometheusLogger): + return callback # type: ignore _prometheus_logger = PrometheusLogger() _in_memory_loggers.append(_prometheus_logger) @@ -3493,7 +3487,7 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( VectorStorePreCallHook, ) - + for callback in _in_memory_loggers: if isinstance(callback, VectorStorePreCallHook): return callback @@ -3536,6 +3530,15 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 humanloop_logger = HumanloopLogger() _in_memory_loggers.append(humanloop_logger) return humanloop_logger # type: ignore + elif logging_integration == "dotprompt": + for callback in _in_memory_loggers: + if isinstance(callback, DotpromptManager): + return callback + + dotprompt_logger = DotpromptManager() + _in_memory_loggers.append(dotprompt_logger) + return dotprompt_logger # type: ignore + return None except Exception as e: verbose_logger.exception( f"[Non-Blocking Error] Error initializing custom logger: {e}" @@ -3582,11 +3585,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915 for callback in _in_memory_loggers: if isinstance(callback, LiteralAILogger): return callback - elif logging_integration == "prometheus": - if PrometheusLogger is not None: - for callback in _in_memory_loggers: - if isinstance(callback, PrometheusLogger): - return callback + elif logging_integration == "prometheus" and PrometheusLogger is not None: + for callback in _in_memory_loggers: + if isinstance(callback, PrometheusLogger): + return callback elif logging_integration == "datadog": for callback in _in_memory_loggers: if isinstance(callback, DataDogLogger): @@ -3686,7 +3688,7 @@ def get_custom_logger_compatible_class( # noqa: PLR0915 from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( VectorStorePreCallHook, ) - + for callback in _in_memory_loggers: if isinstance(callback, VectorStorePreCallHook): return callback diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 258601ff5a..827d28598e 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -822,3 +822,41 @@ def set_last_user_message( messages.reverse() messages.append({"role": "user", "content": content}) return messages + + +def convert_prefix_message_to_non_prefix_messages( + messages: List[AllMessageValues], +) -> List[AllMessageValues]: + """ + For models that don't support {prefix: true} in messages, we need to convert the prefix message to a non-prefix message. + + Use prompt: + + {"role": "assistant", "content": "value", "prefix": true} -> [ + { + "role": "system", + "content": "You are a helpful assistant. You are given a message and you need to respond to it. You are also given a generated content. You need to respond to the message in continuation of the generated content. Do not repeat the same content. Your response should be in continuation of this text: ", + }, + { + "role": "assistant", + "content": message["content"], + }, + ] + + do this in place + """ + new_messages: List[AllMessageValues] = [] + for message in messages: + if message.get("prefix"): + new_messages.append( + { + "role": "system", + "content": "You are a helpful assistant. You are given a message and you need to respond to it. You are also given a generated content. You need to respond to the message in continuation of the generated content. Do not repeat the same content. Your response should be in continuation of this text: ", + } + ) + new_messages.append( + {**{k: v for k, v in message.items() if k != "prefix"}} # type: ignore + ) + else: + new_messages.append(message) + return new_messages diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 91a5b317fd..b4ace1545d 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1121,13 +1121,14 @@ def convert_to_gemini_tool_call_result( } """ content_str: str = "" - if isinstance(message["content"], str): - content_str = message["content"] - elif isinstance(message["content"], List): - content_list = message["content"] - for content in content_list: - if content["type"] == "text": - content_str += content["text"] + if "content" in message: + if isinstance(message["content"], str): + content_str = message["content"] + elif isinstance(message["content"], List): + content_list = message["content"] + for content in content_list: + if content["type"] == "text": + content_str += content["text"] name: Optional[str] = message.get("name", "") # type: ignore # Recover name from last message with tool calls diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 17f5a8d1de..2e9e6770a1 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -940,8 +940,8 @@ class CustomStreamWrapper: and not self.sent_last_thinking_block and model_response.choices[0].delta.content ): - model_response.choices[0].delta.content = ( - "" + (model_response.choices[0].delta.content or "") + model_response.choices[0].delta.content = "" + ( + model_response.choices[0].delta.content or "" ) self.sent_last_thinking_block = True @@ -1841,13 +1841,25 @@ class CustomStreamWrapper: self.logging_obj.async_failure_handler(e, traceback_exception) # type: ignore ) ## Map to OpenAI Exception - raise exception_type( - model=self.model, - custom_llm_provider=self.custom_llm_provider, - original_exception=e, - completion_kwargs={}, - extra_kwargs={}, - ) + try: + exception_type( + model=self.model, + custom_llm_provider=self.custom_llm_provider, + original_exception=e, + completion_kwargs={}, + extra_kwargs={}, + ) + except Exception as e: + from litellm.exceptions import MidStreamFallbackError + + raise MidStreamFallbackError( + message=str(e), + model=self.model, + llm_provider=self.custom_llm_provider or "anthropic", + original_exception=e, + generated_content=self.response_uptil_now, + is_pre_first_chunk=not self.sent_first_chunk, + ) @staticmethod def _strip_sse_data_from_chunk(chunk: Optional[str]) -> Optional[str]: diff --git a/litellm/main.py b/litellm/main.py index 253e382a36..b3850e0438 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -147,8 +147,8 @@ from .llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from .llms.custom_llm import CustomLLM, custom_chat_llm_router from .llms.databricks.embed.handler import DatabricksEmbeddingHandler from .llms.deprecated_providers import aleph_alpha, palm -from .llms.groq.chat.handler import GroqChatCompletion from .llms.gemini.common_utils import get_api_key_from_env +from .llms.groq.chat.handler import GroqChatCompletion from .llms.huggingface.embedding.handler import HuggingFaceEmbedding from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion from .llms.ollama.completion import handler as ollama @@ -1049,11 +1049,13 @@ def completion( # type: ignore # noqa: PLR0915 non_default_params = get_non_default_completion_params(kwargs=kwargs) litellm_params = {} # used to prevent unbound var errors ## PROMPT MANAGEMENT HOOKS ## + if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and ( litellm_logging_obj.should_run_prompt_management_hooks( prompt_id=prompt_id, non_default_params=non_default_params ) ): + ( model, messages, @@ -4063,9 +4065,7 @@ def embedding( # noqa: PLR0915 litellm_params={}, ) elif custom_llm_provider == "gemini": - gemini_api_key = ( - api_key or get_api_key_from_env() or litellm.api_key - ) + gemini_api_key = api_key or get_api_key_from_env() or litellm.api_key api_base = api_base or litellm.api_base or get_secret_str("GEMINI_API_BASE") @@ -5495,6 +5495,7 @@ def speech( # noqa: PLR0915 ##### Health Endpoints ####################### + async def ahealth_check( model_params: dict, mode: Optional[ @@ -5540,7 +5541,11 @@ async def ahealth_check( log_raw_request_response=True, ) model_params["litellm_logging_obj"] = litellm_logging_obj - model_params = HealthCheckHelpers._update_model_params_with_health_check_tracking_information(model_params=model_params) + model_params = ( + HealthCheckHelpers._update_model_params_with_health_check_tracking_information( + model_params=model_params + ) + ) ######################################################### try: model: Optional[str] = model_params.get("model", None) diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub_table.html rename to litellm/proxy/_experimental/out/model_hub_table/index.html diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html deleted file mode 100644 index 8d2b8fa254..0000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 8a63ed7b1f..9f203055ae 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -1,9 +1,19 @@ model_list: - - model_name: genai/test/* +<<<<<<< HEAD + - model_name: "gpt-4o-mini-openai" litellm_params: - model: openai/* - api_base: https://api.openai.com + model: gpt-4o-mini + api_key: os.environ/OPENAI_API_KEY + +router_settings: + model_group_alias: {"gpt-4o": "gpt-4o-mini-openai"} +======= + - model_name: openai-test + litellm_params: + model: dotprompt/gpt-3.5-turbo + prompt_id: test_hello_world_prompt api_key: os.environ/OPENAI_API_KEY litellm_settings: - check_provider_endpoint: true \ No newline at end of file + global_prompt_directory: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts +>>>>>>> litellm_dev_07_31_2025_p1 diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index fb70426bd5..31395fe7da 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -579,7 +579,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 elif api_key == "": # missing 'Bearer ' prefix raise Exception( - f"Malformed API Key passed in. Ensure Key has `Bearer ` prefix. Passed in: {passed_in_key}" + "Malformed API Key passed in. Ensure Key has `Bearer ` prefix." ) if route == "/user/auth": @@ -1237,7 +1237,7 @@ def get_api_key_from_custom_header( api_key = _get_bearer_token(api_key=custom_api_key) verbose_proxy_logger.debug( "Found custom API key using header: {}, setting api_key={}".format( - custom_litellm_key_header_name, api_key + custom_litellm_key_header_name, abbreviate_api_key(api_key) ) ) else: diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/__init__.py index 747ef3cc91..f7c05fb8c4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/__init__.py @@ -24,6 +24,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" **{ **litellm_params.model_dump(), "guardrail_name": guardrail_name, + "event_hook": litellm_params.mode, "default_on": litellm_params.default_on or False, } ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index f75f6db82a..41b3c1368a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -293,7 +293,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): return None @log_guardrail_information - async def async_post_call_hook( + async def async_post_call_success_hook( self, data: Dict[str, Any], user_api_key_dict: UserAPIKeyAuth, diff --git a/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py b/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py index f24997d96e..994247ea94 100644 --- a/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py +++ b/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py @@ -169,23 +169,6 @@ async def _calculate_dau_wau_mau( return result -def _extract_user_agent_from_tag(tag: str) -> Optional[str]: - """ - Extract user agent name from tag. - Tags are in format "User-Agent: " or "User-Agent: /" - """ - if not tag.startswith("User-Agent: "): - return None - - user_agent = tag[12:] # Remove "User-Agent: " prefix - - # If it contains a version, extract just the name part - if "/" in user_agent: - return user_agent.split("/")[0] - - return user_agent - - @router.get( "/tag/user-agent/analytics", response_model=UserAgentAnalyticsResponse, @@ -203,7 +186,7 @@ async def get_user_agent_analytics( ), user_agent_filter: Optional[str] = Query( default=None, - description="Filter by specific user agent (e.g., 'curl', 'litellm')", + description="Filter by specific user agent tag", ), page: int = Query(default=1, description="Page number for pagination", ge=1), page_size: int = Query( @@ -212,20 +195,20 @@ async def get_user_agent_analytics( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ - Get user agent analytics including DAU, WAU, MAU, successful requests, and completed tokens by user agent tags. + Get user agent analytics including DAU, WAU, MAU, successful requests, and completed tokens by tags. - This endpoint analyzes the user-agent tags that are automatically tracked by the system - and provides analytics broken down by user agent. + This endpoint analyzes all tags that are tracked by the system and provides analytics + broken down by individual tags. Args: start_date: Start date for the analytics period (YYYY-MM-DD) end_date: End date for the analytics period (YYYY-MM-DD) - user_agent_filter: Filter results to specific user agent name + user_agent_filter: Filter results to specific tag page: Page number for pagination page_size: Number of items per page Returns: - UserAgentAnalyticsResponse: Analytics data broken down by user agent and date + UserAgentAnalyticsResponse: Analytics data broken down by tag and date """ from litellm.proxy.proxy_server import prisma_client @@ -242,25 +225,19 @@ async def get_user_agent_analytics( ) try: - # Get all user-agent tags from the database - user_agent_tags_records = await prisma_client.db.litellm_dailytagspend.find_many( - where={ - "tag": {"startswith": "User-Agent: "}, - "date": {"gte": start_date, "lte": end_date}, - }, + # Get all tags from the database + where_clause = {"date": {"gte": start_date, "lte": end_date}} + if user_agent_filter: + where_clause["tag"] = {"contains": user_agent_filter} + + tag_records = await prisma_client.db.litellm_dailytagspend.find_many( + where=where_clause, distinct=["tag"], ) - user_agent_tags = [record.tag for record in user_agent_tags_records] + tags = [record.tag for record in tag_records] - # Filter by user agent if specified - if user_agent_filter: - user_agent_tags = [ - tag for tag in user_agent_tags - if user_agent_filter.lower() in tag.lower() - ] - - if not user_agent_tags: + if not tags: return UserAgentAnalyticsResponse( results=[], total_count=0, @@ -269,12 +246,12 @@ async def get_user_agent_analytics( total_pages=0, ) - # Get daily activity data for user-agent tags + # Get daily activity data for tags daily_activity_response = await get_daily_activity( prisma_client=prisma_client, table_name="litellm_dailytagspend", entity_id_field="tag", - entity_id=user_agent_tags, + entity_id=tags, entity_metadata_field=None, start_date=start_date, end_date=end_date, @@ -284,7 +261,7 @@ async def get_user_agent_analytics( page_size=10000, # Large page size to get all data ) - # Process the results to calculate DAU/WAU/MAU and organize by user agent + # Process the results to calculate DAU/WAU/MAU and organize by tag results = [] daily_data_by_tag_and_date: Dict[str, Dict[str, DailySpendData]] = {} @@ -294,19 +271,16 @@ async def get_user_agent_analytics( # Get tag from breakdown data for tag, tag_metrics in daily_data.breakdown.entities.items(): - if tag.startswith("User-Agent: "): - if tag not in daily_data_by_tag_and_date: - daily_data_by_tag_and_date[tag] = {} - daily_data_by_tag_and_date[tag][date_str] = daily_data + if tag not in daily_data_by_tag_and_date: + daily_data_by_tag_and_date[tag] = {} + daily_data_by_tag_and_date[tag][date_str] = daily_data # Calculate DAU/WAU/MAU for each date and tag combination unique_dates: set[str] = set() for tag_data in daily_data_by_tag_and_date.values(): unique_dates.update(tag_data.keys()) - for tag in user_agent_tags: - user_agent = _extract_user_agent_from_tag(tag) - + for tag in tags: for date_str in sorted(unique_dates): if tag in daily_data_by_tag_and_date and date_str in daily_data_by_tag_and_date[tag]: daily_data = daily_data_by_tag_and_date[tag][date_str] @@ -334,13 +308,13 @@ async def get_user_agent_analytics( UserAgentActivityData( date=date_str, tag=tag, - user_agent=user_agent, + user_agent=tag, # Use the full tag as user_agent metrics=metrics, ) ) - # Sort results by date (most recent first) and then by user agent - results.sort(key=lambda x: (x.date, x.user_agent or ""), reverse=True) + # Sort results by date (most recent first) and then by tag + results.sort(key=lambda x: (x.date, x.tag), reverse=True) # Apply pagination total_count = len(results) @@ -381,9 +355,9 @@ async def get_user_agent_summary( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ - Get summary statistics for user agent activity. + Get summary statistics for tag activity. - Returns aggregated metrics across all user agents for the specified time period. + Returns aggregated metrics across all tags for the specified time period. """ from litellm.proxy.proxy_server import prisma_client @@ -400,26 +374,25 @@ async def get_user_agent_summary( ) try: - # Get all user-agent tags - user_agent_tags_records = await prisma_client.db.litellm_dailytagspend.find_many( + # Get all tags + tag_records = await prisma_client.db.litellm_dailytagspend.find_many( where={ - "tag": {"startswith": "User-Agent: "}, "date": {"gte": start_date, "lte": end_date}, }, distinct=["tag"], ) - user_agent_tags = [record.tag for record in user_agent_tags_records] + tags = [record.tag for record in tag_records] - if not user_agent_tags: + if not tags: return { - "total_user_agents": 0, + "total_tags": 0, "total_requests": 0, "total_successful_requests": 0, "total_failed_requests": 0, "total_tokens": 0, "total_spend": 0.0, - "top_user_agents": [], + "top_tags": [], } # Get aggregated data @@ -427,7 +400,7 @@ async def get_user_agent_summary( prisma_client=prisma_client, table_name="litellm_dailytagspend", entity_id_field="tag", - entity_id=user_agent_tags, + entity_id=tags, entity_metadata_field=None, start_date=start_date, end_date=end_date, @@ -437,57 +410,54 @@ async def get_user_agent_summary( page_size=10000, ) - # Aggregate metrics by user agent - user_agent_totals: Dict[str, UserAgentMetrics] = {} + # Aggregate metrics by tag + tag_totals: Dict[str, UserAgentMetrics] = {} for daily_data in daily_activity_response.results: for tag, tag_metrics in daily_data.breakdown.entities.items(): - if tag.startswith("User-Agent: "): - user_agent = _extract_user_agent_from_tag(tag) - if user_agent is not None and user_agent not in user_agent_totals: - user_agent_totals[user_agent] = UserAgentMetrics() - - if user_agent is not None: - totals = user_agent_totals[user_agent] - totals.successful_requests += tag_metrics.metrics.successful_requests - totals.failed_requests += tag_metrics.metrics.failed_requests - totals.total_requests += tag_metrics.metrics.api_requests - totals.completed_tokens += tag_metrics.metrics.completion_tokens - totals.total_tokens += tag_metrics.metrics.total_tokens - totals.spend += tag_metrics.metrics.spend + if tag not in tag_totals: + tag_totals[tag] = UserAgentMetrics() + + totals = tag_totals[tag] + totals.successful_requests += tag_metrics.metrics.successful_requests + totals.failed_requests += tag_metrics.metrics.failed_requests + totals.total_requests += tag_metrics.metrics.api_requests + totals.completed_tokens += tag_metrics.metrics.completion_tokens + totals.total_tokens += tag_metrics.metrics.total_tokens + totals.spend += tag_metrics.metrics.spend # Calculate summary statistics - total_requests = sum(ua.total_requests for ua in user_agent_totals.values()) - total_successful_requests = sum(ua.successful_requests for ua in user_agent_totals.values()) - total_failed_requests = sum(ua.failed_requests for ua in user_agent_totals.values()) - total_tokens = sum(ua.total_tokens for ua in user_agent_totals.values()) - total_spend = sum(ua.spend for ua in user_agent_totals.values()) + total_requests = sum(tag.total_requests for tag in tag_totals.values()) + total_successful_requests = sum(tag.successful_requests for tag in tag_totals.values()) + total_failed_requests = sum(tag.failed_requests for tag in tag_totals.values()) + total_tokens = sum(tag.total_tokens for tag in tag_totals.values()) + total_spend = sum(tag.spend for tag in tag_totals.values()) - # Get top user agents by request count - top_user_agents = sorted( + # Get top tags by request count + top_tags = sorted( [ { - "user_agent": ua, + "tag": tag, "requests": metrics.total_requests, "successful_requests": metrics.successful_requests, "failed_requests": metrics.failed_requests, "tokens": metrics.total_tokens, "spend": metrics.spend, } - for ua, metrics in user_agent_totals.items() + for tag, metrics in tag_totals.items() ], key=lambda x: cast(int, x["requests"]), reverse=True, )[:10] # Top 10 return { - "total_user_agents": len(user_agent_totals), + "total_tags": len(tag_totals), "total_requests": total_requests, "total_successful_requests": total_successful_requests, "total_failed_requests": total_failed_requests, "total_tokens": total_tokens, "total_spend": total_spend, - "top_user_agents": top_user_agents, + "top_tags": top_tags, } except Exception as e: @@ -522,7 +492,7 @@ async def get_per_user_analytics( Get per-user analytics including successful requests, tokens, and spend by individual users. This endpoint provides usage metrics broken down by individual users based on their - user-agent activity during the specified time period. + tag activity during the specified time period. Args: start_date: Start date for the analytics period (YYYY-MM-DD) @@ -548,30 +518,9 @@ async def get_per_user_analytics( ) try: - # Get all user-agent tags from the database - user_agent_tags_records = await prisma_client.db.litellm_dailytagspend.find_many( - where={ - "tag": {"startswith": "User-Agent: "}, - "date": {"gte": start_date, "lte": end_date}, - }, - distinct=["tag"], - ) - - user_agent_tags = [record.tag for record in user_agent_tags_records] - - if not user_agent_tags: - return PerUserAnalyticsResponse( - results=[], - total_count=0, - page=page, - page_size=page_size, - total_pages=0, - ) - - # Get all records for user-agent tags in the date range + # Get all tag records in the date range tag_records = await prisma_client.db.litellm_dailytagspend.find_many( where={ - "tag": {"in": user_agent_tags}, "date": {"gte": start_date, "lte": end_date} } ) @@ -618,18 +567,18 @@ async def get_per_user_analytics( for record in tag_records: if record.api_key in api_key_to_user_id: user_id = api_key_to_user_id[record.api_key] - user_agent = _extract_user_agent_from_tag(record.tag) + tag = record.tag # Use the full tag as user_agent if user_id not in user_metrics: user_metrics[user_id] = PerUserMetrics( user_id=user_id, user_email=user_id_to_email.get(user_id), - user_agent=user_agent + user_agent=tag ) else: - # If user agent is different, keep the first one or prioritize certain ones - if user_agent and not user_metrics[user_id].user_agent: - user_metrics[user_id].user_agent = user_agent + # If tag is different, keep the first one or prioritize certain ones + if tag and not user_metrics[user_id].user_agent: + user_metrics[user_id].user_agent = tag # Aggregate metrics user_metrics[user_id].successful_requests += record.successful_requests or 0 diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index eb06d7c8f4..d743d41ac6 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -172,7 +172,7 @@ async def gemini_proxy_route( request=request, api_key=f"Bearer {google_ai_studio_api_key}" ) - base_target_url = "https://generativelanguage.googleapis.com" + base_target_url = os.getenv("GEMINI_API_BASE") or "https://generativelanguage.googleapis.com" encoded_endpoint = httpx.URL(endpoint).path # Ensure endpoint starts with '/' for proper URL construction @@ -231,7 +231,7 @@ async def cohere_proxy_route( """ [Docs](https://docs.litellm.ai/docs/pass_through/cohere) """ - base_target_url = "https://api.cohere.com" + base_target_url = os.getenv("COHERE_API_BASE") or "https://api.cohere.com" encoded_endpoint = httpx.URL(endpoint).path # Ensure endpoint starts with '/' for proper URL construction @@ -427,7 +427,7 @@ async def anthropic_proxy_route( """ [Docs](https://docs.litellm.ai/docs/anthropic_completion) """ - base_target_url = "https://api.anthropic.com" + base_target_url = os.getenv("ANTHROPIC_API_BASE") or "https://api.anthropic.com" encoded_endpoint = httpx.URL(endpoint).path # Ensure endpoint starts with '/' for proper URL construction @@ -1017,7 +1017,7 @@ async def openai_proxy_route( """ - base_target_url = "https://api.openai.com/" + base_target_url = os.getenv("OPENAI_API_BASE") or "https://api.openai.com/" # Add or update query parameters openai_api_key = passthrough_endpoint_router.get_credentials( custom_llm_provider=litellm.LlmProviders.OPENAI.value, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6b4b4c25dd..98355afdfd 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -272,9 +272,6 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import scim_router from litellm.proxy.management_endpoints.tag_management_endpoints import ( router as tag_management_router, ) -from litellm.proxy.management_endpoints.user_agent_analytics_endpoints import ( - router as user_agent_analytics_router, -) from litellm.proxy.management_endpoints.team_callback_endpoints import ( router as team_callback_router, ) @@ -287,6 +284,9 @@ from litellm.proxy.management_endpoints.ui_sso import ( get_disabled_non_admin_personal_key_creation, ) from litellm.proxy.management_endpoints.ui_sso import router as ui_sso_router +from litellm.proxy.management_endpoints.user_agent_analytics_endpoints import ( + router as user_agent_analytics_router, +) from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMiddleware from litellm.proxy.openai_files_endpoints.files_endpoints import ( @@ -1822,6 +1822,15 @@ class ProxyConfig: ) litellm.guardrail_name_config_map = guardrail_name_config_map + elif key == "global_prompt_directory": + from litellm.integrations.dotprompt import ( + set_global_prompt_directory, + ) + + set_global_prompt_directory(value) + verbose_proxy_logger.info( + f"{blue_color_code}Set Global Prompt Directory on LiteLLM Proxy{reset_color_code}" + ) elif key == "callbacks": initialize_callbacks_on_proxy( value=value, @@ -2220,7 +2229,9 @@ class ProxyConfig: litellm_settings = config.get("litellm_settings", {}) mcp_aliases = litellm_settings.get("mcp_aliases", None) - global_mcp_server_manager.load_servers_from_config(mcp_servers_config, mcp_aliases) + global_mcp_server_manager.load_servers_from_config( + mcp_servers_config, mcp_aliases + ) ## VECTOR STORES vector_store_registry_config = config.get("vector_store_registry", None) @@ -3253,7 +3264,6 @@ async def async_data_generator( "async_data_generator: received streaming chunk - {}".format(chunk) ) - ### CALL HOOKS ### - modify outgoing data chunk = await proxy_logging_obj.async_post_call_streaming_hook( user_api_key_dict=user_api_key_dict, @@ -3262,7 +3272,6 @@ async def async_data_generator( str_so_far=str_so_far, ) - if isinstance(chunk, (ModelResponse, ModelResponseStream)): response_str = litellm.get_response_string(response_obj=chunk) str_so_far += response_str diff --git a/litellm/proxy/test_prompts/test_hello_world_prompt.prompt b/litellm/proxy/test_prompts/test_hello_world_prompt.prompt new file mode 100644 index 0000000000..b8fbc6e3a0 --- /dev/null +++ b/litellm/proxy/test_prompts/test_hello_world_prompt.prompt @@ -0,0 +1,10 @@ +--- +model: gpt-3.5-turbo +input: + schema: + text: string +--- + +Extract the requested information from the given text. If a piece of information is not present, omit that field from the output. + +Text: {{text}} diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 3124da6b4b..4486b99d2e 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -52,11 +52,6 @@ from litellm import ( ModelResponseStream, Router, ) -from litellm.types.mcp import ( - MCPPreCallRequestObject, - MCPPreCallResponseObject, - MCPDuringCallResponseObject, -) from litellm._logging import verbose_proxy_logger from litellm._service_logger import ServiceLogging, ServiceTypes from litellm.caching.caching import DualCache, RedisCache @@ -93,6 +88,11 @@ from litellm.proxy.hooks.parallel_request_limiter import ( from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.secret_managers.main import str_to_bool from litellm.types.integrations.slack_alerting import DEFAULT_ALERT_TYPES +from litellm.types.mcp import ( + MCPDuringCallResponseObject, + MCPPreCallRequestObject, + MCPPreCallResponseObject, +) from litellm.types.utils import CallTypes, LLMResponseTypes, LoggedLiteLLMParams if TYPE_CHECKING: @@ -118,33 +118,6 @@ def print_verbose(print_statement): print(f"LiteLLM Proxy: {print_statement}") # noqa -def safe_deep_copy(data): - """ - Safe Deep Copy - - The LiteLLM Request has some object that can-not be pickled / deep copied - - Use this function to safely deep copy the LiteLLM Request - """ - if litellm.safe_memory_mode is True: - return data - - litellm_parent_otel_span: Optional[Any] = None - # Step 1: Remove the litellm_parent_otel_span - litellm_parent_otel_span = None - if isinstance(data, dict): - # remove litellm_parent_otel_span since this is not picklable - if "metadata" in data and "litellm_parent_otel_span" in data["metadata"]: - litellm_parent_otel_span = data["metadata"].pop("litellm_parent_otel_span") - new_data = copy.deepcopy(data) - - # Step 2: re-add the litellm_parent_otel_span after doing a deep copy - if isinstance(data, dict) and litellm_parent_otel_span is not None: - if "metadata" in data: - data["metadata"]["litellm_parent_otel_span"] = litellm_parent_otel_span - return new_data - - class InternalUsageCache: def __init__(self, dual_cache: DualCache): self.dual_cache: DualCache = dual_cache @@ -474,11 +447,11 @@ class ProxyLogging: ) async def async_pre_mcp_tool_call_hook( - self, - kwargs: dict, - request_obj: Any, - start_time: datetime, - end_time: datetime, + self, + kwargs: dict, + request_obj: Any, + start_time: datetime, + end_time: datetime, ) -> Optional[Any]: """ Pre MCP Tool Call Hook @@ -489,7 +462,7 @@ class ProxyLogging: from litellm.types.mcp import MCPPreCallRequestObject, MCPPreCallResponseObject callbacks = self.get_combined_callback_list( - dynamic_success_callbacks=getattr(self, 'dynamic_success_callbacks', None), + dynamic_success_callbacks=getattr(self, "dynamic_success_callbacks", None), global_callbacks=litellm.success_callback, ) @@ -500,7 +473,7 @@ class ProxyLogging: arguments=kwargs.get("arguments", {}), server_name=kwargs.get("server_name"), user_api_key_auth=kwargs.get("user_api_key_auth"), - hidden_params=HiddenParams() + hidden_params=HiddenParams(), ) for callback in callbacks: @@ -537,10 +510,10 @@ class ProxyLogging: return global_callbacks return list(set(dynamic_success_callbacks + global_callbacks)) - - def _parse_pre_mcp_call_hook_response( - self, response: MCPPreCallResponseObject, original_request: MCPPreCallRequestObject + self, + response: MCPPreCallResponseObject, + original_request: MCPPreCallRequestObject, ) -> Dict[str, Any]: """ Parse the response from the pre_mcp_tool_call_hook @@ -551,18 +524,19 @@ class ProxyLogging: """ result = { "should_proceed": response.should_proceed, - "modified_arguments": response.modified_arguments or original_request.arguments, + "modified_arguments": response.modified_arguments + or original_request.arguments, "error_message": response.error_message, "hidden_params": response.hidden_params, } return result async def async_during_mcp_tool_call_hook( - self, - kwargs: dict, - request_obj: Any, - start_time: datetime, - end_time: datetime, + self, + kwargs: dict, + request_obj: Any, + start_time: datetime, + end_time: datetime, ) -> Optional[Any]: """ During MCP Tool Call Hook @@ -570,10 +544,13 @@ class ProxyLogging: Use this for concurrent monitoring and validation during tool execution. """ from litellm.types.llms.base import HiddenParams - from litellm.types.mcp import MCPDuringCallResponseObject, MCPDuringCallRequestObject + from litellm.types.mcp import ( + MCPDuringCallRequestObject, + MCPDuringCallResponseObject, + ) callbacks = self.get_combined_callback_list( - dynamic_success_callbacks=getattr(self, 'dynamic_success_callbacks', None), + dynamic_success_callbacks=getattr(self, "dynamic_success_callbacks", None), global_callbacks=litellm.success_callback, ) @@ -584,7 +561,7 @@ class ProxyLogging: arguments=kwargs.get("arguments", {}), server_name=kwargs.get("server_name"), start_time=start_time.timestamp() if start_time else None, - hidden_params=HiddenParams() + hidden_params=HiddenParams(), ) for callback in callbacks: @@ -603,7 +580,9 @@ class ProxyLogging: # this allows for execution control decisions ###################################################################### if response is not None: - return self._parse_during_mcp_call_hook_response(response=response) + return self._parse_during_mcp_call_hook_response( + response=response + ) except Exception as e: verbose_proxy_logger.exception( "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format( @@ -613,7 +592,7 @@ class ProxyLogging: return None def _parse_during_mcp_call_hook_response( - self, response: MCPDuringCallResponseObject + self, response: MCPDuringCallResponseObject ) -> Dict[str, Any]: """ Parse the response from the during_mcp_tool_call_hook @@ -1382,9 +1361,15 @@ class PrismaClient: from prisma import Prisma # type: ignore except Exception as e: verbose_proxy_logger.error(f"Failed to import Prisma client: {e}") - verbose_proxy_logger.error("This usually means 'prisma generate' hasn't been run yet.") - verbose_proxy_logger.error("Please run 'prisma generate' to generate the Prisma client.") - raise Exception("Unable to find Prisma binaries. Please run 'prisma generate' first.") + verbose_proxy_logger.error( + "This usually means 'prisma generate' hasn't been run yet." + ) + verbose_proxy_logger.error( + "Please run 'prisma generate' to generate the Prisma client." + ) + raise Exception( + "Unable to find Prisma binaries. Please run 'prisma generate' first." + ) if http_client is not None: self.db = PrismaWrapper( original_prisma=Prisma(http=http_client), diff --git a/litellm/router.py b/litellm/router.py index d02392c59f..ca724006fd 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -23,6 +23,7 @@ from functools import lru_cache from typing import ( TYPE_CHECKING, Any, + AsyncGenerator, Callable, Dict, List, @@ -146,7 +147,7 @@ from litellm.types.services import ServiceTypes from litellm.types.utils import GenericBudgetConfigType, LiteLLMBatch from litellm.types.utils import ModelInfo from litellm.types.utils import ModelInfo as ModelMapInfo -from litellm.types.utils import StandardLoggingPayload +from litellm.types.utils import ModelResponseStream, StandardLoggingPayload, Usage from litellm.utils import ( CustomStreamWrapper, EmbeddingResponse, @@ -1078,9 +1079,144 @@ class Router: ) raise e + async def _acompletion_streaming_iterator( + self, + model_response: CustomStreamWrapper, + messages: List[Dict[str, str]], + initial_kwargs: dict, + ) -> CustomStreamWrapper: + """ + Helper to iterate over a streaming response. + + Catches errors for fallbacks using the router's fallback system + """ + from litellm.exceptions import MidStreamFallbackError + + class FallbackStreamWrapper(CustomStreamWrapper): + def __init__(self, async_generator: AsyncGenerator): + # Copy attributes from the original model_response + super().__init__( + completion_stream=async_generator, + model=model_response.model, + custom_llm_provider=model_response.custom_llm_provider, + logging_obj=model_response.logging_obj, + ) + self._async_generator = async_generator + + def __aiter__(self): + return self + + async def __anext__(self): + return await self._async_generator.__anext__() + + async def stream_with_fallbacks(): + try: + async for item in model_response: + yield item + except MidStreamFallbackError as e: + from litellm.main import stream_chunk_builder + + complete_response_object = stream_chunk_builder( + chunks=model_response.chunks + ) + complete_response_object_usage = cast( + Optional[Usage], + getattr(complete_response_object, "usage", None), + ) + try: + # Use the router's fallback system + model_group = cast(str, initial_kwargs.get("model")) + fallbacks: Optional[List] = initial_kwargs.get( + "fallbacks", self.fallbacks + ) + context_window_fallbacks: Optional[List] = initial_kwargs.get( + "context_window_fallbacks", self.context_window_fallbacks + ) + content_policy_fallbacks: Optional[List] = initial_kwargs.get( + "content_policy_fallbacks", self.content_policy_fallbacks + ) + initial_kwargs["original_function"] = self._acompletion + initial_kwargs["messages"] = messages + [ + { + "role": "system", + "content": "You are a helpful assistant. You are given a message and you need to respond to it. You are also given a generated content. You need to respond to the message in continuation of the generated content. Do not repeat the same content. Your response should be in continuation of this text: ", + }, + { + "role": "assistant", + "content": e.generated_content, + "prefix": True, + }, + ] + self._update_kwargs_before_fallbacks( + model=model_group, kwargs=initial_kwargs + ) + fallback_response = ( + await self.async_function_with_fallbacks_common_utils( + e=e, + disable_fallbacks=False, + fallbacks=fallbacks, + context_window_fallbacks=context_window_fallbacks, + content_policy_fallbacks=content_policy_fallbacks, + model_group=model_group, + args=(), + kwargs=initial_kwargs, + ) + ) + + # If fallback returns a streaming response, iterate over it + if hasattr(fallback_response, "__aiter__"): + async for fallback_item in fallback_response: # type: ignore + if ( + fallback_item + and isinstance(fallback_item, ModelResponseStream) + and hasattr(fallback_item, "usage") + ): + from litellm.cost_calculator import ( + BaseTokenUsageProcessor, + ) + + usage = cast( + Optional[Usage], + getattr(fallback_item, "usage", None), + ) + if usage is not None: + usage_objects = [usage] + else: + usage_objects = [] + + if ( + complete_response_object_usage is not None + and hasattr(complete_response_object_usage, "usage") + and complete_response_object_usage.usage is not None # type: ignore + ): + usage_objects.append(complete_response_object_usage) + + combined_usage = ( + BaseTokenUsageProcessor.combine_usage_objects( + usage_objects=usage_objects + ) + ) + setattr(fallback_item, "usage", combined_usage) + yield fallback_item + else: + # If fallback returns a non-streaming response, yield None + yield None + + except Exception as fallback_error: + # If fallback also fails, log and re-raise original error + verbose_router_logger.error( + f"Fallback also failed: {fallback_error}" + ) + raise fallback_error + + return FallbackStreamWrapper(stream_with_fallbacks()) + async def _acompletion( self, model: str, messages: List[Dict[str, str]], **kwargs - ) -> Union[ModelResponse, CustomStreamWrapper]: + ) -> Union[ + ModelResponse, + CustomStreamWrapper, + ]: """ - Get an available deployment - call it with a semaphore over the call @@ -1092,9 +1228,9 @@ class Router: {} ) # this is a temporary dict to debug timeout issues try: - verbose_router_logger.debug( - f"Inside _acompletion()- model: {model}; kwargs: {kwargs}" - ) + input_kwargs_for_streaming_fallback = kwargs.copy() + input_kwargs_for_streaming_fallback["model"] = model + parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) start_time = time.time() deployment = await self.async_get_available_deployment( @@ -1134,15 +1270,15 @@ class Router: ) self.total_calls[model_name] += 1 - _response = litellm.acompletion( - **{ - **data, - "messages": messages, - "caching": self.cache_responses, - "client": model_client, - **kwargs, - } - ) + input_kwargs = { + **data, + "messages": messages, + "caching": self.cache_responses, + "client": model_client, + **kwargs, + } + + _response = litellm.acompletion(**input_kwargs) logging_obj: Optional[LiteLLMLogging] = kwargs.get( "litellm_logging_obj", None @@ -1199,6 +1335,13 @@ class Router: parent_otel_span=parent_otel_span, ) + if isinstance(response, CustomStreamWrapper): + return await self._acompletion_streaming_iterator( + model_response=response, + messages=messages, + initial_kwargs=input_kwargs_for_streaming_fallback, + ) + return response except litellm.Timeout as e: deployment_request_timeout_param = _timeout_debug_deployment_dict.get( @@ -1577,7 +1720,8 @@ class Router: Wrapper around self.acompletion that catches exceptions and returns them as a result """ try: - return await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) # type: ignore + result = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) # type: ignore + return result except asyncio.CancelledError: verbose_router_logger.debug( "Received 'task.cancel'. Cancelling call w/ model={}.".format(model) @@ -1625,6 +1769,7 @@ class Router: ) for completed_task in done: result = await check_response(completed_task) + if result is not None: # Return the first successful result result._hidden_params["fastest_response_batch_completion"] = True @@ -2914,7 +3059,9 @@ class Router: ) async def create_file_for_deployment(deployment: dict) -> OpenAIFileObject: - kwargs_copy = copy.deepcopy(kwargs) + from litellm.litellm_core_utils.core_helpers import safe_deep_copy + + kwargs_copy = safe_deep_copy(kwargs) self._update_kwargs_with_deployment( deployment=deployment, kwargs=kwargs_copy, @@ -3165,6 +3312,8 @@ class Router: async def try_retrieve_batch(model_name: DeploymentTypedDict): try: + from litellm.litellm_core_utils.core_helpers import safe_deep_copy + model = model_name["litellm_params"].get("model") data = model_name["litellm_params"].copy() custom_llm_provider = data.get("custom_llm_provider") @@ -3178,7 +3327,7 @@ class Router: _, custom_llm_provider, _, _ = get_llm_provider( # type: ignore model=model ) - new_kwargs = copy.deepcopy(kwargs) + new_kwargs = safe_deep_copy(kwargs) self._update_kwargs_with_deployment( deployment=cast(dict, model_name), kwargs=new_kwargs, @@ -3513,8 +3662,199 @@ class Router: #### [END] ASSISTANTS API #### + async def async_function_with_fallbacks_common_utils( # noqa: PLR0915 + self, + e: Exception, + disable_fallbacks: Optional[bool], + fallbacks: Optional[List], + context_window_fallbacks: Optional[List], + content_policy_fallbacks: Optional[List], + model_group: Optional[str], + args: tuple, + kwargs: dict, + ): + """ + Common utilities for async_function_with_fallbacks + """ + verbose_router_logger.debug(f"Traceback{traceback.format_exc()}") + original_exception = e + fallback_model_group = None + original_model_group: Optional[str] = kwargs.get("model") # type: ignore + fallback_failure_exception_str = "" + + if disable_fallbacks is True or original_model_group is None: + raise e + + input_kwargs = { + "litellm_router": self, + "original_exception": original_exception, + **kwargs, + } + + if "max_fallbacks" not in input_kwargs: + input_kwargs["max_fallbacks"] = self.max_fallbacks + if "fallback_depth" not in input_kwargs: + input_kwargs["fallback_depth"] = 0 + + try: + verbose_router_logger.info("Trying to fallback b/w models") + + # check if client-side fallbacks are used (e.g. fallbacks = ["gpt-3.5-turbo", "claude-3-haiku"] or fallbacks=[{"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hey, how's it going?"}]}] + is_non_standard_fallback_format = _check_non_standard_fallback_format( + fallbacks=fallbacks + ) + + if is_non_standard_fallback_format: + input_kwargs.update( + { + "fallback_model_group": fallbacks, + "original_model_group": original_model_group, + } + ) + + response = await run_async_fallback( + *args, + **input_kwargs, + ) + + return response + + if isinstance(e, litellm.ContextWindowExceededError): + if context_window_fallbacks is not None: + context_window_fallback_model_group: Optional[List[str]] = ( + self._get_fallback_model_group_from_fallbacks( + fallbacks=context_window_fallbacks, + model_group=model_group, + ) + ) + if context_window_fallback_model_group is None: + raise original_exception + + input_kwargs.update( + { + "fallback_model_group": context_window_fallback_model_group, + "original_model_group": original_model_group, + } + ) + + response = await run_async_fallback( + *args, + **input_kwargs, + ) + return response + + else: + error_message = "model={}. context_window_fallbacks={}. fallbacks={}.\n\nSet 'context_window_fallback' - https://docs.litellm.ai/docs/routing#fallbacks".format( + model_group, context_window_fallbacks, fallbacks + ) + verbose_router_logger.info( + msg="Got 'ContextWindowExceededError'. No context_window_fallback set. Defaulting \ + to fallbacks, if available.{}".format( + error_message + ) + ) + + e.message += "\n{}".format(error_message) + elif isinstance(e, litellm.ContentPolicyViolationError): + if content_policy_fallbacks is not None: + content_policy_fallback_model_group: Optional[List[str]] = ( + self._get_fallback_model_group_from_fallbacks( + fallbacks=content_policy_fallbacks, + model_group=model_group, + ) + ) + if content_policy_fallback_model_group is None: + raise original_exception + + input_kwargs.update( + { + "fallback_model_group": content_policy_fallback_model_group, + "original_model_group": original_model_group, + } + ) + + response = await run_async_fallback( + *args, + **input_kwargs, + ) + return response + else: + error_message = "model={}. content_policy_fallback={}. fallbacks={}.\n\nSet 'content_policy_fallback' - https://docs.litellm.ai/docs/routing#fallbacks".format( + model_group, content_policy_fallbacks, fallbacks + ) + verbose_router_logger.info( + msg="Got 'ContentPolicyViolationError'. No content_policy_fallback set. Defaulting \ + to fallbacks, if available.{}".format( + error_message + ) + ) + + e.message += "\n{}".format(error_message) + if fallbacks is not None and model_group is not None: + verbose_router_logger.debug(f"inside model fallbacks: {fallbacks}") + ( + fallback_model_group, + generic_fallback_idx, + ) = get_fallback_model_group( + fallbacks=fallbacks, # if fallbacks = [{"gpt-3.5-turbo": ["claude-3-haiku"]}] + model_group=cast(str, model_group), + ) + ## if none, check for generic fallback + if fallback_model_group is None and generic_fallback_idx is not None: + fallback_model_group = fallbacks[generic_fallback_idx]["*"] + + if fallback_model_group is None: + verbose_router_logger.info( + f"No fallback model group found for original model_group={model_group}. Fallbacks={fallbacks}" + ) + if hasattr(original_exception, "message"): + original_exception.message += f"No fallback model group found for original model_group={model_group}. Fallbacks={fallbacks}" # type: ignore + raise original_exception + + input_kwargs.update( + { + "fallback_model_group": fallback_model_group, + "original_model_group": original_model_group, + } + ) + + response = await run_async_fallback( + *args, + **input_kwargs, + ) + + return response + except Exception as new_exception: + parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) + verbose_router_logger.error( + "litellm.router.py::async_function_with_fallbacks() - Error occurred while trying to do fallbacks - {}\n{}\n\nDebug Information:\nCooldown Deployments={}".format( + str(new_exception), + traceback.format_exc(), + await _async_get_cooldown_deployments_with_debug_info( + litellm_router_instance=self, + parent_otel_span=parent_otel_span, + ), + ) + ) + fallback_failure_exception_str = str(new_exception) + + if hasattr(original_exception, "message"): + # add the available fallbacks to the exception + original_exception.message += ". Received Model Group={}\nAvailable Model Group Fallbacks={}".format( # type: ignore + model_group, + fallback_model_group, + ) + if len(fallback_failure_exception_str) > 0: + original_exception.message += ( # type: ignore + "\nError doing the fallback: {}".format( + fallback_failure_exception_str + ) + ) + + raise original_exception + @tracer.wrap() - async def async_function_with_fallbacks(self, *args, **kwargs): # noqa: PLR0915 + async def async_function_with_fallbacks(self, *args, **kwargs): """ Try calling the function_with_retries If it fails after num_retries, fall back to another model group @@ -3553,185 +3893,16 @@ class Router: ) return response except Exception as e: - verbose_router_logger.debug(f"Traceback{traceback.format_exc()}") - original_exception = e - fallback_model_group = None - original_model_group: Optional[str] = kwargs.get("model") # type: ignore - fallback_failure_exception_str = "" - - if disable_fallbacks is True or original_model_group is None: - raise e - - input_kwargs = { - "litellm_router": self, - "original_exception": original_exception, - **kwargs, - } - - if "max_fallbacks" not in input_kwargs: - input_kwargs["max_fallbacks"] = self.max_fallbacks - if "fallback_depth" not in input_kwargs: - input_kwargs["fallback_depth"] = 0 - - try: - verbose_router_logger.info("Trying to fallback b/w models") - - # check if client-side fallbacks are used (e.g. fallbacks = ["gpt-3.5-turbo", "claude-3-haiku"] or fallbacks=[{"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hey, how's it going?"}]}] - is_non_standard_fallback_format = _check_non_standard_fallback_format( - fallbacks=fallbacks - ) - - if is_non_standard_fallback_format: - input_kwargs.update( - { - "fallback_model_group": fallbacks, - "original_model_group": original_model_group, - } - ) - - response = await run_async_fallback( - *args, - **input_kwargs, - ) - - return response - - if isinstance(e, litellm.ContextWindowExceededError): - if context_window_fallbacks is not None: - fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=context_window_fallbacks, - model_group=model_group, - ) - ) - if fallback_model_group is None: - raise original_exception - - input_kwargs.update( - { - "fallback_model_group": fallback_model_group, - "original_model_group": original_model_group, - } - ) - - response = await run_async_fallback( - *args, - **input_kwargs, - ) - return response - - else: - error_message = "model={}. context_window_fallbacks={}. fallbacks={}.\n\nSet 'context_window_fallback' - https://docs.litellm.ai/docs/routing#fallbacks".format( - model_group, context_window_fallbacks, fallbacks - ) - verbose_router_logger.info( - msg="Got 'ContextWindowExceededError'. No context_window_fallback set. Defaulting \ - to fallbacks, if available.{}".format( - error_message - ) - ) - - e.message += "\n{}".format(error_message) - elif isinstance(e, litellm.ContentPolicyViolationError): - if content_policy_fallbacks is not None: - fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=content_policy_fallbacks, - model_group=model_group, - ) - ) - if fallback_model_group is None: - raise original_exception - - input_kwargs.update( - { - "fallback_model_group": fallback_model_group, - "original_model_group": original_model_group, - } - ) - - response = await run_async_fallback( - *args, - **input_kwargs, - ) - return response - else: - error_message = "model={}. content_policy_fallback={}. fallbacks={}.\n\nSet 'content_policy_fallback' - https://docs.litellm.ai/docs/routing#fallbacks".format( - model_group, content_policy_fallbacks, fallbacks - ) - verbose_router_logger.info( - msg="Got 'ContentPolicyViolationError'. No content_policy_fallback set. Defaulting \ - to fallbacks, if available.{}".format( - error_message - ) - ) - - e.message += "\n{}".format(error_message) - if fallbacks is not None and model_group is not None: - verbose_router_logger.debug(f"inside model fallbacks: {fallbacks}") - ( - fallback_model_group, - generic_fallback_idx, - ) = get_fallback_model_group( - fallbacks=fallbacks, # if fallbacks = [{"gpt-3.5-turbo": ["claude-3-haiku"]}] - model_group=cast(str, model_group), - ) - ## if none, check for generic fallback - if ( - fallback_model_group is None - and generic_fallback_idx is not None - ): - fallback_model_group = fallbacks[generic_fallback_idx]["*"] - - if fallback_model_group is None: - verbose_router_logger.info( - f"No fallback model group found for original model_group={model_group}. Fallbacks={fallbacks}" - ) - if hasattr(original_exception, "message"): - original_exception.message += f"No fallback model group found for original model_group={model_group}. Fallbacks={fallbacks}" # type: ignore - raise original_exception - - input_kwargs.update( - { - "fallback_model_group": fallback_model_group, - "original_model_group": original_model_group, - } - ) - - response = await run_async_fallback( - *args, - **input_kwargs, - ) - - return response - except Exception as new_exception: - parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) - verbose_router_logger.error( - "litellm.router.py::async_function_with_fallbacks() - Error occurred while trying to do fallbacks - {}\n{}\n\nDebug Information:\nCooldown Deployments={}".format( - str(new_exception), - traceback.format_exc(), - await _async_get_cooldown_deployments_with_debug_info( - litellm_router_instance=self, - parent_otel_span=parent_otel_span, - ), - ) - ) - fallback_failure_exception_str = str(new_exception) - - if hasattr(original_exception, "message"): - # add the available fallbacks to the exception - original_exception.message += ". Received Model Group={}\nAvailable Model Group Fallbacks={}".format( # type: ignore - model_group, - fallback_model_group, - ) - if len(fallback_failure_exception_str) > 0: - original_exception.message += ( # type: ignore - "\nError doing the fallback: {}".format( - fallback_failure_exception_str - ) - ) - - raise original_exception + return await self.async_function_with_fallbacks_common_utils( + e, + disable_fallbacks, + fallbacks, + context_window_fallbacks, + content_policy_fallbacks, + model_group, + args, + kwargs, + ) def _handle_mock_testing_fallbacks( self, @@ -6008,6 +6179,7 @@ class Router: "context_window_fallbacks", "model_group_retry_policy", "retry_policy", + "model_group_alias", ] for var in vars_to_include: @@ -6037,6 +6209,7 @@ class Router: "fallbacks", "context_window_fallbacks", "model_group_retry_policy", + "model_group_alias", ] _int_settings = [ diff --git a/litellm/types/router.py b/litellm/types/router.py index 1107620533..864fdbf79b 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -89,6 +89,7 @@ class UpdateRouterConfig(BaseModel): retry_after: Optional[float] = None fallbacks: Optional[List[dict]] = None context_window_fallbacks: Optional[List[dict]] = None + model_group_alias: Optional[Dict[str, Union[str, Dict]]] = {} model_config = ConfigDict(protected_namespaces=()) @@ -209,7 +210,6 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): model_info: Optional[Dict] = None mock_response: Optional[Union[str, ModelResponse, Exception, Any]] = None - # auto-router params auto_router_config_path: Optional[str] = None auto_router_config: Optional[str] = None @@ -343,7 +343,7 @@ class LiteLLM_Params(GenericLiteLLMParams): if max_retries is not None and isinstance(max_retries, str): max_retries = int(max_retries) # cast to int args["max_retries"] = max_retries - super().__init__(**{ **args, **params }) + super().__init__(**{**args, **params}) def __contains__(self, key): # Define custom behavior for the 'in' operator @@ -776,9 +776,11 @@ class MockRouterTestingParams: ), ) + class ModelGroupSettings(BaseModel): forward_client_headers_to_llm_api: Optional[List[str]] = None + class PreRoutingHookResponse(BaseModel): """ Response object from the pre-routing hook. @@ -787,5 +789,6 @@ class PreRoutingHookResponse(BaseModel): Add fields that you expect to be modified by the pre-routing hook. """ + model: str - messages: Optional[List[Dict[str, str]]] \ No newline at end of file + messages: Optional[List[Dict[str, str]]] diff --git a/litellm/types/utils.py b/litellm/types/utils.py index f0bc2aa732..8417b5910a 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2320,6 +2320,7 @@ class LlmProviders(str, Enum): RECRAFT = "recraft" AUTO_ROUTER = "auto_router" VERCEL_AI_GATEWAY = "vercel_ai_gateway" + DOTPROMPT = "dotprompt" # Create a set of all provider values for quick lookup LlmProvidersSet = {provider.value for provider in LlmProviders} diff --git a/litellm/utils.py b/litellm/utils.py index 21499d5056..ab34e0eb18 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -681,7 +681,9 @@ def function_setup( # noqa: PLR0915 if add_breadcrumb: try: - details_to_log = copy.deepcopy(kwargs) + from litellm.litellm_core_utils.core_helpers import safe_deep_copy + + details_to_log = safe_deep_copy(kwargs) except Exception: details_to_log = kwargs diff --git a/pyproject.toml b/pyproject.toml index 02718554e4..33e554025c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm" -version = "1.74.12" +version = "1.74.14" description = "Library to easily interface with LLM API providers" authors = ["BerriAI"] license = "MIT" @@ -152,7 +152,7 @@ requires = ["poetry-core", "wheel"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "1.74.12" +version = "1.74.14" version_files = [ "pyproject.toml:^version" ] diff --git a/tests/code_coverage_tests/ban_copy_deepcopy_kwargs.py b/tests/code_coverage_tests/ban_copy_deepcopy_kwargs.py new file mode 100644 index 0000000000..77c019c7dd --- /dev/null +++ b/tests/code_coverage_tests/ban_copy_deepcopy_kwargs.py @@ -0,0 +1,142 @@ +import ast +import os + + +class CopyDeepcopyKwargsDetector(ast.NodeVisitor): + def __init__(self): + self.violations = [] + + def visit_Call(self, node): + # Check if this is a copy.deepcopy call + if self._is_copy_deepcopy_call(node): + # Check if any argument contains 'kwargs' in its name + for arg in node.args: + if self._is_kwargs_related(arg): + # Get line number and argument name for reporting + arg_name = self._get_arg_name(arg) + self.violations.append( + { + "line": node.lineno, + "arg_name": arg_name, + "full_call": ( + ast.unparse(node) + if hasattr(ast, "unparse") + else str(node) + ), + } + ) + + self.generic_visit(node) + + def _is_copy_deepcopy_call(self, node): + """Check if this is a copy.deepcopy() call""" + if isinstance(node.func, ast.Attribute): + # Case: copy.deepcopy() + if ( + isinstance(node.func.value, ast.Name) + and node.func.value.id == "copy" + and node.func.attr == "deepcopy" + ): + return True + elif isinstance(node.func, ast.Name): + # Case: deepcopy() (if imported as 'from copy import deepcopy') + if node.func.id == "deepcopy": + return True + return False + + def _is_kwargs_related(self, arg): + """Check if the argument is kwargs-related""" + if isinstance(arg, ast.Name): + # Direct variable names containing 'kwargs' + return "kwargs" in arg.id.lower() + elif isinstance(arg, ast.Subscript): + # Handle cases like kwargs['key'] + if isinstance(arg.value, ast.Name): + return "kwargs" in arg.value.id.lower() + elif isinstance(arg, ast.Attribute): + # Handle cases like self.kwargs + return "kwargs" in arg.attr.lower() + return False + + def _get_arg_name(self, arg): + """Get a readable name for the argument""" + if isinstance(arg, ast.Name): + return arg.id + elif isinstance(arg, ast.Subscript) and isinstance(arg.value, ast.Name): + return f"{arg.value.id}[...]" + elif isinstance(arg, ast.Attribute): + return f"...{arg.attr}" + else: + return "unknown_kwargs_variable" + + +def find_copy_deepcopy_kwargs_in_file(file_path): + """Find copy.deepcopy usage with kwargs in a single file""" + try: + with open(file_path, "r", encoding="utf-8") as file: + tree = ast.parse(file.read(), filename=file_path) + detector = CopyDeepcopyKwargsDetector() + detector.visit(tree) + return detector.violations + except Exception as e: + print(f"Error parsing {file_path}: {e}") + return [] + + +def find_copy_deepcopy_kwargs_in_directory(directory): + """Find copy.deepcopy usage with kwargs in all Python files in directory""" + violations = {} + + for root, _, files in os.walk(directory): + for file in files: + if file.endswith(".py"): + file_path = os.path.join(root, file) + print(f"Checking file: {file_path}") + file_violations = find_copy_deepcopy_kwargs_in_file(file_path) + if file_violations: + violations[file_path] = file_violations + + return violations + + +if __name__ == "__main__": + # Check for copy.deepcopy(kwargs) usage in the litellm directory + directory_path = "./litellm" + violations = find_copy_deepcopy_kwargs_in_directory(directory_path) + + print("\n" + "=" * 80) + print("COPY.DEEPCOPY KWARGS VIOLATIONS FOUND:") + print("=" * 80) + + if violations: + total_violations = 0 + for file_path, file_violations in violations.items(): + print(f"\n📁 File: {file_path}") + for violation in file_violations: + total_violations += 1 + print( + f" ❌ Line {violation['line']}: copy.deepcopy({violation['arg_name']})" + ) + print(f" Full call: {violation['full_call']}") + + print(f"\n{'='*80}") + print(f"🚨 TOTAL VIOLATIONS: {total_violations}") + print("🚨 USE safe_deep_copy() INSTEAD OF copy.deepcopy() FOR KWARGS!") + print("🚨 Available imports:") + print(" - from litellm.proxy.utils import safe_deep_copy") + print(" - from litellm.litellm_core_utils.core_helpers import safe_deep_copy") + print("=" * 80) + + # Get first violation for the exception message + first_file = list(violations.keys())[0] + first_violation = violations[first_file][0] + + raise Exception( + f"🚨 Found {total_violations} copy.deepcopy(kwargs) violations! " + f"First violation: {first_file}:{first_violation['line']} - " + f"copy.deepcopy({first_violation['arg_name']}). " + f"Use safe_deep_copy() instead to handle non-serializable objects like OTEL spans." + ) + else: + print("✅ No copy.deepcopy(kwargs) violations found!") + print("✅ All kwargs copying appears to use safe_deep_copy() correctly.") diff --git a/tests/local_testing/test_router_batch_completion.py b/tests/local_testing/test_router_batch_completion.py index 6fedb82a55..534db7ed01 100644 --- a/tests/local_testing/test_router_batch_completion.py +++ b/tests/local_testing/test_router_batch_completion.py @@ -130,6 +130,7 @@ async def test_batch_completion_fastest_response_unit_test(): @pytest.mark.asyncio async def test_batch_completion_fastest_response_streaming(): litellm.set_verbose = True + litellm._turn_on_debug() router = litellm.Router( model_list=[ diff --git a/tests/test_litellm/integrations/dotprompt/chat_prompt.prompt b/tests/test_litellm/integrations/dotprompt/chat_prompt.prompt new file mode 100644 index 0000000000..ef3c76d300 --- /dev/null +++ b/tests/test_litellm/integrations/dotprompt/chat_prompt.prompt @@ -0,0 +1,13 @@ +--- +model: gpt-4 +temperature: 0.7 +max_tokens: 150 +input: + schema: + user_message: string + system_context?: string +--- + +{% if system_context %}System: {{system_context}} + +{% endif %}User: {{user_message}} \ No newline at end of file diff --git a/tests/test_litellm/integrations/dotprompt/coding_assistant.prompt b/tests/test_litellm/integrations/dotprompt/coding_assistant.prompt new file mode 100644 index 0000000000..a128c7828d --- /dev/null +++ b/tests/test_litellm/integrations/dotprompt/coding_assistant.prompt @@ -0,0 +1,33 @@ +--- +model: claude-3-5-sonnet-20241022 +temperature: 0.2 +max_tokens: 2000 +input: + schema: + language: string + task: string + code?: string + requirements?: array +output: + format: text +--- + +You are a helpful coding assistant. {% if language %}Focus on {{language}} programming.{% endif %} + +Task: {{task}} + +{% if code %} +Current code: +```{{language}} +{{code}} +``` +{% endif %} + +{% if requirements %} +Requirements: +{% for req in requirements %} +- {{req}} +{% endfor %} +{% endif %} + +Please provide a clear and well-documented solution. \ No newline at end of file diff --git a/tests/test_litellm/integrations/dotprompt/sample_prompt.prompt b/tests/test_litellm/integrations/dotprompt/sample_prompt.prompt new file mode 100644 index 0000000000..7329e92c4c --- /dev/null +++ b/tests/test_litellm/integrations/dotprompt/sample_prompt.prompt @@ -0,0 +1,16 @@ +--- +model: gemini/gemini-1.5-pro +input: + schema: + text: string +output: + format: json + schema: + title?: string, the title of the article if it has one + summary: string, a 3-sentence summary of the text + tags?(array, a list of string tag category for the text): string +--- + +Extract the requested information from the given text. If a piece of information is not present, omit that field from the output. + +Text: {{text}} diff --git a/tests/test_litellm/integrations/dotprompt/test_dotprompt_manager.py b/tests/test_litellm/integrations/dotprompt/test_dotprompt_manager.py new file mode 100644 index 0000000000..28f7e85bf8 --- /dev/null +++ b/tests/test_litellm/integrations/dotprompt/test_dotprompt_manager.py @@ -0,0 +1,238 @@ +import json +import os +import sys +import tempfile +from pathlib import Path + +import pytest +from fastapi.testclient import TestClient + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + + +from unittest.mock import MagicMock, patch + +import litellm +from litellm.integrations.dotprompt import DotpromptManager +from litellm.types.utils import StandardCallbackDynamicParams + + +def test_dotprompt_manager_initialization(): + """Test basic DotpromptManager initialization.""" + prompt_dir = "." # Current directory when running from tests/test_litellm/prompts + manager = DotpromptManager(prompt_dir) + + assert manager.integration_name == "dotprompt" + assert manager.prompt_directory == prompt_dir + + +def test_should_run_prompt_management(): + """Test should_run_prompt_management method.""" + prompt_dir = "." + manager = DotpromptManager(prompt_dir) + + # Test with existing prompt + assert ( + manager.should_run_prompt_management( + "sample_prompt", StandardCallbackDynamicParams() + ) + == True + ) + + # Test with non-existing prompt + assert ( + manager.should_run_prompt_management( + "nonexistent_prompt", StandardCallbackDynamicParams() + ) + == False + ) + + +def test_convert_to_messages_simple(): + """Test converting simple text to messages.""" + prompt_dir = "." + manager = DotpromptManager(prompt_dir) + + # Test simple text + messages = manager._convert_to_messages("Hello world!") + assert len(messages) == 1 + assert messages[0]["role"] == "user" + assert messages[0]["content"] == "Hello world!" + + +def test_convert_to_messages_with_roles(): + """Test converting text with role prefixes to messages.""" + prompt_dir = "." + manager = DotpromptManager(prompt_dir) + + # Test text with role prefixes + content = """System: You are a helpful assistant. + +User: What is the capital of France?""" + + messages = manager._convert_to_messages(content) + assert len(messages) == 2 + + assert messages[0]["role"] == "system" + assert messages[0]["content"] == "You are a helpful assistant." + + assert messages[1]["role"] == "user" + assert messages[1]["content"] == "What is the capital of France?" + + +def test_compile_prompt_helper(): + """Test the _compile_prompt_helper method.""" + prompt_dir = "." + manager = DotpromptManager(prompt_dir) + + # Test compiling a simple prompt + result = manager._compile_prompt_helper( + prompt_id="sample_prompt", + prompt_variables={"text": "This is a test article."}, + dynamic_callback_params=StandardCallbackDynamicParams(), + ) + + assert result["prompt_id"] == "sample_prompt" + assert result["prompt_template_model"] == "gemini/gemini-1.5-pro" + assert len(result["prompt_template"]) >= 1 + assert "This is a test article." in result["prompt_template"][0]["content"] + + +def test_compile_prompt_helper_with_chat_format(): + """Test compiling a prompt that generates role-based messages.""" + prompt_dir = "." + manager = DotpromptManager(prompt_dir) + + # Test with chat_prompt that has system context + result = manager._compile_prompt_helper( + prompt_id="chat_prompt", + prompt_variables={ + "user_message": "Hello there!", + "system_context": "You are a helpful assistant.", + }, + dynamic_callback_params=StandardCallbackDynamicParams(), + ) + + assert result["prompt_id"] == "chat_prompt" + assert result["prompt_template_model"] == "gpt-4" + assert len(result["prompt_template"]) == 2 + + # Should have system message first + assert result["prompt_template"][0]["role"] == "system" + assert "You are a helpful assistant." in result["prompt_template"][0]["content"] + + # Then user message + assert result["prompt_template"][1]["role"] == "user" + assert "Hello there!" in result["prompt_template"][1]["content"] + + +def test_extract_optional_params(): + """Test extracting optional parameters from template metadata.""" + prompt_dir = "." + manager = DotpromptManager(prompt_dir) + + # Get a template with optional params + template = manager.prompt_manager.get_prompt("chat_prompt") + params = manager._extract_optional_params(template) + + assert "temperature" in params + assert params["temperature"] == 0.7 + assert "max_tokens" in params + assert params["max_tokens"] == 150 + + +def test_error_handling(): + """Test error handling for invalid prompts.""" + prompt_dir = "." + manager = DotpromptManager(prompt_dir) + + # Test with non-existent prompt + with pytest.raises(ValueError, match="Prompt 'nonexistent' not found"): + manager._compile_prompt_helper( + prompt_id="nonexistent", + prompt_variables={}, + dynamic_callback_params=StandardCallbackDynamicParams(), + ) + + +def test_integration_with_prompt_management(): + """Test integration with the prompt management system.""" + with tempfile.TemporaryDirectory() as temp_dir: + # Create a test prompt + prompt_file = Path(temp_dir) / "test_integration.prompt" + prompt_file.write_text( + """--- +model: gpt-3.5-turbo +temperature: 0.5 +--- +System: You are a {{role}}. + +User: {{question}}""" + ) + + manager = DotpromptManager(temp_dir) + + # Test should_run_prompt_management + assert ( + manager.should_run_prompt_management( + "test_integration", StandardCallbackDynamicParams() + ) + == True + ) + + # Test compile_prompt_helper + result = manager._compile_prompt_helper( + prompt_id="test_integration", + prompt_variables={"role": "helpful assistant", "question": "What is AI?"}, + dynamic_callback_params=StandardCallbackDynamicParams(), + ) + + assert result["prompt_template_model"] == "gpt-3.5-turbo" + assert result["prompt_template_optional_params"]["temperature"] == 0.5 + assert len(result["prompt_template"]) == 2 + + assert result["prompt_template"][0]["role"] == "system" + assert "helpful assistant" in result["prompt_template"][0]["content"] + + assert result["prompt_template"][1]["role"] == "user" + assert "What is AI?" in result["prompt_template"][1]["content"] + + +def test_set_prompt_directory(): + """Test setting and changing prompt directory.""" + with tempfile.TemporaryDirectory() as temp_dir: + manager = DotpromptManager(temp_dir) + + # Initially should be empty + assert not manager.should_run_prompt_management( + "test_prompt", StandardCallbackDynamicParams() + ) + + # Create a prompt file + prompt_file = Path(temp_dir) / "test_prompt.prompt" + prompt_file.write_text("Hello {{name}}!") + + # Set directory to force reload + manager.set_prompt_directory(temp_dir) + + # Now should find the prompt + assert manager.should_run_prompt_management( + "test_prompt", StandardCallbackDynamicParams() + ) + + +def test_no_prompt_directory_error(): + """Test error when no prompt directory is set.""" + manager = DotpromptManager(None) + + # should_run_prompt_management returns False when there's an error + result = manager.should_run_prompt_management( + "any_prompt", StandardCallbackDynamicParams() + ) + assert result == False + + # But accessing prompt_manager property should raise an error + with pytest.raises(ValueError, match="prompt_directory must be set"): + _ = manager.prompt_manager diff --git a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py new file mode 100644 index 0000000000..be5ab55166 --- /dev/null +++ b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py @@ -0,0 +1,246 @@ +import json +import os +import sys +import tempfile +from pathlib import Path + +import pytest +from fastapi.testclient import TestClient + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + + +from unittest.mock import MagicMock, patch + +import litellm +from litellm.integrations.dotprompt.prompt_manager import PromptManager, PromptTemplate + + +def test_prompt_manager_initialization(): + """Test basic PromptManager initialization and loading.""" + # Test with the existing prompts directory + prompt_dir = "." # Current directory when running from tests/test_litellm/prompts + manager = PromptManager(prompt_dir) + + # Should have loaded at least the sample prompts + assert len(manager.prompts) >= 3 + assert "sample_prompt" in manager.prompts + assert "chat_prompt" in manager.prompts + assert "coding_assistant" in manager.prompts + + +def test_prompt_template_creation(): + """Test PromptTemplate creation and metadata extraction.""" + metadata = { + "model": "gpt-4", + "temperature": 0.7, + "input": {"schema": {"text": "string"}}, + "output": {"format": "json"}, + } + + template = PromptTemplate( + content="Hello {{name}}!", metadata=metadata, template_id="test_template" + ) + + assert template.content == "Hello {{name}}!" + assert template.model == "gpt-4" + assert template.optional_params["temperature"] == 0.7 + assert template.input_schema == {"text": "string"} + assert template.output_format == "json" + + +def test_render_simple_template(): + """Test rendering a simple template with variables.""" + prompt_dir = "." # Current directory when running from tests/test_litellm/prompts + manager = PromptManager(prompt_dir) + + # Test sample_prompt rendering + rendered = manager.render( + "sample_prompt", {"text": "This is a test article about AI."} + ) + + expected_content = "Extract the requested information from the given text. If a piece of information is not present, omit that field from the output.\n\nText: This is a test article about AI." + assert rendered == expected_content + + +def test_render_chat_prompt(): + """Test rendering the chat prompt with conditional content.""" + prompt_dir = "." # Current directory when running from tests/test_litellm/prompts + manager = PromptManager(prompt_dir) + + # Test with system context + rendered = manager.render( + "chat_prompt", + { + "user_message": "Hello there!", + "system_context": "You are a helpful assistant.", + }, + ) + + assert "System: You are a helpful assistant." in rendered + assert "User: Hello there!" in rendered + + # Test without system context + rendered_no_system = manager.render("chat_prompt", {"user_message": "Hello there!"}) + + assert "System:" not in rendered_no_system + assert "User: Hello there!" in rendered_no_system + + +def test_render_coding_assistant(): + """Test rendering the coding assistant prompt with complex logic.""" + prompt_dir = "." # Current directory when running from tests/test_litellm/prompts + manager = PromptManager(prompt_dir) + + rendered = manager.render( + "coding_assistant", + { + "language": "Python", + "task": "Create a function to calculate fibonacci numbers", + "code": "def fib(n):\n pass", + "requirements": ["Use recursion", "Handle edge cases", "Add documentation"], + }, + ) + + assert "Focus on Python programming." in rendered + assert "Create a function to calculate fibonacci numbers" in rendered + assert "def fib(n):" in rendered + assert "Use recursion" in rendered + assert "Handle edge cases" in rendered + assert "Add documentation" in rendered + + +def test_input_validation(): + """Test input validation against schema.""" + # Create a temporary directory with a test prompt + with tempfile.TemporaryDirectory() as temp_dir: + prompt_file = Path(temp_dir) / "test_validation.prompt" + prompt_file.write_text( + """--- +input: + schema: + name: string + age: integer + active: boolean +--- +Hello {{name}}, you are {{age}} years old and {'active' if active else 'inactive'}.""" + ) + + manager = PromptManager(temp_dir) + + # Valid input should work + rendered = manager.render( + "test_validation", {"name": "Alice", "age": 30, "active": True} + ) + assert "Hello Alice, you are 30 years old" in rendered + + # Invalid type should raise error + with pytest.raises(ValueError, match="Invalid type for field 'age'"): + manager.render( + "test_validation", + { + "name": "Alice", + "age": "thirty", # string instead of int + "active": True, + }, + ) + + +def test_prompt_not_found(): + """Test error handling for non-existent prompts.""" + prompt_dir = "." # Current directory when running from tests/test_litellm/prompts + manager = PromptManager(prompt_dir) + + with pytest.raises(KeyError, match="Prompt 'nonexistent' not found"): + manager.render("nonexistent", {"some": "variable"}) + + +def test_list_prompts(): + """Test listing available prompts.""" + prompt_dir = "." # Current directory when running from tests/test_litellm/prompts + manager = PromptManager(prompt_dir) + + prompts = manager.list_prompts() + assert isinstance(prompts, list) + assert "sample_prompt" in prompts + assert "chat_prompt" in prompts + assert "coding_assistant" in prompts + + +def test_get_prompt_metadata(): + """Test retrieving prompt metadata.""" + prompt_dir = "." # Current directory when running from tests/test_litellm/prompts + manager = PromptManager(prompt_dir) + + metadata = manager.get_prompt_metadata("sample_prompt") + assert metadata is not None + assert metadata["model"] == "gemini/gemini-1.5-pro" + assert "input" in metadata + assert "output" in metadata + + +def test_add_prompt_programmatically(): + """Test adding prompts programmatically.""" + prompt_dir = "." # Current directory when running from tests/test_litellm/prompts + manager = PromptManager(prompt_dir) + + initial_count = len(manager.prompts) + + manager.add_prompt( + "dynamic_prompt", + "Hello {{name}}! Welcome to {{place}}.", + {"model": "gpt-3.5-turbo", "temperature": 0.5}, + ) + + assert len(manager.prompts) == initial_count + 1 + assert "dynamic_prompt" in manager.prompts + + rendered = manager.render("dynamic_prompt", {"name": "World", "place": "Earth"}) + assert rendered == "Hello World! Welcome to Earth." + + +def test_frontmatter_parsing(): + """Test YAML frontmatter parsing.""" + # Create a temporary directory with a test prompt + with tempfile.TemporaryDirectory() as temp_dir: + # Test with frontmatter + prompt_with_frontmatter = Path(temp_dir) / "with_frontmatter.prompt" + prompt_with_frontmatter.write_text( + """--- +model: gpt-4 +temperature: 0.8 +input: + schema: + topic: string +--- +Write about {{topic}}.""" + ) + + # Test without frontmatter + prompt_without_frontmatter = Path(temp_dir) / "without_frontmatter.prompt" + prompt_without_frontmatter.write_text("Simple template: {{message}}") + + manager = PromptManager(temp_dir) + + # Check frontmatter was parsed correctly + with_meta = manager.get_prompt("with_frontmatter") + assert with_meta.model == "gpt-4" + assert with_meta.optional_params["temperature"] == 0.8 + + # Check template without frontmatter still works + without_meta = manager.get_prompt("without_frontmatter") + assert without_meta.metadata == {} + + rendered = manager.render("without_frontmatter", {"message": "Hello!"}) + assert rendered == "Simple template: Hello!" + + +def test_prompt_main(): + """ + Integration test placeholder for litellm completion integration. + This would be implemented once the PromptManager is integrated with litellm. + """ + # TODO: Implement once PromptManager is integrated with litellm completion + pass diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 1d349a44e5..980693aa73 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -125,3 +125,21 @@ def test_handle_any_messages_to_chat_completion_str_messages_conversion_complex( result = handle_any_messages_to_chat_completion_str_messages_conversion(message) assert len(result) == 1 assert result[0]["input"] == json.dumps(message) + + +def test_convert_prefix_message_to_non_prefix_messages(): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_prefix_message_to_non_prefix_messages, + ) + + messages = [ + {"role": "assistant", "content": "value", "prefix": True}, + ] + result = convert_prefix_message_to_non_prefix_messages(messages) + assert result == [ + { + "role": "system", + "content": "You are a helpful assistant. You are given a message and you need to respond to it. You are also given a generated content. You need to respond to the message in continuation of the generated content. Do not repeat the same content. Your response should be in continuation of this text: ", + }, + {"role": "assistant", "content": "value"}, + ] diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 0253f00c7e..5cc30f3918 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -87,44 +87,33 @@ def test_convert_to_azure_openai_messages(): """Test coverting image_url to azure_openai spec""" from typing import List + + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_azure_openai_messages, + ) from litellm.types.llms.openai import AllMessageValues - from litellm.litellm_core_utils.prompt_templates.factory import convert_to_azure_openai_messages input: List[AllMessageValues] = [ { "role": "user", "content": [ - { - "type": "text", - "text": "What is in this image?" - }, - { - "type": "image_url", - "image_url": "www.mock.com" - } - ] + {"type": "text", "text": "What is in this image?"}, + {"type": "image_url", "image_url": "www.mock.com"}, + ], } ] expected_content = [ - { - "type": "text", - "text": "What is in this image?" - }, - { - "type": "image_url", - "image_url": {"url": "www.mock.com"} - } + {"type": "text", "text": "What is in this image?"}, + {"type": "image_url", "image_url": {"url": "www.mock.com"}}, ] output = convert_to_azure_openai_messages(input) - content = output[0].get('content') + content = output[0].get("content") assert content == expected_content - - def test_bedrock_validate_format_image_or_video(): """Test the _validate_format method for images, videos, and documents""" @@ -405,12 +394,44 @@ def test_unpack_defs_resolves_nested_ref_inside_anyof_items(): unpack_defs(schema, schema["$defs"]) # Extract the items schema after unpacking - items_schema = ( - schema["properties"]["vatAmounts"]["anyOf"][0]["items"] - ) + items_schema = schema["properties"]["vatAmounts"]["anyOf"][0]["items"] # Assertions: items_schema should now be the resolved object, not an empty dict - assert isinstance(items_schema, dict), "Items schema should be a dict after unpacking" + assert isinstance( + items_schema, dict + ), "Items schema should be a dict after unpacking" assert items_schema.get("type") == "object" # Ensure essential properties are present assert set(items_schema.get("properties", {}).keys()) == {"vatRate", "vatAmount"} + + +def test_convert_gemini_messages(): + """ + Handle 'content' not being present in the message - https://github.com/BerriAI/litellm/issues/13169 + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_result, + ) + from litellm.types.llms.openai import ChatCompletionToolMessage + + message = ChatCompletionToolMessage( + role="tool", + tool_call_id="call_d5b2e3fe-d2c0-451d-b034-cf4fbb22e66c", + ) + last_message_with_tool_calls = { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_d5b2e3fe-d2c0-451d-b034-cf4fbb22e66c", + "type": "function", + "index": 0, + "function": {"name": "tool_MAX_Data__get_issues", "arguments": "{}"}, + } + ], + } + + convert_to_gemini_tool_call_result( + message=message, + last_message_with_tool_calls=last_message_with_tool_calls, + ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py index e4c9d056dd..12d84f9530 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py @@ -280,7 +280,7 @@ class TestPanwAirsResponseScanning: mock_response = {"action": "allow", "category": "benign"} with patch.object(handler, "_call_panw_api", return_value=mock_response): - result = await handler.async_post_call_hook( + result = await handler.async_post_call_success_hook( data=request_data, user_api_key_dict=user_api_key_dict, response=safe_response, @@ -299,7 +299,7 @@ class TestPanwAirsResponseScanning: with patch.object(handler, "_call_panw_api", return_value=mock_response): with pytest.raises(HTTPException) as exc_info: - await handler.async_post_call_hook( + await handler.async_post_call_success_hook( data=request_data, user_api_key_dict=user_api_key_dict, response=harmful_response, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index cacdd51ee9..4795fdc372 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -2,7 +2,7 @@ import copy import json import os import sys -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi.testclient import TestClient @@ -13,6 +13,7 @@ sys.path.insert( import litellm +from litellm.router_utils.fallback_event_handlers import run_async_fallback def test_update_kwargs_does_not_mutate_defaults_and_merges_metadata(): @@ -1064,3 +1065,321 @@ def test_router_get_model_access_groups_team_only_models(): model_name="gpt-3.5-turbo", team_id="team_1" ) assert list(access_groups.keys()) == ["default-models"] + + +@pytest.mark.asyncio +async def test_acompletion_streaming_iterator(): + """Test _acompletion_streaming_iterator for normal streaming and fallback behavior.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.exceptions import MidStreamFallbackError + from litellm.types.utils import ModelResponseStream + + # Helper class for creating async iterators + class AsyncIterator: + def __init__(self, items, error_after=None): + self.items = items + self.index = 0 + self.error_after = error_after + + def __aiter__(self): + return self + + async def __anext__(self): + if self.error_after is not None and self.index >= self.error_after: + raise self.error_after + if self.index >= len(self.items): + raise StopAsyncIteration + item = self.items[self.index] + self.index += 1 + return item + + # Set up router with fallback configuration + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake-key-1"}, + }, + { + "model_name": "gpt-3.5-turbo", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake-key-2"}, + }, + ], + fallbacks=[{"gpt-4": ["gpt-3.5-turbo"]}], + set_verbose=True, + ) + + # Test data + messages = [{"role": "user", "content": "Hello"}] + initial_kwargs = {"model": "gpt-4", "stream": True, "temperature": 0.7} + + # Test 1: Successful streaming (no errors) + print("\n=== Test 1: Successful streaming ===") + + # Mock successful streaming response + mock_chunks = [ + MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello"))]), + MagicMock(choices=[MagicMock(delta=MagicMock(content=" there"))]), + MagicMock(choices=[MagicMock(delta=MagicMock(content="!"))]), + ] + + mock_response = AsyncIterator(mock_chunks) + + setattr(mock_response, "model", "gpt-4") + setattr(mock_response, "custom_llm_provider", "openai") + setattr(mock_response, "logging_obj", MagicMock()) + + result = await router._acompletion_streaming_iterator( + model_response=mock_response, messages=messages, initial_kwargs=initial_kwargs + ) + + # Collect streamed chunks + collected_chunks = [] + async for chunk in result: + collected_chunks.append(chunk) + + assert len(collected_chunks) == 3 + assert all(chunk in mock_chunks for chunk in collected_chunks) + print("✓ Successfully streamed all chunks") + + # Test 2: MidStreamFallbackError with fallback + print("\n=== Test 2: MidStreamFallbackError with fallback ===") + + # Create error that should trigger after first chunk + error = MidStreamFallbackError( + message="Connection lost", + model="gpt-4", + llm_provider="openai", + generated_content="Hello", + ) + + class AsyncIteratorWithError: + def __init__(self, items, error_after_index): + self.items = items + self.index = 0 + self.error_after_index = error_after_index + + def __aiter__(self): + return self + + async def __anext__(self): + if self.index >= len(self.items): + raise StopAsyncIteration + if self.index == self.error_after_index: + raise error + item = self.items[self.index] + self.index += 1 + return item + + mock_error_response = AsyncIteratorWithError( + mock_chunks, 1 + ) # Error after first chunk + + setattr(mock_error_response, "model", "gpt-4") + setattr(mock_error_response, "custom_llm_provider", "openai") + setattr(mock_error_response, "logging_obj", MagicMock()) + + # Mock the fallback response + fallback_chunks = [ + MagicMock(choices=[MagicMock(delta=MagicMock(content=" world"))]), + MagicMock(choices=[MagicMock(delta=MagicMock(content="!"))]), + ] + + mock_fallback_response = AsyncIterator(fallback_chunks) + + # Mock the fallback function + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=mock_fallback_response, + ) as mock_fallback_utils: + + collected_chunks = [] + result = await router._acompletion_streaming_iterator( + model_response=mock_error_response, + messages=messages, + initial_kwargs=initial_kwargs, + ) + + async for chunk in result: + collected_chunks.append(chunk) + + # Verify fallback was called + assert mock_fallback_utils.called + call_args = mock_fallback_utils.call_args + + # Check that generated content was added to messages + fallback_kwargs = call_args.kwargs["kwargs"] + modified_messages = fallback_kwargs["messages"] + + # Should have original message + system message + assistant message with prefix + assert len(modified_messages) == 3 + assert modified_messages[0] == {"role": "user", "content": "Hello"} + assert modified_messages[1]["role"] == "system" + assert "continuation" in modified_messages[1]["content"] + assert modified_messages[2]["role"] == "assistant" + assert modified_messages[2]["content"] == "Hello" + assert modified_messages[2]["prefix"] == True + + # Verify fallback parameters + assert call_args.kwargs["disable_fallbacks"] == False + assert call_args.kwargs["model_group"] == "gpt-4" + + # Should get original chunk + fallback chunks + assert len(collected_chunks) == 3 # 1 original + 2 fallback + print("✓ Fallback system called correctly with proper message modification") + + # Test 3: Fallback failure + print("\n=== Test 3: Fallback failure ===") + + mock_error_response_2 = AsyncIteratorWithError(mock_chunks, 1) # Same error pattern + + # Mock fallback failure + fallback_error = Exception("Fallback also failed") + with patch.object( + router, "async_function_with_fallbacks_common_utils", side_effect=fallback_error + ): + + collected_chunks = [] + original_error = None + setattr(mock_error_response_2, "model", "gpt-4") + setattr(mock_error_response_2, "custom_llm_provider", "openai") + setattr(mock_error_response_2, "logging_obj", MagicMock()) + + try: + result = await router._acompletion_streaming_iterator( + model_response=mock_error_response_2, + messages=messages, + initial_kwargs=initial_kwargs, + ) + + async for chunk in result: + collected_chunks.append(chunk) + except MidStreamFallbackError as e: + original_error = e + + # Should re-raise original MidStreamFallbackError, not fallback error + assert original_error is not None + assert isinstance(original_error, MidStreamFallbackError) + assert original_error.generated_content == "Hello" + print("✓ Original error re-raised when fallback fails") + + print("\n=== All tests passed! ===") + + +@pytest.mark.asyncio +async def test_acompletion_streaming_iterator_edge_cases(): + """Test edge cases for _acompletion_streaming_iterator.""" + from unittest.mock import MagicMock + + from litellm.exceptions import MidStreamFallbackError + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, + } + ], + set_verbose=True, + ) + + messages = [{"role": "user", "content": "Test"}] + initial_kwargs = {"model": "gpt-4", "stream": True} + + # Test: Empty generated content + empty_error = MidStreamFallbackError( + message="Error", + model="gpt-4", + llm_provider="openai", + generated_content="", # Empty content + ) + + class AsyncIteratorImmediateError: + def __aiter__(self): + return self + + async def __anext__(self): + raise empty_error + + mock_response = AsyncIteratorImmediateError() + + # Mock empty fallback response using AsyncIterator + class EmptyAsyncIterator: + def __aiter__(self): + return self + + async def __anext__(self): + raise StopAsyncIteration + + mock_fallback_response = EmptyAsyncIterator() + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=mock_fallback_response, + ) as mock_fallback_utils: + + collected_chunks = [] + async for chunk in router._acompletion_streaming_iterator( + model_response=mock_response, + messages=messages, + initial_kwargs=initial_kwargs, + ): + collected_chunks.append(chunk) + + # Should still call fallback even with empty content + assert mock_fallback_utils.called + fallback_kwargs = mock_fallback_utils.call_args.kwargs["kwargs"] + modified_messages = fallback_kwargs["messages"] + + # Should have assistant message with empty content + assert modified_messages[2]["content"] == "" + print("✓ Handles empty generated content correctly") + + print("✓ Edge case tests passed!") + + +@pytest.mark.asyncio +async def test_async_function_with_fallbacks_common_utils(): + """Test the async_function_with_fallbacks_common_utils method""" + # Create a basic router for testing + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": { + "model": "gpt-3.5-turbo", + }, + } + ], + max_fallbacks=5, + ) + + # Test case 1: disable_fallbacks=True should raise original exception + test_exception = Exception("Test error") + with pytest.raises(Exception, match="Test error"): + await router.async_function_with_fallbacks_common_utils( + e=test_exception, + disable_fallbacks=True, + fallbacks=None, + context_window_fallbacks=None, + content_policy_fallbacks=None, + model_group="gpt-3.5-turbo", + args=(), + kwargs=MagicMock(), + ) + + # Test case 2: original_model_group=None should raise original exception + with pytest.raises(Exception, match="Test error"): + await router.async_function_with_fallbacks_common_utils( + e=test_exception, + disable_fallbacks=False, + fallbacks=None, + context_window_fallbacks=None, + content_policy_fallbacks=None, + model_group="gpt-3.5-turbo", + args=(), + kwargs={}, # No model key + ) diff --git a/ui/litellm-dashboard/src/components/generic_key_value_manager.tsx b/ui/litellm-dashboard/src/components/generic_key_value_manager.tsx new file mode 100644 index 0000000000..520893d14c --- /dev/null +++ b/ui/litellm-dashboard/src/components/generic_key_value_manager.tsx @@ -0,0 +1,339 @@ +import React, { useState, useEffect, useCallback } from "react"; +import { + Card, + Title, + Text, + Table, + TableHead, + TableRow, + TableHeaderCell, + TableCell, + TableBody, +} from "@tremor/react"; +import { message, Input } from "antd"; +import { EditOutlined, DeleteOutlined, SaveOutlined, CloseOutlined } from "@ant-design/icons"; +import { ChevronDownIcon, ChevronRightIcon, PlusCircleIcon } from "@heroicons/react/outline"; + +interface KeyValueItem { + id?: string; + key: string; + value: string; +} + +interface GenericKeyValueManagerProps { + title: string; + description: string; + keyLabel: string; + valueLabel: string; + keyPlaceholder: string; + valuePlaceholder: string; + items: KeyValueItem[]; + onItemsChange: (items: KeyValueItem[]) => void; + onSave?: () => Promise; + showSaveButton?: boolean; + isCollapsible?: boolean; + defaultExpanded?: boolean; + configExample?: React.ReactNode; + additionalActions?: (item: KeyValueItem) => React.ReactNode; +} + +const GenericKeyValueManager: React.FC = ({ + title, + description, + keyLabel, + valueLabel, + keyPlaceholder, + valuePlaceholder, + items, + onItemsChange, + onSave, + showSaveButton = true, + isCollapsible = false, + defaultExpanded = true, + configExample, + additionalActions, +}) => { + const [newKey, setNewKey] = useState(""); + const [newValue, setNewValue] = useState(""); + const [editingItem, setEditingItem] = useState(null); + const [editingKey, setEditingKey] = useState(""); + const [editingValue, setEditingValue] = useState(""); + const [isExpanded, setIsExpanded] = useState(defaultExpanded); + + const generateId = () => Math.random().toString(36).substr(2, 9); + + const handleAddItem = useCallback(() => { + if (newKey.trim() && newValue.trim()) { + const newItem: KeyValueItem = { + id: generateId(), + key: newKey.trim(), + value: newValue.trim(), + }; + onItemsChange([...items, newItem]); + setNewKey(""); + setNewValue(""); + } else { + message.error(`Please provide both ${keyLabel.toLowerCase()} and ${valueLabel.toLowerCase()}`); + } + }, [newKey, newValue, items, onItemsChange, keyLabel, valueLabel]); + + const handleEditItem = useCallback((item: KeyValueItem) => { + setEditingItem({ ...item }); + setEditingKey(item.key); + setEditingValue(item.value); + }, []); + + const handleSaveEdit = useCallback(() => { + if (editingKey.trim() && editingValue.trim()) { + const updatedItems = items.map((item) => + item.id === editingItem?.id ? { ...item, key: editingKey.trim(), value: editingValue.trim() } : item + ); + onItemsChange(updatedItems); + setEditingItem(null); + setEditingKey(""); + setEditingValue(""); + } else { + message.error(`Please provide both ${keyLabel.toLowerCase()} and ${valueLabel.toLowerCase()}`); + } + }, [editingKey, editingValue, items, editingItem, onItemsChange, keyLabel, valueLabel]); + + const handleCancelEdit = useCallback(() => { + setEditingItem(null); + setEditingKey(""); + setEditingValue(""); + }, []); + + const handleDeleteItem = useCallback((id: string) => { + const updatedItems = items.filter((item) => item.id !== id); + onItemsChange(updatedItems); + }, [items, onItemsChange]); + + const handleSave = useCallback(async () => { + if (onSave) { + try { + await onSave(); + } catch (error) { + console.error("Failed to save:", error); + } + } + }, [onSave]); + + const ContentSection = useCallback(() => ( +
+ {/* Add New Item Section */} + + Add New {keyLabel} +
+
+ + setNewKey(e.target.value)} + placeholder={keyPlaceholder} + size="middle" + /> +
+
+ + setNewValue(e.target.value)} + placeholder={valuePlaceholder} + size="middle" + /> +
+
+ +
+
+
+ + {/* Manage Existing Items Section */} + +
+ Manage Existing {keyLabel}s + {showSaveButton && ( + + )} +
+ +
+
+ + + + {keyLabel} + {valueLabel} + Actions + + + + {items.map((item) => ( + + {editingItem && editingItem.id === item.id ? ( + <> + + setEditingKey(e.target.value)} + size="small" + /> + + + setEditingValue(e.target.value)} + size="small" + /> + + +
+ + +
+
+ + ) : ( + <> + + {item.key} + + + {item.value} + + +
+ {additionalActions && additionalActions(item)} + + +
+
+ + )} +
+ ))} + {items.length === 0 && ( + + + No {keyLabel.toLowerCase()}s added yet. Add a new {keyLabel.toLowerCase()} above. + + + )} +
+
+
+
+
+ + {/* Configuration Example */} + {configExample && ( + + Configuration Example + {configExample} + + )} +
+ ), [ + keyLabel, + valueLabel, + keyPlaceholder, + valuePlaceholder, + newKey, + newValue, + items, + editingItem, + editingKey, + editingValue, + showSaveButton, + configExample, + additionalActions, + handleAddItem, + handleSave, + handleEditItem, + handleSaveEdit, + handleCancelEdit, + handleDeleteItem, + ]); + + if (isCollapsible) { + return ( + +
setIsExpanded(!isExpanded)} + > +
+ {title} +

{description}

+
+
+ {isExpanded ? ( + + ) : ( + + )} +
+
+ + {isExpanded && ( +
+ +
+ )} +
+ ); + } + + return ( +
+
+ {title} + {description} +
+
+ +
+
+ ); +}; + +export default GenericKeyValueManager; \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/model_dashboard.tsx b/ui/litellm-dashboard/src/components/model_dashboard.tsx index 98cc3b1430..c66ec76844 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard.tsx @@ -73,6 +73,7 @@ import { ModelDataTable } from "./model_dashboard/table"; import { columns } from "./model_dashboard/columns"; import HealthCheckComponent from "./model_dashboard/HealthCheckComponent"; import PassThroughSettings from "./pass_through_settings"; +import ModelGroupAliasSettings from "./model_group_alias_settings"; import { all_admin_roles } from "@/utils/roles"; import { Table as TableInstance } from "@tanstack/react-table"; @@ -197,6 +198,9 @@ const ModelDashboard: React.FC = ({ const [credentialsList, setCredentialsList] = useState([]); + // Model Group Alias state + const [modelGroupAlias, setModelGroupAlias] = useState<{[key: string]: string}>({}); + // Add state for advanced settings visibility const [showAdvancedSettings, setShowAdvancedSettings] = useState(false); @@ -479,6 +483,8 @@ const ModelDashboard: React.FC = ({ } }; + + useEffect(() => { if (!accessToken || !token || !userRole || !userID) { return; @@ -646,6 +652,10 @@ const ModelDashboard: React.FC = ({ setModelGroupRetryPolicy(model_group_retry_policy); setGlobalRetryPolicy(router_settings.retry_policy); setDefaultRetry(default_retries); + + // Set model group alias + const model_group_alias = router_settings.model_group_alias || {}; + setModelGroupAlias(model_group_alias); } catch (error) { console.error("There was an error fetching the model data", error); } @@ -1095,6 +1105,9 @@ const ModelDashboard: React.FC = ({ {all_admin_roles.includes(userRole) && ( Model Retry Settings )} + {all_admin_roles.includes(userRole) && ( + Model Group Alias + )}
@@ -1859,6 +1872,13 @@ const ModelDashboard: React.FC = ({ Save + + + )} diff --git a/ui/litellm-dashboard/src/components/model_group_alias_settings.tsx b/ui/litellm-dashboard/src/components/model_group_alias_settings.tsx new file mode 100644 index 0000000000..b1131c3611 --- /dev/null +++ b/ui/litellm-dashboard/src/components/model_group_alias_settings.tsx @@ -0,0 +1,370 @@ +import React, { useState, useEffect } from "react"; +import { message } from "antd"; +import { PlusCircleIcon, PencilIcon, TrashIcon, ChevronDownIcon, ChevronRightIcon } from "@heroicons/react/outline"; +import { setCallbacksCall } from "./networking"; +import { + Card, + Title, + Text, + Table, + TableHead, + TableHeaderCell, + TableBody, + TableRow, + TableCell +} from "@tremor/react"; + +interface ModelGroupAliasSettingsProps { + accessToken: string; + initialModelGroupAlias?: { [key: string]: string }; + onAliasUpdate?: (updatedAlias: { [key: string]: string }) => void; +} + +interface AliasItem { + id: string; + aliasName: string; + targetModelGroup: string; +} + +const ModelGroupAliasSettings: React.FC = ({ + accessToken, + initialModelGroupAlias = {}, + onAliasUpdate, +}) => { + const [aliases, setAliases] = useState([]); + const [newAlias, setNewAlias] = useState({ aliasName: "", targetModelGroup: "" }); + const [editingAlias, setEditingAlias] = useState(null); + const [isExpanded, setIsExpanded] = useState(true); + + useEffect(() => { + // Convert object to array for display + const aliasArray = Object.entries(initialModelGroupAlias).map(([aliasName, targetModelGroup], index) => ({ + id: `${index}-${aliasName}`, + aliasName, + targetModelGroup, + })); + setAliases(aliasArray); + }, [initialModelGroupAlias]); + + const saveAliasesToBackend = async (updatedAliases: AliasItem[]) => { + if (!accessToken) { + console.error("Access token is missing"); + return false; + } + + try { + // Convert array back to object format + const aliasObject: { [key: string]: string } = {}; + updatedAliases.forEach(alias => { + aliasObject[alias.aliasName] = alias.targetModelGroup; + }); + + const payload = { + router_settings: { + model_group_alias: aliasObject, + }, + }; + + console.log("Saving model group alias:", aliasObject); + await setCallbacksCall(accessToken, payload); + + if (onAliasUpdate) { + onAliasUpdate(aliasObject); + } + + return true; + } catch (error) { + console.error("Failed to save model group alias settings:", error); + message.error("Failed to save model group alias settings"); + return false; + } + }; + + const handleAddAlias = async () => { + if (!newAlias.aliasName || !newAlias.targetModelGroup) { + message.error("Please provide both alias name and target model group"); + return; + } + + // Check for duplicate alias names + if (aliases.some(alias => alias.aliasName === newAlias.aliasName)) { + message.error("An alias with this name already exists"); + return; + } + + const newAliasObj: AliasItem = { + id: `${Date.now()}-${newAlias.aliasName}`, + aliasName: newAlias.aliasName, + targetModelGroup: newAlias.targetModelGroup, + }; + + const updatedAliases = [...aliases, newAliasObj]; + + if (await saveAliasesToBackend(updatedAliases)) { + setAliases(updatedAliases); + setNewAlias({ aliasName: "", targetModelGroup: "" }); + message.success("Alias added successfully"); + } + }; + + const handleEditAlias = (alias: AliasItem) => { + setEditingAlias({ ...alias }); + }; + + const handleUpdateAlias = async () => { + if (!editingAlias) return; + + if (!editingAlias.aliasName || !editingAlias.targetModelGroup) { + message.error("Please provide both alias name and target model group"); + return; + } + + // Check for duplicate alias names (excluding current alias) + if (aliases.some(alias => alias.id !== editingAlias.id && alias.aliasName === editingAlias.aliasName)) { + message.error("An alias with this name already exists"); + return; + } + + const updatedAliases = aliases.map(alias => + alias.id === editingAlias.id ? editingAlias : alias + ); + + if (await saveAliasesToBackend(updatedAliases)) { + setAliases(updatedAliases); + setEditingAlias(null); + message.success("Alias updated successfully"); + } + }; + + const handleCancelEdit = () => { + setEditingAlias(null); + }; + + const deleteAlias = async (aliasId: string) => { + const updatedAliases = aliases.filter(alias => alias.id !== aliasId); + + if (await saveAliasesToBackend(updatedAliases)) { + setAliases(updatedAliases); + message.success("Alias deleted successfully"); + } + }; + + // Convert current aliases to object for config example + const aliasObject = aliases.reduce((acc, alias) => { + acc[alias.aliasName] = alias.targetModelGroup; + return acc; + }, {} as { [key: string]: string }); + + return ( + +
setIsExpanded(!isExpanded)} + > +
+ Model Group Alias Settings +

Create aliases for your model groups to simplify API calls. For example, you can create an alias 'gpt-4o' that points to 'gpt-4o-mini-openai' model group.

+
+
+ {isExpanded ? ( + + ) : ( + + )} +
+
+ + {isExpanded && ( +
+
+ Add New Alias +
+
+ + + setNewAlias({ + ...newAlias, + aliasName: e.target.value, + }) + } + placeholder="e.g., gpt-4o" + className="w-full px-3 py-2 border border-gray-300 rounded-md text-sm" + /> +
+
+ + + setNewAlias({ + ...newAlias, + targetModelGroup: e.target.value, + }) + } + placeholder="e.g., gpt-4o-mini-openai" + className="w-full px-3 py-2 border border-gray-300 rounded-md text-sm" + /> +
+
+ +
+
+
+ + + Manage Existing Aliases + +
+
+ + + + + Alias Name + + + Target Model Group + + + Actions + + + + + {aliases.map((alias) => ( + + {editingAlias && editingAlias.id === alias.id ? ( + <> + + + setEditingAlias({ + ...editingAlias, + aliasName: e.target.value, + }) + } + className="w-full px-2 py-1 border border-gray-300 rounded-md text-sm" + /> + + + + setEditingAlias({ + ...editingAlias, + targetModelGroup: e.target.value, + }) + } + className="w-full px-2 py-1 border border-gray-300 rounded-md text-sm" + /> + + +
+ + +
+
+ + ) : ( + <> + + {alias.aliasName} + + + {alias.targetModelGroup} + + +
+ + +
+
+ + )} +
+ ))} + {aliases.length === 0 && ( + + + No aliases added yet. Add a new alias above. + + + )} +
+
+
+
+ + {/* Configuration Example */} + + Configuration Example + + Here's how your current aliases would look in the config.yaml: + +
+
+ router_settings: +
+   model_group_alias: + {Object.keys(aliasObject).length === 0 ? ( + +
+     # No aliases configured yet +
+ ) : ( + Object.entries(aliasObject).map(([key, value]) => ( + +
+     "{key}": "{value}" +
+ )) + )} +
+
+
+
+ )} +
+ ); +}; + +export default ModelGroupAliasSettings; \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/new_usage.tsx b/ui/litellm-dashboard/src/components/new_usage.tsx index b2ea9150db..3dda412e20 100644 --- a/ui/litellm-dashboard/src/components/new_usage.tsx +++ b/ui/litellm-dashboard/src/components/new_usage.tsx @@ -61,7 +61,7 @@ import { EntityList } from "./entity_usage" import { formatNumberWithCommas } from "@/utils/dataUtils" import { valueFormatterSpend } from "./usage/utils/value_formatters" import CloudZeroExportModal from "./cloudzero_export_modal" -import { UiLoadingSpinner } from "./ui/ui-loading-spinner" +import { ChartLoader } from "./shared/chart_loader" interface NewUsagePageProps { accessToken: string | null @@ -368,23 +368,6 @@ const NewUsagePage: React.FC = ({ accessToken, userRole, user return () => clearTimeout(timeoutId) }, [fetchUserSpendData]) - // Enhanced loading component with better visual feedback - const ChartLoader = () => ( -
-
- -
- - {isDateChanging ? "Processing date selection..." : "Loading chart data..."} - - - {isDateChanging ? "This will only take a moment" : "Fetching your data"} - -
-
-
- ) - const modelMetrics = processActivityData(userSpendData, "models") const keyMetrics = processActivityData(userSpendData, "api_keys") const mcpServerMetrics = processActivityData(userSpendData, "mcp_servers") @@ -530,7 +513,7 @@ const NewUsagePage: React.FC = ({ accessToken, userRole, user Daily Spend {loading ? ( - + ) : ( = ({ accessToken, userRole, user
{loading ? ( - + ) : ( = ({ accessToken, userRole, user Spend by Provider {loading ? ( - + ) : ( diff --git a/ui/litellm-dashboard/src/components/shared/chart_loader.tsx b/ui/litellm-dashboard/src/components/shared/chart_loader.tsx new file mode 100644 index 0000000000..c6ab0c552b --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/chart_loader.tsx @@ -0,0 +1,24 @@ +import React from "react"; +import { UiLoadingSpinner } from "../ui/ui-loading-spinner"; + +interface ChartLoaderProps { + isDateChanging?: boolean; +} + +export const ChartLoader: React.FC = ({ isDateChanging = false }) => ( +
+
+ +
+ + {isDateChanging ? "Processing date selection..." : "Loading chart data..."} + + + {isDateChanging ? "This will only take a moment" : "Fetching your data"} + +
+
+
+); + +export default ChartLoader; \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/user_agent_activity.tsx b/ui/litellm-dashboard/src/components/user_agent_activity.tsx index ec2d1dfb15..0d0eda3615 100644 --- a/ui/litellm-dashboard/src/components/user_agent_activity.tsx +++ b/ui/litellm-dashboard/src/components/user_agent_activity.tsx @@ -26,9 +26,10 @@ import { } from "@tremor/react"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { userAgentAnalyticsCall, userAgentSummaryCall } from "./networking"; -import UsageDatePicker from "./shared/usage_date_picker"; +import AdvancedDatePicker from "./shared/advanced_date_picker"; import PerUserUsage from "./per_user_usage"; import { DateRangePickerValue } from "@tremor/react"; +import { ChartLoader } from "./shared/chart_loader"; interface UserAgentMetrics { dau: number; @@ -58,14 +59,14 @@ interface UserAgentAnalyticsResponse { } interface UserAgentSummaryData { - total_user_agents: number; + total_tags: number; total_requests: number; total_successful_requests: number; total_failed_requests: number; total_tokens: number; total_spend: number; - top_user_agents: Array<{ - user_agent: string; + top_tags: Array<{ + tag: string; requests: number; successful_requests: number; failed_requests: number; @@ -92,13 +93,13 @@ const UserAgentActivity: React.FC = ({ }); const [summaryData, setSummaryData] = useState({ - total_user_agents: 0, + total_tags: 0, total_requests: 0, total_successful_requests: 0, total_failed_requests: 0, total_tokens: 0, total_spend: 0, - top_user_agents: [], + top_tags: [], }); const [dateValue, setDateValue] = useState({ @@ -107,50 +108,84 @@ const UserAgentActivity: React.FC = ({ }); const [userAgentFilter, setUserAgentFilter] = useState(""); - const [loading, setLoading] = useState(false); + const [analyticsLoading, setAnalyticsLoading] = useState(false); + const [summaryLoading, setSummaryLoading] = useState(false); + const [isDateChanging, setIsDateChanging] = useState(false); const [currentPage, setCurrentPage] = useState(1); - const fetchData = async () => { + const fetchAnalyticsData = async () => { if (!accessToken || !dateValue.from || !dateValue.to) return; - setLoading(true); + setAnalyticsLoading(true); try { - const [analytics, summary] = await Promise.all([ - userAgentAnalyticsCall( - accessToken, - dateValue.from, - dateValue.to, - currentPage, - 50, - userAgentFilter || undefined - ), - userAgentSummaryCall(accessToken, dateValue.from, dateValue.to), - ]); + const analytics = await userAgentAnalyticsCall( + accessToken, + dateValue.from, + dateValue.to, + currentPage, + 50, + userAgentFilter || undefined + ); setAnalyticsData(analytics); + } catch (error) { + console.error("Failed to fetch user agent analytics data:", error); + } finally { + setAnalyticsLoading(false); + setIsDateChanging(false); + } + }; + + const fetchSummaryData = async () => { + if (!accessToken || !dateValue.from || !dateValue.to) return; + + setSummaryLoading(true); + try { + const summary = await userAgentSummaryCall(accessToken, dateValue.from, dateValue.to); setSummaryData(summary); } catch (error) { - console.error("Failed to fetch user agent data:", error); + console.error("Failed to fetch user agent summary data:", error); } finally { - setLoading(false); + setSummaryLoading(false); + setIsDateChanging(false); } }; + // Super responsive date change handler + const handleDateChange = (newValue: DateRangePickerValue) => { + // Instant visual feedback + setIsDateChanging(true); + setAnalyticsLoading(true); + setSummaryLoading(true); + + // Update date immediately for UI responsiveness + setDateValue(newValue); + setCurrentPage(1); // Reset to first page when date changes + }; + + // Debounced effect for data fetching useEffect(() => { - fetchData(); - }, [accessToken, dateValue, userAgentFilter, currentPage]); + if (!dateValue.from || !dateValue.to) return; - const handleNextPage = () => { - if (currentPage < analyticsData.total_pages) { - setCurrentPage(currentPage + 1); - } - }; + const timeoutId = setTimeout(() => { + // Call both fetch functions independently + fetchAnalyticsData(); + fetchSummaryData(); + }, 50); // Very short debounce - const handlePrevPage = () => { - if (currentPage > 1) { - setCurrentPage(currentPage - 1); - } - }; + return () => clearTimeout(timeoutId); + }, [accessToken, dateValue, userAgentFilter]); + + // Separate effect for pagination that only affects analytics + useEffect(() => { + if (!dateValue.from || !dateValue.to) return; + + const timeoutId = setTimeout(() => { + fetchAnalyticsData(); + }, 50); + + return () => clearTimeout(timeoutId); + }, [currentPage]); // Aggregate data by user agent for charts const aggregatedByUserAgent = analyticsData.results.reduce((acc, item) => { @@ -184,10 +219,26 @@ const UserAgentActivity: React.FC = ({ (a: any, b: any) => b.total_requests - a.total_requests ); - const successRateData = summaryData.top_user_agents.map((ua) => ({ - user_agent: ua.user_agent, - success_rate: ua.successful_requests / (ua.requests || 1) * 100, - total_requests: ua.requests, + // Helper function to extract user agent from tag + const extractUserAgent = (tag: string): string => { + if (tag.startsWith("User-Agent: ")) { + return tag.replace("User-Agent: ", ""); + } + return tag; + }; + + // Helper function to truncate user agent name with tooltip + const truncateUserAgent = (userAgent: string): string => { + if (userAgent.length > 10) { + return userAgent.substring(0, 10) + "..."; + } + return userAgent; + }; + + const successRateData = (summaryData.top_tags || []).map((tag) => ({ + user_agent: extractUserAgent(tag.tag), + success_rate: tag.successful_requests / (tag.requests || 1) * 100, + total_requests: tag.requests, })); // Get unique user agents for chart @@ -263,12 +314,9 @@ const UserAgentActivity: React.FC = ({ {/* Date Range Picker */} - { - setDateValue(value); - setCurrentPage(1); // Reset to first page when date changes - }} + onValueChange={handleDateChange} /> @@ -289,49 +337,59 @@ const UserAgentActivity: React.FC = ({ {/* Top 4 User Agents Cards */} - - {summaryData.top_user_agents.slice(0, 4).map((ua, index) => ( - - - {ua.user_agent} - -
-
- Success Requests - {formatAbbreviatedNumber(ua.successful_requests)} -
-
- Total Tokens - {formatAbbreviatedNumber(ua.tokens)} -
-
- Total Cost - ${formatAbbreviatedNumber(ua.spend, 4)} -
-
+ {summaryLoading ? ( + + - ))} - {/* Fill remaining slots if less than 4 agents */} - {Array.from({ length: Math.max(0, 4 - summaryData.top_user_agents.length) }).map((_, index) => ( - - No Data -
-
- Success Requests - - + ) : ( + + {(summaryData.top_tags || []).slice(0, 4).map((tag, index) => { + const userAgent = extractUserAgent(tag.tag); + const displayName = truncateUserAgent(userAgent); + return ( + + + {displayName} + +
+
+ Success Requests + {formatAbbreviatedNumber(tag.successful_requests)} +
+
+ Total Tokens + {formatAbbreviatedNumber(tag.tokens)} +
+
+ Total Cost + ${formatAbbreviatedNumber(tag.spend, 4)} +
+
+
+ ); + })} + {/* Fill remaining slots if less than 4 agents */} + {Array.from({ length: Math.max(0, 4 - (summaryData.top_tags || []).length) }).map((_, index) => ( + + No Data +
+
+ Success Requests + - +
+
+ Total Tokens + - +
+
+ Total Cost + - +
-
- Total Tokens - - -
-
- Total Cost - - -
-
- - ))} - + + ))} + + )} {/* Main TabGroup for DAU/WAU/MAU vs Per User Usage */} @@ -361,45 +419,57 @@ const UserAgentActivity: React.FC = ({
Daily Active Users - Last 7 Days
- formatAbbreviatedNumber(value)} - yAxisWidth={60} - showLegend={true} - /> + {analyticsLoading ? ( + + ) : ( + formatAbbreviatedNumber(value)} + yAxisWidth={60} + showLegend={true} + /> + )}
Weekly Active Users - Last 4 Weeks
- formatAbbreviatedNumber(value)} - yAxisWidth={60} - showLegend={true} - /> + {analyticsLoading ? ( + + ) : ( + formatAbbreviatedNumber(value)} + yAxisWidth={60} + showLegend={true} + /> + )}
Monthly Active Users - Last 7 Months
- formatAbbreviatedNumber(value)} - yAxisWidth={60} - showLegend={true} - /> + {analyticsLoading ? ( + + ) : ( + formatAbbreviatedNumber(value)} + yAxisWidth={60} + showLegend={true} + /> + )}