diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 617b08ae0f..482b599d63 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -415,6 +415,8 @@ router_settings: | DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND | Default price per second for Replicate GPU. Default is 0.001400 | DEFAULT_REPLICATE_POLLING_DELAY_SECONDS | Default delay in seconds for Replicate polling. Default is 1 | DEFAULT_REPLICATE_POLLING_RETRIES | Default number of retries for Replicate polling. Default is 5 +| DEFAULT_SQS_BATCH_SIZE | Default batch size for SQS logging. Default is 512 +| DEFAULT_SQS_FLUSH_INTERVAL_SECONDS | Default flush interval for SQS logging. Default is 10 | DEFAULT_S3_BATCH_SIZE | Default batch size for S3 logging. Default is 512 | DEFAULT_S3_FLUSH_INTERVAL_SECONDS | Default flush interval for S3 logging. Default is 10 | DEFAULT_SLACK_ALERTING_THRESHOLD | Default threshold for Slack alerting. Default is 300 diff --git a/docs/my-website/docs/proxy/logging.md b/docs/my-website/docs/proxy/logging.md index 7ec9080dfd..e1926d776e 100644 --- a/docs/my-website/docs/proxy/logging.md +++ b/docs/my-website/docs/proxy/logging.md @@ -9,6 +9,7 @@ Log Proxy input, output, and exceptions using: - Langfuse - OpenTelemetry - GCS, s3, Azure (Blob) Buckets +- AWS SQS - Lunary - MLflow - Deepeval @@ -1384,6 +1385,75 @@ litellm_settings: On s3 bucket, you will see the object key as `my-test-path/my-team-alias/...` +## AWS SQS + + +| Property | Details | +|----------|---------| +| Description | Log LLM Input/Output to AWS SQS Queue | +| AWS Docs on SQS | [AWS SQS](https://aws.amazon.com/sqs/) | +| Fields Logged to SQS | LiteLLM [Standard Logging Payload is logged for each LLM call](../proxy/logging_spec) | + + +Log LLM Logs to [AWS Simple Queue Service (SQS)](https://aws.amazon.com/sqs/) + +We will use the litellm `--config` to set + +- `litellm.callbacks = ["aws_sqs"]` + +This will log all successful LLM calls to AWS SQS Queue + +**Step 1** Set AWS Credentials in .env + +```shell +AWS_ACCESS_KEY_ID = "" +AWS_SECRET_ACCESS_KEY = "" +AWS_REGION_NAME = "" +``` + +**Step 2**: Create a `config.yaml` file and set `litellm_settings`: `callbacks` + +```yaml +model_list: + - model_name: gpt-4o + litellm_params: + model: gpt-4o +litellm_settings: + callbacks: ["aws_sqs"] + aws_sqs_callback_params: + sqs_queue_url: https://sqs.us-west-2.amazonaws.com/123456789012/my-queue # AWS SQS Queue URL + sqs_region_name: us-west-2 # AWS Region Name for SQS + sqs_aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID # use os.environ/ to pass environment variables. This is AWS Access Key ID for SQS + sqs_aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY # AWS Secret Access Key for SQS + sqs_batch_size: 10 # [OPTIONAL] Number of messages to batch before sending (default: 10) + sqs_flush_interval: 30 # [OPTIONAL] Time in seconds to wait before flushing batch (default: 30) +``` + +**Step 3**: Start the proxy, make a test request + +Start proxy + +```shell +litellm --config config.yaml --debug +``` + +Test Request + +```shell +curl --location 'http://0.0.0.0:4000/chat/completions' \ + --header 'Content-Type: application/json' \ + --data ' { + "model": "gpt-4o", + "messages": [ + { + "role": "user", + "content": "what llm are you" + } + ] + }' +``` + + ## Azure Blob Storage Log LLM Logs to [Azure Data Lake Storage](https://learn.microsoft.com/en-us/azure/storage/blobs/data-lake-storage-introduction) diff --git a/litellm/__init__.py b/litellm/__init__.py index 0ec8a1e01b..a772aa6a05 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -118,6 +118,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "smtp_email", "deepeval", "s3_v2", + "aws_sqs", ] logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None _known_custom_logger_compatible_callbacks: List = list( @@ -291,6 +292,7 @@ model_cost_map_url: str = ( suppress_debug_info = False dynamodb_table_name: Optional[str] = None s3_callback_params: Optional[Dict] = None +aws_sqs_callback_params: Optional[Dict] = None generic_logger_headers: Optional[Dict] = None default_key_generate_params: Optional[Dict] = None upperbound_key_generate_params: Optional[LiteLLM_UpperboundKeyGenerateParams] = None diff --git a/litellm/constants.py b/litellm/constants.py index 74ba5d7094..b34e7a5569 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -8,6 +8,12 @@ DEFAULT_S3_FLUSH_INTERVAL_SECONDS = int( os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10) ) DEFAULT_S3_BATCH_SIZE = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512)) +DEFAULT_SQS_FLUSH_INTERVAL_SECONDS = int( + os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10) +) +DEFAULT_SQS_BATCH_SIZE = int(os.getenv("DEFAULT_SQS_BATCH_SIZE", 512)) +SQS_SEND_MESSAGE_ACTION = "SendMessage" +SQS_API_VERSION = "2012-11-05" DEFAULT_MAX_RETRIES = int(os.getenv("DEFAULT_MAX_RETRIES", 2)) DEFAULT_MAX_RECURSE_DEPTH = int(os.getenv("DEFAULT_MAX_RECURSE_DEPTH", 100)) DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER = int( diff --git a/litellm/integrations/sqs.py b/litellm/integrations/sqs.py new file mode 100644 index 0000000000..2a0c73dfdb --- /dev/null +++ b/litellm/integrations/sqs.py @@ -0,0 +1,275 @@ +"""SQS Logging Integration + +This logger sends ``StandardLoggingPayload`` entries to an AWS SQS queue. + +""" + +from __future__ import annotations + +import asyncio +from typing import List, Optional + +import litellm +from litellm._logging import print_verbose, verbose_logger +from litellm.constants import ( + DEFAULT_SQS_BATCH_SIZE, + DEFAULT_SQS_FLUSH_INTERVAL_SECONDS, + SQS_API_VERSION, + SQS_SEND_MESSAGE_ACTION, +) +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.types.utils import StandardLoggingPayload + +from .custom_batch_logger import CustomBatchLogger + + +class SQSLogger(CustomBatchLogger, BaseAWSLLM): + """Batching logger that writes logs to an AWS SQS queue.""" + + def __init__( + self, + sqs_queue_url: Optional[str] = None, + sqs_region_name: Optional[str] = None, + sqs_api_version: Optional[str] = None, + sqs_use_ssl: bool = True, + sqs_verify: Optional[bool] = None, + sqs_endpoint_url: Optional[str] = None, + sqs_aws_access_key_id: Optional[str] = None, + sqs_aws_secret_access_key: Optional[str] = None, + sqs_aws_session_token: Optional[str] = None, + sqs_aws_session_name: Optional[str] = None, + sqs_aws_profile_name: Optional[str] = None, + sqs_aws_role_name: Optional[str] = None, + sqs_aws_web_identity_token: Optional[str] = None, + sqs_aws_sts_endpoint: Optional[str] = None, + sqs_flush_interval: Optional[int] = DEFAULT_SQS_FLUSH_INTERVAL_SECONDS, + sqs_batch_size: Optional[int] = DEFAULT_SQS_BATCH_SIZE, + sqs_config=None, + **kwargs, + ) -> None: + try: + verbose_logger.debug( + f"in init sqs logger - sqs_callback_params {litellm.aws_sqs_callback_params}" + ) + + self.async_httpx_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.LoggingCallback, + ) + + self._init_sqs_params( + sqs_queue_url=sqs_queue_url, + sqs_region_name=sqs_region_name, + sqs_api_version=sqs_api_version, + sqs_use_ssl=sqs_use_ssl, + sqs_verify=sqs_verify, + sqs_endpoint_url=sqs_endpoint_url, + sqs_aws_access_key_id=sqs_aws_access_key_id, + sqs_aws_secret_access_key=sqs_aws_secret_access_key, + sqs_aws_session_token=sqs_aws_session_token, + sqs_aws_session_name=sqs_aws_session_name, + sqs_aws_profile_name=sqs_aws_profile_name, + sqs_aws_role_name=sqs_aws_role_name, + sqs_aws_web_identity_token=sqs_aws_web_identity_token, + sqs_aws_sts_endpoint=sqs_aws_sts_endpoint, + sqs_config=sqs_config, + ) + + asyncio.create_task(self.periodic_flush()) + self.flush_lock = asyncio.Lock() + + verbose_logger.debug( + f"sqs flush interval: {sqs_flush_interval}, sqs batch size: {sqs_batch_size}" + ) + + CustomBatchLogger.__init__( + self, + flush_lock=self.flush_lock, + flush_interval=sqs_flush_interval, + batch_size=sqs_batch_size, + ) + + self.log_queue: List[StandardLoggingPayload] = [] + + BaseAWSLLM.__init__(self) + + except Exception as e: + print_verbose(f"Got exception on init sqs client {str(e)}") + raise e + + def _init_sqs_params( + self, + sqs_queue_url: Optional[str] = None, + sqs_region_name: Optional[str] = None, + sqs_api_version: Optional[str] = None, + sqs_use_ssl: bool = True, + sqs_verify: Optional[bool] = None, + sqs_endpoint_url: Optional[str] = None, + sqs_aws_access_key_id: Optional[str] = None, + sqs_aws_secret_access_key: Optional[str] = None, + sqs_aws_session_token: Optional[str] = None, + sqs_aws_session_name: Optional[str] = None, + sqs_aws_profile_name: Optional[str] = None, + sqs_aws_role_name: Optional[str] = None, + sqs_aws_web_identity_token: Optional[str] = None, + sqs_aws_sts_endpoint: Optional[str] = None, + sqs_config=None, + ) -> None: + litellm.aws_sqs_callback_params = litellm.aws_sqs_callback_params or {} + + # read in .env variables - example os.environ/AWS_BUCKET_NAME + for key, value in litellm.aws_sqs_callback_params.items(): + if isinstance(value, str) and value.startswith("os.environ/"): + litellm.aws_sqs_callback_params[key] = litellm.get_secret(value) + + self.sqs_queue_url = ( + litellm.aws_sqs_callback_params.get("sqs_queue_url") or sqs_queue_url + ) + self.sqs_region_name = ( + litellm.aws_sqs_callback_params.get("sqs_region_name") or sqs_region_name + ) + self.sqs_api_version = ( + litellm.aws_sqs_callback_params.get("sqs_api_version") or sqs_api_version + ) + self.sqs_use_ssl = ( + litellm.aws_sqs_callback_params.get("sqs_use_ssl", True) or sqs_use_ssl + ) + self.sqs_verify = litellm.aws_sqs_callback_params.get("sqs_verify") or sqs_verify + self.sqs_endpoint_url = ( + litellm.aws_sqs_callback_params.get("sqs_endpoint_url") or sqs_endpoint_url + ) + self.sqs_aws_access_key_id = ( + litellm.aws_sqs_callback_params.get("sqs_aws_access_key_id") + or sqs_aws_access_key_id + ) + + self.sqs_aws_secret_access_key = ( + litellm.aws_sqs_callback_params.get("sqs_aws_secret_access_key") + or sqs_aws_secret_access_key + ) + + self.sqs_aws_session_token = ( + litellm.aws_sqs_callback_params.get("sqs_aws_session_token") + or sqs_aws_session_token + ) + + self.sqs_aws_session_name = ( + litellm.aws_sqs_callback_params.get("sqs_aws_session_name") or sqs_aws_session_name + ) + + self.sqs_aws_profile_name = ( + litellm.aws_sqs_callback_params.get("sqs_aws_profile_name") or sqs_aws_profile_name + ) + + self.sqs_aws_role_name = ( + litellm.aws_sqs_callback_params.get("sqs_aws_role_name") or sqs_aws_role_name + ) + + self.sqs_aws_web_identity_token = ( + litellm.aws_sqs_callback_params.get("sqs_aws_web_identity_token") + or sqs_aws_web_identity_token + ) + + self.sqs_aws_sts_endpoint = ( + litellm.aws_sqs_callback_params.get("sqs_aws_sts_endpoint") or sqs_aws_sts_endpoint + ) + + self.sqs_config = litellm.aws_sqs_callback_params.get("sqs_config") or sqs_config + + async def async_log_success_event( + self, kwargs, response_obj, start_time, end_time + ) -> None: + try: + verbose_logger.debug( + "SQS Logging - Enters logging function for model %s", kwargs + ) + standard_logging_payload = kwargs.get("standard_logging_object") + if standard_logging_payload is None: + raise ValueError("standard_logging_payload is None") + + self.log_queue.append(standard_logging_payload) + verbose_logger.debug( + "sqs logging: queue length %s, batch size %s", + len(self.log_queue), + self.batch_size, + ) + except Exception as e: + verbose_logger.exception(f"sqs Layer Error - {str(e)}") + + async def async_send_batch(self) -> None: + verbose_logger.debug( + f"sqs logger - sending batch of {len(self.log_queue)}" + ) + if not self.log_queue: + return + + for payload in self.log_queue: + asyncio.create_task(self.async_send_message(payload)) + + async def async_send_message(self, payload: StandardLoggingPayload) -> None: + try: + from urllib.parse import quote + + import requests + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest + + from litellm.litellm_core_utils.asyncify import asyncify + + asyncified_get_credentials = asyncify(self.get_credentials) + credentials = await asyncified_get_credentials( + aws_access_key_id=self.sqs_aws_access_key_id, + aws_secret_access_key=self.sqs_aws_secret_access_key, + aws_session_token=self.sqs_aws_session_token, + aws_region_name=self.sqs_region_name, + aws_session_name=self.sqs_aws_session_name, + aws_profile_name=self.sqs_aws_profile_name, + aws_role_name=self.sqs_aws_role_name, + aws_web_identity_token=self.sqs_aws_web_identity_token, + aws_sts_endpoint=self.sqs_aws_sts_endpoint, + ) + + if self.sqs_queue_url is None: + raise ValueError("sqs_queue_url not set") + + json_string = safe_dumps(payload) + + body = ( + f"Action={SQS_SEND_MESSAGE_ACTION}&Version={SQS_API_VERSION}&MessageBody=" + + quote(json_string, safe="") + ) + + headers = { + "Content-Type": "application/x-www-form-urlencoded", + } + + req = requests.Request( + "POST", self.sqs_queue_url, data=body, headers=headers + ) + prepped = req.prepare() + + aws_request = AWSRequest( + method=prepped.method, + url=prepped.url, + data=prepped.body, + headers=prepped.headers, + ) + SigV4Auth(credentials, "sqs", self.sqs_region_name).add_auth( + aws_request + ) + + signed_headers = dict(aws_request.headers.items()) + + response = await self.async_httpx_client.post( + self.sqs_queue_url, + data=body, + headers=signed_headers, + ) + response.raise_for_status() + except Exception as e: + verbose_logger.exception(f"Error sending to SQS: {str(e)}") + diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index 1b75cc3e3d..e6180fa8d9 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -32,6 +32,7 @@ from litellm.integrations.opentelemetry import OpenTelemetry from litellm.integrations.opik.opik import OpikLogger from litellm.integrations.prometheus import PrometheusLogger from litellm.integrations.s3_v2 import S3Logger +from litellm.integrations.sqs import SQSLogger from litellm.integrations.vector_store_integrations.bedrock_vector_store import ( BedrockVectorStore, ) @@ -73,6 +74,7 @@ class CustomLoggerRegistry: "bedrock_vector_store": BedrockVectorStore, "deepeval": DeepEvalLogger, "s3_v2": S3Logger, + "aws_sqs": SQSLogger, "dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler, } diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index f7aa59db97..88c39a665a 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -56,6 +56,7 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.deepeval.deepeval import DeepEvalLogger from litellm.integrations.mlflow import MlflowLogger +from litellm.integrations.sqs import SQSLogger from litellm.integrations.vector_store_integrations.bedrock_vector_store import ( BedrockVectorStore, ) @@ -3142,6 +3143,14 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 _s3_v2_logger = S3V2Logger() _in_memory_loggers.append(_s3_v2_logger) return _s3_v2_logger # type: ignore + elif logging_integration == "aws_sqs": + for callback in _in_memory_loggers: + if isinstance(callback, SQSLogger): + return callback # type: ignore + + _aws_sqs_logger = SQSLogger() + _in_memory_loggers.append(_aws_sqs_logger) + return _aws_sqs_logger # type: ignore elif logging_integration == "azure_storage": for callback in _in_memory_loggers: if isinstance(callback, AzureBlobStorageLogger): @@ -3476,6 +3485,13 @@ def get_custom_logger_compatible_class( # noqa: PLR0915 for callback in _in_memory_loggers: if isinstance(callback, S3V2Logger): return callback + elif logging_integration == "aws_sqs": + for callback in _in_memory_loggers: + if isinstance(callback, SQSLogger): + return callback + _aws_sqs_logger = SQSLogger() + _in_memory_loggers.append(_aws_sqs_logger) + return _aws_sqs_logger # type: ignore elif logging_integration == "azure_storage": for callback in _in_memory_loggers: if isinstance(callback, AzureBlobStorageLogger): diff --git a/tests/logging_callback_tests/test_sqs_logger.py b/tests/logging_callback_tests/test_sqs_logger.py new file mode 100644 index 0000000000..a3fe09b2a5 --- /dev/null +++ b/tests/logging_callback_tests/test_sqs_logger.py @@ -0,0 +1,73 @@ +import asyncio +import json +from unittest.mock import AsyncMock, MagicMock +from urllib.parse import unquote + +import litellm +import pytest + +from litellm.integrations.sqs import SQSLogger +from litellm.types.utils import StandardLoggingPayload + + +@pytest.mark.asyncio +async def test_async_sqs_logger_flush(): + expected_queue_url = "https://sqs.us-east-1.amazonaws.com/123456789012/test-queue" + expected_region = "us-east-1" + + sqs_logger = SQSLogger( + sqs_queue_url=expected_queue_url, + sqs_region_name=expected_region, + sqs_flush_interval=1, + ) + + # Mock the httpx client + mock_response = MagicMock() + mock_response.raise_for_status = MagicMock() + sqs_logger.async_httpx_client.post = AsyncMock(return_value=mock_response) + + litellm.callbacks = [sqs_logger] + + await litellm.acompletion( + model="gpt-4o", + messages=[{"role": "user", "content": "hello"}], + mock_response="hi", + ) + + await asyncio.sleep(2) + + # Verify that httpx post was called + sqs_logger.async_httpx_client.post.assert_called() + + # Get the call arguments + call_args = sqs_logger.async_httpx_client.post.call_args + + # Verify the URL is correct + called_url = call_args[0][0] # First positional argument + assert called_url == expected_queue_url, f"Expected URL {expected_queue_url}, got {called_url}" + + # Verify the payload contains StandardLoggingPayload data + called_data = call_args.kwargs['data'] + + # Extract the MessageBody from the URL-encoded data + # Format: "Action=SendMessage&Version=2012-11-05&MessageBody=" + assert "Action=SendMessage" in called_data + assert "Version=2012-11-05" in called_data + assert "MessageBody=" in called_data + + # Extract and decode the message body + message_body_start = called_data.find("MessageBody=") + len("MessageBody=") + message_body_encoded = called_data[message_body_start:] + message_body_json = unquote(message_body_encoded) + + # Parse the JSON to verify it's a StandardLoggingPayload + payload_data = json.loads(message_body_json) + + # Verify it has the expected StandardLoggingPayload structure + assert "model" in payload_data + assert "messages" in payload_data + assert "response" in payload_data + assert payload_data["model"] == "gpt-4o" + assert len(payload_data["messages"]) == 1 + assert payload_data["messages"][0]["role"] == "user" + assert payload_data["messages"][0]["content"] == "hello" diff --git a/tests/logging_callback_tests/test_unit_tests_init_callbacks.py b/tests/logging_callback_tests/test_unit_tests_init_callbacks.py index aff622201f..3c9d31890c 100644 --- a/tests/logging_callback_tests/test_unit_tests_init_callbacks.py +++ b/tests/logging_callback_tests/test_unit_tests_init_callbacks.py @@ -28,7 +28,6 @@ for collector in collectors: ###################################### - expected_env_vars = { "LAGO_API_KEY": "api_key", "LAGO_API_BASE": "mock_base", @@ -56,6 +55,7 @@ expected_env_vars = { "AWS_SECRET_ACCESS_KEY": "aws_secret_access_key", "AWS_ACCESS_KEY_ID": "aws_access_key_id", "AWS_REGION": "aws_region", + "AWS_SQS_QUEUE_URL": "https://sqs.us-east-1.amazonaws.com/123456789012/test-queue", }