diff --git a/src/textql_sdk/_hooks/registration.py b/src/textql_sdk/_hooks/registration.py index 1a6daed6..4ff0be1b 100644 --- a/src/textql_sdk/_hooks/registration.py +++ b/src/textql_sdk/_hooks/registration.py @@ -2,6 +2,7 @@ import httpx +from .._version import __version__ from ..sdkconfiguration import SDKConfiguration from .types import BeforeRequestContext, BeforeRequestHook, Hooks, SDKInitHook @@ -30,6 +31,14 @@ def sdk_init(self, config: SDKConfiguration) -> SDKConfiguration: return config +class _SDKOriginHook(BeforeRequestHook): + def before_request( + self, hook_ctx: BeforeRequestContext, request: httpx.Request + ) -> httpx.Request: + request.headers["X-TextQL-SDK"] = f"python/{__version__}" + return request + + class _RPCPublicPrefixHook(BeforeRequestHook): """Connect RPCs are mounted under ``/rpc/public`` on the host, but the generated operations build paths like ``/textql.rpc.public./`` @@ -52,4 +61,5 @@ def init_hooks(hooks: Hooks): with an instance of a hook that implements that specific Hook interface Hooks are registered per SDK instance, and are valid for the lifetime of the SDK instance""" hooks.register_sdk_init_hook(_ServerURLFromEnvHook()) + hooks.register_before_request_hook(_SDKOriginHook()) hooks.register_before_request_hook(_RPCPublicPrefixHook()) diff --git a/src/textql_sdk/streaming.py b/src/textql_sdk/streaming.py index b39cd419..b7f50fbc 100644 --- a/src/textql_sdk/streaming.py +++ b/src/textql_sdk/streaming.py @@ -49,6 +49,7 @@ from .sdk import Textql from ._hooks.registration import server_url_from_env +from ._version import __version__ from .sdkconfiguration import SERVERS from ._connect.public.agent_connect import AgentServiceClient, AgentServiceClientSync from ._connect.public.apps_connect import AppServiceClient, AppServiceClientSync @@ -105,6 +106,7 @@ def __init__(self, api_key: str) -> None: async def on_start(self, ctx: RequestContext) -> None: ctx.request_headers()["tql_api_key"] = self._api_key + ctx.request_headers()["X-TextQL-SDK"] = f"python/{__version__}" # pylint: disable=unused-argument # Parameter names must match MetadataInterceptor exactly: the protocol does @@ -125,6 +127,7 @@ def __init__(self, api_key: str) -> None: def on_start_sync(self, ctx: RequestContext) -> None: ctx.request_headers()["tql_api_key"] = self._api_key + ctx.request_headers()["X-TextQL-SDK"] = f"python/{__version__}" # pylint: disable=unused-argument # See _ApiKeyInterceptor.on_end -- names are load-bearing for the protocol. diff --git a/tests/unit/test_sdk_origin.py b/tests/unit/test_sdk_origin.py new file mode 100644 index 00000000..81d40ee7 --- /dev/null +++ b/tests/unit/test_sdk_origin.py @@ -0,0 +1,161 @@ +import pyqwest +import pytest + +from textql_sdk._connect.public.agent_pb2 import ( + GetAgentRequest, + StreamAgentStatusRequest, +) +from textql_sdk._version import __user_agent__, __version__ +from textql_sdk.streaming import create_streaming_client, create_streaming_client_sync +from tests.conftest import AUTH_HEADER_NAME, FAKE_API_KEY, json_response + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("user_agent", [None, "customer-application/2.0"]) +@pytest.mark.parametrize("rpc_prefix", ["", "/rpc/public"]) +@pytest.mark.parametrize("sdk_marker", [None, "stale-client/0.0"]) +async def test_sdk_origin_header_on_transport( + make_sdk, asynchronous, user_agent, rpc_prefix, sdk_marker +): + bundle = make_sdk(lambda req: json_response(200, {})) + headers = { + "X-TextQL-Client": "customer-application", + "X-TextQL-Client-Version": "2.0", + "X-TextQL-Agent": "customer-agent", + } + if user_agent is not None: + headers["user-agent"] = user_agent + request_headers = dict(headers) + if sdk_marker is not None: + request_headers["x-textql-sdk"] = sdk_marker + + operation = ( + bundle.sdk.agents.get_agent_async + if asynchronous + else bundle.sdk.agents.get_agent + ) + result = operation( + agent_id="a1", + server_url=f"https://textql-sdk-tests.invalid{rpc_prefix}", + http_headers=request_headers, + ) + if asynchronous: + await result + + request = bundle.transport.last_request + assert request.headers["X-TextQL-SDK"] == f"python/{__version__}" + assert request.headers["User-Agent"] == (user_agent or __user_agent__) + for name, value in headers.items(): + assert request.headers[name] == value + assert request.headers[AUTH_HEADER_NAME] == FAKE_API_KEY + assert ( + request.url.path == "/rpc/public/textql.rpc.public.agent.AgentService/GetAgent" + ) + assert bundle.transport.body_json() == {"agentId": "a1"} + + +class RecordingConnectTransport: + def __init__(self): + self.requests = [] + + async def execute(self, request): + self.requests.append(request) + streaming = request.url.endswith("/StreamAgentStatus") + return pyqwest.Response( + status=200, + headers=pyqwest.Headers( + { + "content-type": ( + "application/connect+proto" + if streaming + else "application/proto" + ) + } + ), + content=b"\x02\x00\x00\x00\x02{}" if streaming else b"", + ) + + def execute_sync(self, request): + self.requests.append(request) + streaming = request.url.endswith("/StreamAgentStatus") + return pyqwest.SyncResponse( + status=200, + headers=pyqwest.Headers( + { + "content-type": ( + "application/connect+proto" + if streaming + else "application/proto" + ) + } + ), + content=b"\x02\x00\x00\x00\x02{}" if streaming else b"", + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("sdk_marker", [None, "stale-client/0.0"]) +async def test_connect_sdk_origin_header_on_transport( + asynchronous, streaming, sdk_marker +): + transport = RecordingConnectTransport() + headers = { + "X-TextQL-Client": "customer-application", + "X-TextQL-Client-Version": "2.0", + "X-TextQL-Agent": "customer-agent", + "user-agent": "customer-application/2.0", + } + request_headers = dict(headers) + if sdk_marker is not None: + request_headers["x-textql-sdk"] = sdk_marker + if asynchronous: + client = create_streaming_client( + api_key=FAKE_API_KEY, + server_url="https://textql-sdk-tests.invalid", + http_client=pyqwest.Client(transport), + ) + if streaming: + messages = [ + message + async for message in client.agents.stream_agent_status( + StreamAgentStatusRequest(), headers=request_headers + ) + ] + assert messages == [] + else: + await client.agents.get_agent( + GetAgentRequest(agent_id="a1"), headers=request_headers + ) + else: + client = create_streaming_client_sync( + api_key=FAKE_API_KEY, + server_url="https://textql-sdk-tests.invalid", + http_client=pyqwest.SyncClient(transport), + ) + if streaming: + assert ( + list( + client.agents.stream_agent_status( + StreamAgentStatusRequest(), headers=request_headers + ) + ) + == [] + ) + else: + client.agents.get_agent( + GetAgentRequest(agent_id="a1"), headers=request_headers + ) + + assert len(transport.requests) == 1 + request = transport.requests[0] + assert request.headers["X-TextQL-SDK"] == f"python/{__version__}" + for name, value in headers.items(): + assert request.headers[name] == value + assert request.headers[AUTH_HEADER_NAME] == FAKE_API_KEY + operation = "StreamAgentStatus" if streaming else "GetAgent" + assert request.url == ( + f"https://textql-sdk-tests.invalid/rpc/public/textql.rpc.public.agent.AgentService/{operation}" + )