Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
8 changes: 8 additions & 0 deletions livekit-api/livekit/api/twirp_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

import asyncio
import logging
import uuid
from typing import Dict, List, Optional, Type, TypeVar

import aiohttp
Expand All @@ -37,6 +38,11 @@
# Identifies the SDK and version to the server on every request.
_USER_AGENT = f"livekit-server-sdk-python/{__version__}"

# Carries a per-request idempotency key. The SDK's auto-retries (see _failover)
# keep the same key across attempts, so the server can identify and deduplicate
# repeated requests.
REQUEST_ID_HEADER = "X-Livekit-Request-Id"

# Shared across all clients in the process so the region list is fetched once.
_REGION_CACHE = RegionCache()

Expand Down Expand Up @@ -207,6 +213,8 @@ async def request(
headers["User-Agent"] = _USER_AGENT
forward_headers = dict(headers) # for the discovery fetch (no content-type yet)
headers["Content-Type"] = "application/protobuf"
if not any(h.lower() == REQUEST_ID_HEADER.lower() for h in headers):
headers[REQUEST_ID_HEADER] = str(uuid.uuid4())
serialized_data = data.SerializeToString()

# The effective per-attempt timeout is the per-call override, or the
Expand Down
50 changes: 47 additions & 3 deletions tests/api/test_failover.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,12 +28,13 @@
import json
import os
import urllib.request
from typing import List, Optional

import aiohttp
import pytest

from livekit.api import CreateRoomRequest, Room, ServerError
from livekit.api.twirp_client import TwirpClient
from livekit.api.twirp_client import REQUEST_ID_HEADER, TwirpClient

BASE = os.getenv("LK_TEST_SERVER_URL", "http://127.0.0.1:9999")

Expand All @@ -53,8 +54,15 @@ def _server_up() -> bool:

# _failover_force bypasses the cloud-host check (the mock is on 127.0.0.1) and a
# tiny backoff keeps the tests fast — both are internal, test-only knobs.
async def _call(mock: dict, *, failover: bool = True, force: bool = True) -> Room:
async with aiohttp.ClientSession() as session:
async def _call(
mock: dict,
*,
failover: bool = True,
force: bool = True,
extra_headers: Optional[dict] = None,
trace_configs: Optional[List[aiohttp.TraceConfig]] = None,
) -> Room:
async with aiohttp.ClientSession(trace_configs=trace_configs) as session:
client = TwirpClient(
session,
BASE,
Expand All @@ -67,6 +75,7 @@ async def _call(mock: dict, *, failover: bool = True, force: bool = True) -> Roo
"authorization": "Bearer test-token",
# These tests exercise failover, not authz; skip the mock's permission check.
"X-Lk-Mock": json.dumps({"skipAuth": True, **mock}),
**(extra_headers or {}),
}
return await client.request("RoomService", "CreateRoom", CreateRoomRequest(), headers, Room)

Expand Down Expand Up @@ -113,3 +122,38 @@ def test_disabled():
# failover=False disables failover entirely.
with pytest.raises(ServerError):
asyncio.run(_call({"failRegions": [0]}, failover=False))


# Records the request id header(s) the SDK put on the wire for each Twirp
# attempt. Region discovery is a separate request, so it is not recorded.
def _request_id_recorder(seen: List[List[str]]) -> aiohttp.TraceConfig:
trace = aiohttp.TraceConfig()

async def on_request_start(_session, _ctx, params) -> None:
if not params.url.path.endswith("/settings/regions"):
seen.append(list(params.headers.getall(REQUEST_ID_HEADER, [])))

trace.on_request_start.append(on_request_start)
return trace


def test_request_id_stable_across_attempts():
# The id is generated once per logical call, so a replayed request carries
# the same idempotency key on every attempt and the server can dedup it.
seen: List[List[str]] = []
asyncio.run(_call({"failRegions": [0, 1]}, trace_configs=[_request_id_recorder(seen)]))
assert len(seen) == 3 # primary + two fallbacks
assert all(len(ids) == 1 for ids in seen) # never duplicated
assert seen[0][0]
assert len({ids[0] for ids in seen}) == 1


def test_request_id_unique_per_call():
# A new logical call is a new request, so it gets its own id.
seen: List[List[str]] = []
recorder = _request_id_recorder(seen)
asyncio.run(_call({}, trace_configs=[recorder]))
asyncio.run(_call({}, trace_configs=[recorder]))
assert len(seen) == 2
assert seen[0][0] and seen[1][0]
assert seen[0][0] != seen[1][0]
Loading