mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-09 16:24:38 +00:00
[Feat] Add new AWS SQS Logging Integration (#12176)
* add aws_sqs * add sqs controls * add SQS to registry * fix url lib parse * fixes AWS SQS * test_async_sqs_logger_flush * fix test * fix SQS logger auth * add AWS SQS * add aws sqs * docs logging * test_async_sqs_logger_flush * test_async_sqs_logger_flush * add SQS logger * update SQS logging * use constants for SQS
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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/<variable name> 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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)}")
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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=<url_encoded_json>"
|
||||
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"
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user