Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -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
239 changes: 239 additions & 0 deletions py/core/billing.py
Original file line number Diff line number Diff line change
@@ -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)
30 changes: 30 additions & 0 deletions py/core/main/api/v3/documents_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
WrappedIngestionResponse,
WrappedRelationshipsResponse,
)
from core.billing import InterloomBillingContext
from core.utils import update_settings_from_dict
from shared.abstractions import IngestionMode

Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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=(
Expand Down Expand Up @@ -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 = [
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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),
Expand All @@ -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"]
Expand Down
19 changes: 18 additions & 1 deletion py/core/main/app_entry.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import asyncio
import logging
import os
from contextlib import asynccontextmanager
Expand All @@ -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
Expand Down Expand Up @@ -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()
Expand Down
Loading
Loading