diff --git a/.env.example b/.env.example index 7ae9837b3..5c85a0650 100644 --- a/.env.example +++ b/.env.example @@ -13,3 +13,7 @@ export R2R_POSTGRES_HOST=your_host export R2R_POSTGRES_PORT=your_port export R2R_POSTGRES_DBNAME=your_db export R2R_PROJECT_NAME=your_project_name + +# Optional Interloom billing usage delivery +export R2R_BILLING_CALLBACK_URL=https://interloom.example.com/api/v1/webhooks/r2r-usage +export R2R_BILLING_CALLBACK_SERVICE_TOKEN=your_service_token diff --git a/py/core/billing.py b/py/core/billing.py new file mode 100644 index 000000000..46e438efe --- /dev/null +++ b/py/core/billing.py @@ -0,0 +1,239 @@ +import asyncio +import contextvars +import functools +import logging +import os +from datetime import datetime, timezone +from typing import Any, Awaitable, Callable +from uuid import UUID, uuid4 + +import httpx +from pydantic import BaseModel, model_validator + +logger = logging.getLogger(__name__) + + +class InterloomBillingContext(BaseModel): + ingestion_id: UUID + organization_id: UUID + file_id: UUID | None = None + note_id: UUID | None = None + + @model_validator(mode="after") + def validate_source(self): + if (self.file_id is None) == (self.note_id is None): + raise ValueError("Exactly one of file_id or note_id is required") + return self + + @property + def source_id(self) -> UUID: + if self.file_id is not None: + return self.file_id + if self.note_id is None: + raise ValueError("Billing context source is missing") + return self.note_id + + +_billing_context: contextvars.ContextVar[InterloomBillingContext | None] = ( + contextvars.ContextVar("interloom_billing_context", default=None) +) + + +def _context_from_input( + input_data: dict[str, Any], +) -> InterloomBillingContext | None: + value = input_data.get("interloom_billing_context") + if value is None: + return None + return InterloomBillingContext.model_validate(value) + + +def with_billing_context_from_input(function: Callable[..., Awaitable[Any]]): + @functools.wraps(function) + async def wrapped(input_data: dict[str, Any], *args: Any, **kwargs: Any): + token = _billing_context.set(_context_from_input(input_data)) + try: + return await function(input_data, *args, **kwargs) + finally: + _billing_context.reset(token) + + return wrapped + + +def with_billing_context_from_hatchet(function: Callable[..., Awaitable[Any]]): + @functools.wraps(function) + async def wrapped(instance: Any, context: Any, *args: Any, **kwargs: Any): + input_data = context.workflow_input()["request"] + token = _billing_context.set(_context_from_input(input_data)) + try: + return await function(instance, context, *args, **kwargs) + finally: + _billing_context.reset(token) + + return wrapped + + +def _field(value: Any, name: str, default: Any = None) -> Any: + if value is None: + return default + if isinstance(value, dict): + return value.get(name, default) + return getattr(value, name, default) + + +def _usage_from_response(response: Any) -> dict[str, int | None]: + usage = _field(response, "usage") + input_tokens = _field(usage, "prompt_tokens") + if input_tokens is None: + input_tokens = _field(usage, "input_tokens") + output_tokens = _field(usage, "completion_tokens") + if output_tokens is None: + output_tokens = _field(usage, "output_tokens") + + input_details = _field(usage, "prompt_tokens_details") + if input_details is None: + input_details = _field(usage, "input_tokens_details") + + return { + "input_tokens": input_tokens, + "output_tokens": output_tokens, + "cached_input_tokens": _field(input_details, "cached_tokens"), + } + + +class BillingUsageRecorder: + def __init__(self, outbox_handler: Any): + self.outbox_handler = outbox_handler + + async def start_call(self, *, operation: str, model: str) -> UUID | None: + context = _billing_context.get() + if context is None: + return None + + event_id = uuid4() + await self.outbox_handler.create_pending( + event_id=event_id, + context=context, + operation=operation, + requested_model=model, + started_at=datetime.now(timezone.utc), + ) + return event_id + + def add_litellm_metadata( + self, + *, + kwargs: dict[str, Any], + event_id: UUID | None, + ) -> dict[str, Any]: + if event_id is None: + return kwargs + + context = _billing_context.get() + if context is None: + return kwargs + + result = dict(kwargs) + metadata = dict(result.get("metadata") or {}) + spend_logs_metadata = dict(metadata.get("spend_logs_metadata") or {}) + spend_logs_metadata.update( + { + "source": "r2r", + "billing_event_id": str(event_id), + "ingestion_id": str(context.ingestion_id), + "organization_id": str(context.organization_id), + "file_id": str(context.file_id) if context.file_id else None, + "note_id": str(context.note_id) if context.note_id else None, + } + ) + metadata["spend_logs_metadata"] = spend_logs_metadata + result["metadata"] = metadata + return result + + async def complete_call( + self, event_id: UUID | None, response: Any + ) -> None: + if event_id is None: + return + await self.outbox_handler.complete( + event_id=event_id, + provider_request_id=_field(response, "id"), + resolved_model=_field(response, "model"), + ended_at=datetime.now(timezone.utc), + **_usage_from_response(response), + ) + + async def fail_call(self, event_id: UUID | None, error: Exception) -> None: + if event_id is None: + return + await self.outbox_handler.fail( + event_id=event_id, + ended_at=datetime.now(timezone.utc), + error=str(error), + ) + + +class BillingOutboxDispatcher: + def __init__( + self, + *, + outbox_handler: Any, + callback_url: str, + service_token: str, + poll_interval_seconds: float = 5, + ): + self.outbox_handler = outbox_handler + self.callback_url = callback_url + self.service_token = service_token + self.poll_interval_seconds = poll_interval_seconds + self._stopped = asyncio.Event() + + @classmethod + def from_environment(cls, outbox_handler: Any): + callback_url = os.getenv("R2R_BILLING_CALLBACK_URL") + service_token = os.getenv("R2R_BILLING_CALLBACK_SERVICE_TOKEN") + if not callback_url or not service_token: + return None + return cls( + outbox_handler=outbox_handler, + callback_url=callback_url, + service_token=service_token, + ) + + async def run(self) -> None: + async with httpx.AsyncClient(timeout=30) as client: + while not self._stopped.is_set(): + try: + await self._deliver_batch(client) + except Exception: + logger.exception("Failed to dispatch R2R billing usage") + + try: + await asyncio.wait_for( + self._stopped.wait(), + timeout=self.poll_interval_seconds, + ) + except TimeoutError: + pass + + async def stop(self) -> None: + self._stopped.set() + + async def _deliver_batch(self, client: httpx.AsyncClient) -> None: + events = await self.outbox_handler.claim_for_delivery(limit=100) + for event in events: + event_id = event["billing_event_id"] + try: + response = await client.post( + self.callback_url, + json=event, + headers={"il-service-token": self.service_token}, + ) + response.raise_for_status() + except Exception as error: + await self.outbox_handler.mark_delivery_failed( + event_id=event_id, + error=str(error), + ) + else: + await self.outbox_handler.mark_delivered(event_id=event_id) diff --git a/py/core/main/api/v3/documents_router.py b/py/core/main/api/v3/documents_router.py index 68c7b81fa..d7ea682c2 100644 --- a/py/core/main/api/v3/documents_router.py +++ b/py/core/main/api/v3/documents_router.py @@ -38,6 +38,7 @@ WrappedIngestionResponse, WrappedRelationshipsResponse, ) +from core.billing import InterloomBillingContext from core.utils import update_settings_from_dict from shared.abstractions import IngestionMode @@ -178,6 +179,20 @@ def _prepare_ingestion_config( effective_config.validate_config() return effective_config + @staticmethod + def _prepare_billing_context( + billing_context: InterloomBillingContext | None, + document_id: UUID, + ) -> dict[str, Any] | None: + if billing_context is None: + return None + if billing_context.source_id != document_id: + raise R2RException( + status_code=422, + message="Billing context source ID must match the document ID.", + ) + return billing_context.model_dump(mode="json") + def _setup_routes(self): @self.router.post( "/documents", @@ -258,6 +273,12 @@ async def create_document( None, description="Metadata to associate with the document, such as title, description, or custom fields.", ), + interloom_billing_context: Optional[ + Json[InterloomBillingContext] + ] = Form( + None, + description="Internal Interloom attribution for billing usage generated by this ingestion.", + ), ingestion_mode: IngestionMode = Form( default=IngestionMode.custom, description=( @@ -385,6 +406,9 @@ async def create_document( document_id = id or generate_document_id( "".join(chunks), auth_user.id ) + billing_context_payload = self._prepare_billing_context( + interloom_billing_context, document_id + ) # FIXME: Metadata doesn't seem to be getting passed through raw_chunks_for_doc = [ @@ -413,6 +437,7 @@ async def create_document( "ingestion_config": effective_ingestion_config.model_dump( mode="json" ), + "interloom_billing_context": billing_context_payload, } if run_with_orchestration: @@ -507,6 +532,10 @@ async def create_document( message="Either a file or content must be provided.", ) + billing_context_payload = self._prepare_billing_context( + interloom_billing_context, document_id + ) + workflow_input = { "file_data": file_data, "document_id": str(document_id), @@ -522,6 +551,7 @@ async def create_document( "user": auth_user.model_dump_json(), "size_in_bytes": content_length, "version": "v0", + "interloom_billing_context": billing_context_payload, } file_name = file_data["filename"] diff --git a/py/core/main/app_entry.py b/py/core/main/app_entry.py index cc216f7ad..ca9f72691 100644 --- a/py/core/main/app_entry.py +++ b/py/core/main/app_entry.py @@ -1,3 +1,4 @@ +import asyncio import logging import os from contextlib import asynccontextmanager @@ -9,6 +10,7 @@ from fastapi.responses import JSONResponse from core.base import R2RException +from core.billing import BillingOutboxDispatcher from core.utils.logging_config import configure_logging from .app import R2RApp @@ -42,7 +44,22 @@ async def lifespan(app: FastAPI): # Start the Hatchet worker await r2r_app.orchestration_provider.start_worker() - yield + billing_dispatcher = BillingOutboxDispatcher.from_environment( + r2r_app.providers.database.billing_outbox_handler + ) + billing_dispatcher_task = ( + asyncio.create_task(billing_dispatcher.run()) + if billing_dispatcher is not None + else None + ) + + try: + yield + finally: + if billing_dispatcher is not None: + await billing_dispatcher.stop() + if billing_dispatcher_task is not None: + await billing_dispatcher_task # # Shutdown scheduler.shutdown() diff --git a/py/core/main/assembly/factory.py b/py/core/main/assembly/factory.py index b01795107..41e00cc7e 100644 --- a/py/core/main/assembly/factory.py +++ b/py/core/main/assembly/factory.py @@ -18,6 +18,7 @@ OrchestrationConfig, SchedulerConfig, ) +from core.billing import BillingUsageRecorder from core.providers import ( AnthropicCompletionProvider, APSchedulerProvider, @@ -250,7 +251,10 @@ def create_file_provider( @staticmethod def create_embedding_provider( - embedding: EmbeddingConfig, *args, **kwargs + embedding: EmbeddingConfig, + *args, + billing_usage_recorder: BillingUsageRecorder | None = None, + **kwargs, ) -> ( LiteLLMEmbeddingProvider | OllamaEmbeddingProvider @@ -270,7 +274,10 @@ def create_embedding_provider( elif embedding.provider == "litellm": from core.providers import LiteLLMEmbeddingProvider - embedding_provider = LiteLLMEmbeddingProvider(embedding) + embedding_provider = LiteLLMEmbeddingProvider( + embedding, + billing_usage_recorder=billing_usage_recorder, + ) elif embedding.provider == "ollama": from core.providers import OllamaEmbeddingProvider @@ -286,7 +293,10 @@ def create_embedding_provider( @staticmethod def create_llm_provider( - llm_config: CompletionConfig, *args, **kwargs + llm_config: CompletionConfig, + *args, + billing_usage_recorder: BillingUsageRecorder | None = None, + **kwargs, ) -> ( AnthropicCompletionProvider | LiteLLMCompletionProvider @@ -297,7 +307,10 @@ def create_llm_provider( if llm_config.provider == "anthropic": llm_provider = AnthropicCompletionProvider(llm_config) elif llm_config.provider == "litellm": - llm_provider = LiteLLMCompletionProvider(llm_config) + llm_provider = LiteLLMCompletionProvider( + llm_config, + billing_usage_recorder=billing_usage_recorder, + ) elif llm_config.provider == "openai": llm_provider = OpenAICompletionProvider(llm_config) elif llm_config.provider == "r2r": @@ -398,34 +411,46 @@ async def create_providers( f"Both embedding configurations must use the same dimensions. Got {self.config.embedding.base_dimension} and {self.config.completion_embedding.base_dimension}" ) + crypto_provider = ( + crypto_provider_override + or self.create_crypto_provider(self.config.crypto, *args, **kwargs) + ) + + database_provider = ( + database_provider_override + or await self.create_database_provider( + self.config.database, crypto_provider, *args, **kwargs + ) + ) + billing_usage_recorder = BillingUsageRecorder( + database_provider.billing_outbox_handler + ) + embedding_provider = ( embedding_provider_override or self.create_embedding_provider( - self.config.embedding, *args, **kwargs + self.config.embedding, + *args, + billing_usage_recorder=billing_usage_recorder, + **kwargs, ) ) completion_embedding_provider = ( embedding_provider_override or self.create_embedding_provider( - self.config.completion_embedding, *args, **kwargs + self.config.completion_embedding, + *args, + billing_usage_recorder=billing_usage_recorder, + **kwargs, ) ) llm_provider = llm_provider_override or self.create_llm_provider( - self.config.completion, *args, **kwargs - ) - - crypto_provider = ( - crypto_provider_override - or self.create_crypto_provider(self.config.crypto, *args, **kwargs) - ) - - database_provider = ( - database_provider_override - or await self.create_database_provider( - self.config.database, crypto_provider, *args, **kwargs - ) + self.config.completion, + *args, + billing_usage_recorder=billing_usage_recorder, + **kwargs, ) file_provider = self.create_file_provider( diff --git a/py/core/main/orchestration/hatchet/ingestion_workflow.py b/py/core/main/orchestration/hatchet/ingestion_workflow.py index 85f031199..98c28dbf3 100644 --- a/py/core/main/orchestration/hatchet/ingestion_workflow.py +++ b/py/core/main/orchestration/hatchet/ingestion_workflow.py @@ -17,6 +17,7 @@ generate_extraction_id, ) from core.base.abstractions import DocumentResponse, R2RException +from core.billing import with_billing_context_from_hatchet from core.utils import ( generate_default_user_collection_id, num_tokens, @@ -58,6 +59,7 @@ def concurrency(self, context: Context) -> str: return str(uuid.uuid4()) @orchestration_provider.step(retries=0, timeout="60m") + @with_billing_context_from_hatchet async def parse(self, context: Context) -> dict: try: logger.info("Initiating ingestion workflow, step: parse") @@ -385,6 +387,7 @@ async def ingest(self, context: Context) -> dict: } @orchestration_provider.step(parents=["ingest"], timeout="60m") + @with_billing_context_from_hatchet async def embed(self, context: Context) -> dict: document_info_dict = context.step_output("ingest")["document_info"] document_info = DocumentResponse(**document_info_dict) diff --git a/py/core/main/orchestration/simple/ingestion_workflow.py b/py/core/main/orchestration/simple/ingestion_workflow.py index d30e24b42..e358e0d9f 100644 --- a/py/core/main/orchestration/simple/ingestion_workflow.py +++ b/py/core/main/orchestration/simple/ingestion_workflow.py @@ -10,6 +10,7 @@ GraphConstructionStatus, R2RException, ) +from core.billing import with_billing_context_from_input from core.utils import ( generate_default_user_collection_id, generate_extraction_id, @@ -23,6 +24,7 @@ def simple_ingestion_factory(service: IngestionService): + @with_billing_context_from_input async def ingest_files(input_data): document_info = None try: @@ -283,6 +285,7 @@ async def _ensure_collections_exists( ) raise e + @with_billing_context_from_input async def ingest_chunks(input_data): document_info = None try: diff --git a/py/core/parsers/media/audio_parser.py b/py/core/parsers/media/audio_parser.py index 7d5f9f1d4..af4aa6fe6 100644 --- a/py/core/parsers/media/audio_parser.py +++ b/py/core/parsers/media/audio_parser.py @@ -43,6 +43,11 @@ async def ingest( # type: ignore Yields: Chunks of transcribed text """ + billing_usage_recorder = getattr( + self.llm_provider, "billing_usage_recorder", None + ) + billing_event_id = None + temp_file_path = None try: # Create a temporary file to store the audio data with tempfile.NamedTemporaryFile( @@ -52,23 +57,47 @@ async def ingest( # type: ignore temp_file_path = temp_file.name # Call Whisper transcription - response = await self.atranscription( - model=self.config.audio_transcription_model - or self.config.app.audio_lm, - file=open(temp_file_path, "rb"), - **kwargs, + model = ( + self.config.audio_transcription_model + or self.config.app.audio_lm ) + transcription_kwargs = kwargs + if billing_usage_recorder is not None: + billing_event_id = await billing_usage_recorder.start_call( + operation="transcription", + model=model, + ) + transcription_kwargs = ( + billing_usage_recorder.add_litellm_metadata( + kwargs=kwargs, + event_id=billing_event_id, + ) + ) + + with open(temp_file_path, "rb") as audio_file: + response = await self.atranscription( + model=model, + file=audio_file, + **transcription_kwargs, + ) + if billing_usage_recorder is not None: + await billing_usage_recorder.complete_call( + billing_event_id, response + ) # The response should contain the transcribed text directly yield response.text except Exception as e: + if billing_usage_recorder is not None: + await billing_usage_recorder.fail_call(billing_event_id, e) logger.error(f"Error processing audio with Whisper: {str(e)}") raise finally: # Clean up the temporary file try: - os.unlink(temp_file_path) + if temp_file_path is not None: + os.unlink(temp_file_path) except Exception as e: logger.warning(f"Failed to delete temporary file: {str(e)}") diff --git a/py/core/providers/database/billing_outbox.py b/py/core/providers/database/billing_outbox.py new file mode 100644 index 000000000..ccfce8a01 --- /dev/null +++ b/py/core/providers/database/billing_outbox.py @@ -0,0 +1,197 @@ +from datetime import datetime +from typing import Any +from uuid import UUID + +from core.base import Handler +from core.billing import InterloomBillingContext + +from .base import PostgresConnectionManager + + +class PostgresBillingOutboxHandler(Handler): + TABLE_NAME = "interloom_billing_outbox" + + def __init__( + self, project_name: str, connection_manager: PostgresConnectionManager + ): + super().__init__(project_name, connection_manager) + + async def create_tables(self) -> None: + table = self._get_table_name(self.TABLE_NAME) + await self.connection_manager.execute_query( + f""" + CREATE TABLE IF NOT EXISTS {table} ( + event_id UUID PRIMARY KEY, + ingestion_id UUID NOT NULL, + organization_id UUID NOT NULL, + file_id UUID, + note_id UUID, + operation TEXT NOT NULL, + requested_model TEXT NOT NULL, + resolved_model TEXT, + provider_request_id TEXT, + status TEXT NOT NULL, + started_at TIMESTAMPTZ NOT NULL, + ended_at TIMESTAMPTZ, + input_tokens BIGINT, + output_tokens BIGINT, + cached_input_tokens BIGINT, + error TEXT, + delivery_attempts INTEGER NOT NULL DEFAULT 0, + next_delivery_at TIMESTAMPTZ, + delivered_at TIMESTAMPTZ, + last_delivery_error TEXT, + CHECK ((file_id IS NOT NULL)::int + (note_id IS NOT NULL)::int = 1), + CHECK (status IN ('pending', 'completed', 'failed')) + ); + CREATE INDEX IF NOT EXISTS idx_{self.project_name}_{self.TABLE_NAME}_delivery + ON {table} (next_delivery_at, started_at) + WHERE status = 'completed' AND delivered_at IS NULL; + CREATE INDEX IF NOT EXISTS idx_{self.project_name}_{self.TABLE_NAME}_ingestion + ON {table} (ingestion_id); + """ + ) + + async def create_pending( + self, + *, + event_id: UUID, + context: InterloomBillingContext, + operation: str, + requested_model: str, + started_at: datetime, + ) -> None: + await self.connection_manager.execute_query( + f""" + INSERT INTO {self._get_table_name(self.TABLE_NAME)} ( + event_id, ingestion_id, organization_id, file_id, note_id, + operation, requested_model, status, started_at + ) VALUES ($1, $2, $3, $4, $5, $6, $7, 'pending', $8) + """, + [ + event_id, + context.ingestion_id, + context.organization_id, + context.file_id, + context.note_id, + operation, + requested_model, + started_at, + ], + ) + + async def complete( + self, + *, + event_id: UUID, + provider_request_id: str | None, + resolved_model: str | None, + ended_at: datetime, + input_tokens: int | None, + output_tokens: int | None, + cached_input_tokens: int | None, + ) -> None: + await self.connection_manager.execute_query( + f""" + UPDATE {self._get_table_name(self.TABLE_NAME)} + SET status = 'completed', provider_request_id = $2, + resolved_model = $3, ended_at = $4, input_tokens = $5, + output_tokens = $6, cached_input_tokens = $7, + next_delivery_at = NOW() + WHERE event_id = $1 AND status = 'pending' + """, + [ + event_id, + provider_request_id, + resolved_model, + ended_at, + input_tokens, + output_tokens, + cached_input_tokens, + ], + ) + + async def fail( + self, + *, + event_id: UUID, + ended_at: datetime, + error: str, + ) -> None: + await self.connection_manager.execute_query( + f""" + UPDATE {self._get_table_name(self.TABLE_NAME)} + SET status = 'failed', ended_at = $2, error = $3 + WHERE event_id = $1 AND status = 'pending' + """, + [event_id, ended_at, error], + ) + + async def claim_for_delivery(self, *, limit: int) -> list[dict[str, Any]]: + rows = await self.connection_manager.fetch_query( + f""" + WITH claimed AS ( + SELECT event_id + FROM {self._get_table_name(self.TABLE_NAME)} + WHERE status = 'completed' AND delivered_at IS NULL + AND next_delivery_at <= NOW() + ORDER BY next_delivery_at, started_at + FOR UPDATE SKIP LOCKED + LIMIT $1 + ) + UPDATE {self._get_table_name(self.TABLE_NAME)} AS outbox + SET delivery_attempts = delivery_attempts + 1, + next_delivery_at = NOW() + INTERVAL '1 minute' + FROM claimed + WHERE outbox.event_id = claimed.event_id + RETURNING outbox.* + """, + [limit], + ) + return [self._serialize_event(dict(row)) for row in rows] + + async def mark_delivered(self, *, event_id: UUID | str) -> None: + await self.connection_manager.execute_query( + f""" + UPDATE {self._get_table_name(self.TABLE_NAME)} + SET delivered_at = NOW(), last_delivery_error = NULL + WHERE event_id = $1 + """, + [UUID(str(event_id))], + ) + + async def mark_delivery_failed( + self, *, event_id: UUID | str, error: str + ) -> None: + await self.connection_manager.execute_query( + f""" + UPDATE {self._get_table_name(self.TABLE_NAME)} + SET last_delivery_error = $2, + next_delivery_at = NOW() + + LEAST(POWER(2, delivery_attempts), 3600) + * INTERVAL '1 second' + WHERE event_id = $1 + """, + [UUID(str(event_id)), error], + ) + + @staticmethod + def _serialize_event(event: dict[str, Any]) -> dict[str, Any]: + serialized = { + key: value.isoformat() + if isinstance(value, datetime) + else str(value) + if isinstance(value, UUID) + else value + for key, value in event.items() + if key + not in { + "delivery_attempts", + "next_delivery_at", + "delivered_at", + "last_delivery_error", + } + } + serialized["schema_version"] = 1 + serialized["billing_event_id"] = serialized.pop("event_id") + return serialized diff --git a/py/core/providers/database/postgres.py b/py/core/providers/database/postgres.py index b921316df..392ff13c0 100644 --- a/py/core/providers/database/postgres.py +++ b/py/core/providers/database/postgres.py @@ -10,6 +10,7 @@ PostgresConfigurationSettings, ) from .base import PostgresConnectionManager, SemaphoreConnectionPool +from .billing_outbox import PostgresBillingOutboxHandler from .chunks import PostgresChunksHandler from .collections import PostgresCollectionsHandler from .conversations import PostgresConversationsHandler @@ -55,6 +56,7 @@ class PostgresDatabaseProvider(DatabaseProvider): default_collection_description: str connection_manager: PostgresConnectionManager + billing_outbox_handler: PostgresBillingOutboxHandler documents_handler: PostgresDocumentsHandler collections_handler: PostgresCollectionsHandler token_handler: PostgresTokensHandler @@ -134,6 +136,9 @@ def __init__( self.connection_manager: PostgresConnectionManager = ( PostgresConnectionManager() ) + self.billing_outbox_handler = PostgresBillingOutboxHandler( + self.project_name, self.connection_manager + ) self.documents_handler = PostgresDocumentsHandler( project_name=self.project_name, connection_manager=self.connection_manager, @@ -222,6 +227,7 @@ async def initialize(self): f'CREATE SCHEMA IF NOT EXISTS "{self.project_name}";' ) + await self.billing_outbox_handler.create_tables() await self.documents_handler.create_tables() await self.collections_handler.create_tables() await self.token_handler.create_tables() diff --git a/py/core/providers/embeddings/litellm.py b/py/core/providers/embeddings/litellm.py index 7322d7de4..7e7fdcffa 100644 --- a/py/core/providers/embeddings/litellm.py +++ b/py/core/providers/embeddings/litellm.py @@ -16,6 +16,7 @@ EmbeddingProvider, R2RException, ) +from core.billing import BillingUsageRecorder from .utils import truncate_texts_to_token_limit @@ -27,12 +28,14 @@ def __init__( self, config: EmbeddingConfig, *args, + billing_usage_recorder: BillingUsageRecorder | None = None, **kwargs, ) -> None: super().__init__(config) self.litellm_embedding = embedding self.litellm_aembedding = aembedding + self.billing_usage_recorder = billing_usage_recorder provider = config.provider if not provider: @@ -79,6 +82,7 @@ def _get_embedding_kwargs(self, **kwargs): async def _execute_task(self, task: dict[str, Any]) -> list[list[float]]: texts = task["texts"] kwargs = self._get_embedding_kwargs(**task.get("kwargs", {})) + billing_event_id = None if "dimensions" in kwargs and math.isnan(kwargs["dimensions"]): kwargs.pop("dimensions") @@ -92,17 +96,42 @@ async def _execute_task(self, task: dict[str, Any]) -> list[list[float]]: texts, kwargs["model"] ) + if self.billing_usage_recorder is not None: + billing_event_id = ( + await self.billing_usage_recorder.start_call( + operation="embedding", + model=kwargs["model"], + ) + ) + kwargs = self.billing_usage_recorder.add_litellm_metadata( + kwargs=kwargs, + event_id=billing_event_id, + ) + response = await self.litellm_aembedding( input=texts, **kwargs, ) + if self.billing_usage_recorder is not None: + await self.billing_usage_recorder.complete_call( + billing_event_id, response + ) return [data["embedding"] for data in response.data] - except AuthenticationError: + except AuthenticationError as e: + if self.billing_usage_recorder is not None: + await self.billing_usage_recorder.fail_call( + billing_event_id, + e, + ) logger.error( "Authentication error: Invalid API key or credentials." ) raise except Exception as e: + if self.billing_usage_recorder is not None: + await self.billing_usage_recorder.fail_call( + billing_event_id, e + ) error_msg = f"Error getting embeddings: {str(e)}" logger.error(error_msg) diff --git a/py/core/providers/llm/litellm.py b/py/core/providers/llm/litellm.py index 44d467c2a..134e65e76 100644 --- a/py/core/providers/llm/litellm.py +++ b/py/core/providers/llm/litellm.py @@ -6,16 +6,24 @@ from core.base.abstractions import GenerationConfig from core.base.providers.llm import CompletionConfig, CompletionProvider +from core.billing import BillingUsageRecorder logger = logging.getLogger() class LiteLLMCompletionProvider(CompletionProvider): - def __init__(self, config: CompletionConfig, *args, **kwargs) -> None: + def __init__( + self, + config: CompletionConfig, + *args, + billing_usage_recorder: BillingUsageRecorder | None = None, + **kwargs, + ) -> None: super().__init__(config) litellm.modify_params = True self.acompletion = acompletion self.completion = completion + self.billing_usage_recorder = billing_usage_recorder # if config.provider != "litellm": # logger.error(f"Invalid provider: {config.provider}") @@ -54,11 +62,35 @@ async def _execute_task(self, task: dict[str, Any]): args["messages"] = messages args = {**args, **kwargs} + billing_event_id = None + if self.billing_usage_recorder is not None: + billing_event_id = await self.billing_usage_recorder.start_call( + operation="completion", + model=args["model"], + ) + args = self.billing_usage_recorder.add_litellm_metadata( + kwargs=args, + event_id=billing_event_id, + ) + logger.debug( f"Executing LiteLLM task with generation_config={generation_config}" ) - return await self.acompletion(**args) + try: + response = await self.acompletion(**args) + except Exception as error: + if self.billing_usage_recorder is not None: + await self.billing_usage_recorder.fail_call( + billing_event_id, error + ) + raise + + if self.billing_usage_recorder is not None: + await self.billing_usage_recorder.complete_call( + billing_event_id, response + ) + return response def _execute_task_sync(self, task: dict[str, Any]): messages = task["messages"] diff --git a/py/tests/unit/test_billing.py b/py/tests/unit/test_billing.py new file mode 100644 index 000000000..86743f33e --- /dev/null +++ b/py/tests/unit/test_billing.py @@ -0,0 +1,189 @@ +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock +from uuid import uuid4 + +import httpx +import pytest +from pydantic import ValidationError + +from core.base import CompletionConfig, EmbeddingConfig, GenerationConfig +from core.billing import ( + BillingOutboxDispatcher, + BillingUsageRecorder, + InterloomBillingContext, + with_billing_context_from_input, +) +from core.providers.embeddings.litellm import LiteLLMEmbeddingProvider +from core.providers.llm.litellm import LiteLLMCompletionProvider + + +def test_billing_context_requires_exactly_one_source(): + values = { + "ingestion_id": uuid4(), + "organization_id": uuid4(), + } + + with pytest.raises(ValidationError): + InterloomBillingContext(**values) + + with pytest.raises(ValidationError): + InterloomBillingContext( + **values, + file_id=uuid4(), + note_id=uuid4(), + ) + + +async def test_recorder_persists_usage_and_adds_reconciliation_metadata(): + handler = SimpleNamespace( + create_pending=AsyncMock(), + complete=AsyncMock(), + fail=AsyncMock(), + ) + recorder = BillingUsageRecorder(handler) + context = InterloomBillingContext( + ingestion_id=uuid4(), + organization_id=uuid4(), + file_id=uuid4(), + ) + + @with_billing_context_from_input + async def record(input_data): + event_id = await recorder.start_call( + operation="completion", + model="litellm_proxy/openai/gpt-5.4-nano", + ) + kwargs = recorder.add_litellm_metadata( + kwargs={"metadata": {"existing": "value"}}, + event_id=event_id, + ) + await recorder.complete_call( + event_id, + { + "id": "provider-request-id", + "model": "openai/gpt-5.4-nano", + "usage": { + "prompt_tokens": 10, + "completion_tokens": 4, + "prompt_tokens_details": {"cached_tokens": 3}, + }, + }, + ) + return event_id, kwargs + + event_id, kwargs = await record( + {"interloom_billing_context": context.model_dump(mode="json")} + ) + + handler.create_pending.assert_awaited_once() + handler.complete.assert_awaited_once_with( + event_id=event_id, + provider_request_id="provider-request-id", + resolved_model="openai/gpt-5.4-nano", + ended_at=handler.complete.await_args.kwargs["ended_at"], + input_tokens=10, + output_tokens=4, + cached_input_tokens=3, + ) + assert kwargs["metadata"]["existing"] == "value" + assert kwargs["metadata"]["spend_logs_metadata"] == { + "source": "r2r", + "billing_event_id": str(event_id), + "ingestion_id": str(context.ingestion_id), + "organization_id": str(context.organization_id), + "file_id": str(context.file_id), + "note_id": None, + } + + assert ( + await recorder.start_call(operation="embedding", model="model") is None + ) + + +async def test_dispatcher_marks_successful_delivery(): + event_id = uuid4() + handler = SimpleNamespace( + claim_for_delivery=AsyncMock( + return_value=[ + {"billing_event_id": str(event_id), "input_tokens": 10} + ] + ), + mark_delivered=AsyncMock(), + mark_delivery_failed=AsyncMock(), + ) + dispatcher = BillingOutboxDispatcher( + outbox_handler=handler, + callback_url="https://interloom.example.test/r2r-usage", + service_token="service-token", + ) + + async def send(request: httpx.Request) -> httpx.Response: + assert request.headers["il-service-token"] == "service-token" + return httpx.Response(200) + + async with httpx.AsyncClient( + transport=httpx.MockTransport(send) + ) as client: + await dispatcher._deliver_batch(client) + + handler.mark_delivered.assert_awaited_once_with(event_id=str(event_id)) + handler.mark_delivery_failed.assert_not_awaited() + + +async def test_litellm_providers_record_completion_and_embedding_usage(): + event_id = uuid4() + recorder = SimpleNamespace( + start_call=AsyncMock(return_value=event_id), + add_litellm_metadata=Mock( + side_effect=lambda *, kwargs, event_id: { + **kwargs, + "metadata": {"billing_event_id": str(event_id)}, + } + ), + complete_call=AsyncMock(), + fail_call=AsyncMock(), + ) + response = SimpleNamespace( + id="provider-request-id", + model="openai/model", + usage=SimpleNamespace(prompt_tokens=2, completion_tokens=1), + data=[{"embedding": [0.1, 0.2]}], + ) + + completion_provider = LiteLLMCompletionProvider( + CompletionConfig(provider="litellm"), + billing_usage_recorder=recorder, + ) + completion_provider.acompletion = AsyncMock(return_value=response) + completion_result = await completion_provider._execute_task( + { + "messages": [{"role": "user", "content": "hello"}], + "generation_config": GenerationConfig(model="openai/model"), + "kwargs": {}, + } + ) + + assert completion_result is response + assert completion_provider.acompletion.await_args.kwargs["metadata"] == { + "billing_event_id": str(event_id) + } + recorder.complete_call.assert_awaited_with(event_id, response) + + embedding_provider = LiteLLMEmbeddingProvider( + EmbeddingConfig( + provider="litellm", + base_model="openai/embedding-model", + base_dimension=2, + ), + billing_usage_recorder=recorder, + ) + embedding_provider.litellm_aembedding = AsyncMock(return_value=response) + embedding_result = await embedding_provider._execute_task( + {"texts": ["hello"], "kwargs": {}} + ) + + assert embedding_result == [[0.1, 0.2]] + assert embedding_provider.litellm_aembedding.await_args.kwargs[ + "metadata" + ] == {"billing_event_id": str(event_id)} + recorder.complete_call.assert_awaited_with(event_id, response)