Compare commits
3
Commits
release/0.0.2
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
148a35b71c | ||
|
|
69762d93df | ||
|
|
aa35e2d30c
|
@@ -56,6 +56,14 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
.venv/bin/python -m mypy -p kaya.cors
|
.venv/bin/python -m mypy -p kaya.cors
|
||||||
.venv/bin/python -m unittest discover -s packages/kaya-cors/tests
|
.venv/bin/python -m unittest discover -s packages/kaya-cors/tests
|
||||||
|
- name: Check kaya-forwarded
|
||||||
|
run: |
|
||||||
|
.venv/bin/python -m mypy -p kaya.forwarded
|
||||||
|
.venv/bin/python -m unittest discover -s packages/kaya-forwarded/tests
|
||||||
|
- name: Check kaya-otel
|
||||||
|
run: |
|
||||||
|
.venv/bin/python -m mypy -p kaya.otel
|
||||||
|
.venv/bin/python -m unittest discover -s packages/kaya-otel/tests
|
||||||
- name: Publish kaya-core artifacts
|
- name: Publish kaya-core artifacts
|
||||||
env:
|
env:
|
||||||
TWINE_REPOSITORY_URL: ${{ vars.PYPI_REGISTRY_URL }}
|
TWINE_REPOSITORY_URL: ${{ vars.PYPI_REGISTRY_URL }}
|
||||||
@@ -120,3 +128,19 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
.venv/bin/pyproject-build packages/kaya-cors
|
.venv/bin/pyproject-build packages/kaya-cors
|
||||||
.venv/bin/twine upload --repository gitea packages/kaya-cors/dist/*.whl packages/kaya-cors/dist/*.tar.gz
|
.venv/bin/twine upload --repository gitea packages/kaya-cors/dist/*.whl packages/kaya-cors/dist/*.tar.gz
|
||||||
|
- name: Publish kaya-forwarded artifacts
|
||||||
|
env:
|
||||||
|
TWINE_REPOSITORY_URL: ${{ vars.PYPI_REGISTRY_URL }}
|
||||||
|
TWINE_USERNAME: ${{ vars.PUBLISHER_USERNAME }}
|
||||||
|
TWINE_PASSWORD: ${{ secrets.PUBLISHER_TOKEN }}
|
||||||
|
run: |
|
||||||
|
.venv/bin/pyproject-build packages/kaya-forwarded
|
||||||
|
.venv/bin/twine upload --repository gitea packages/kaya-forwarded/dist/*.whl packages/kaya-forwarded/dist/*.tar.gz
|
||||||
|
- name: Publish kaya-otel artifacts
|
||||||
|
env:
|
||||||
|
TWINE_REPOSITORY_URL: ${{ vars.PYPI_REGISTRY_URL }}
|
||||||
|
TWINE_USERNAME: ${{ vars.PUBLISHER_USERNAME }}
|
||||||
|
TWINE_PASSWORD: ${{ secrets.PUBLISHER_TOKEN }}
|
||||||
|
run: |
|
||||||
|
.venv/bin/pyproject-build packages/kaya-otel
|
||||||
|
.venv/bin/twine upload --repository gitea packages/kaya-otel/dist/*.whl packages/kaya-otel/dist/*.tar.gz
|
||||||
|
|||||||
@@ -14,6 +14,8 @@ This repository is a monorepo for the Kaya framework. The code is split into ind
|
|||||||
- **kaya-oidc** — OpenID Connect authentication (`packages/kaya-oidc/`)
|
- **kaya-oidc** — OpenID Connect authentication (`packages/kaya-oidc/`)
|
||||||
- **kaya-openapi** — automatic OpenAPI specification generation (`packages/kaya-openapi/`)
|
- **kaya-openapi** — automatic OpenAPI specification generation (`packages/kaya-openapi/`)
|
||||||
- **kaya-cors** — CORS (Cross-Origin Resource Sharing) support (`packages/kaya-cors/`)
|
- **kaya-cors** — CORS (Cross-Origin Resource Sharing) support (`packages/kaya-cors/`)
|
||||||
|
- **kaya-forwarded** — trusted-proxy `Forwarded`/`X-Forwarded-*` client address resolution (`packages/kaya-forwarded/`)
|
||||||
|
- **kaya-otel** — OpenTelemetry tracing and metrics (`packages/kaya-otel/`)
|
||||||
|
|
||||||
Additional `kaya-*` packages can be added as new directories under `packages/`.
|
Additional `kaya-*` packages can be added as new directories under `packages/`.
|
||||||
|
|
||||||
@@ -28,7 +30,7 @@ pip install --index-url https://gitea.woggioni.net/api/packages/woggioni/pypi/si
|
|||||||
Install the packages in development mode:
|
Install the packages in development mode:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
pip install -e packages/kaya-core -e packages/kaya-rsgi -e packages/kaya-session -e packages/kaya-session-redis -e packages/kaya-session-memcache -e packages/kaya-oidc -e packages/kaya-openapi -e packages/kaya-cors
|
pip install -e packages/kaya-core -e packages/kaya-rsgi -e packages/kaya-session -e packages/kaya-session-redis -e packages/kaya-session-memcache -e packages/kaya-oidc -e packages/kaya-openapi -e packages/kaya-cors -e packages/kaya-forwarded -e packages/kaya-otel
|
||||||
```
|
```
|
||||||
|
|
||||||
Run the example:
|
Run the example:
|
||||||
@@ -48,6 +50,8 @@ python -m unittest discover -s packages/kaya-session-memcache/tests
|
|||||||
python -m unittest discover -s packages/kaya-oidc/tests
|
python -m unittest discover -s packages/kaya-oidc/tests
|
||||||
python -m unittest discover -s packages/kaya-openapi/tests
|
python -m unittest discover -s packages/kaya-openapi/tests
|
||||||
python -m unittest discover -s packages/kaya-cors/tests
|
python -m unittest discover -s packages/kaya-cors/tests
|
||||||
|
python -m unittest discover -s packages/kaya-forwarded/tests
|
||||||
|
python -m unittest discover -s packages/kaya-otel/tests
|
||||||
```
|
```
|
||||||
|
|
||||||
## Static analysis
|
## Static analysis
|
||||||
@@ -61,6 +65,8 @@ mypy -p kaya.session.memcache
|
|||||||
mypy -p kaya.oidc
|
mypy -p kaya.oidc
|
||||||
mypy -p kaya.openapi
|
mypy -p kaya.openapi
|
||||||
mypy -p kaya.cors
|
mypy -p kaya.cors
|
||||||
|
mypy -p kaya.forwarded
|
||||||
|
mypy -p kaya.otel
|
||||||
```
|
```
|
||||||
|
|
||||||
## Building packages
|
## Building packages
|
||||||
@@ -74,4 +80,6 @@ python -m build packages/kaya-session-memcache
|
|||||||
python -m build packages/kaya-oidc
|
python -m build packages/kaya-oidc
|
||||||
python -m build packages/kaya-openapi
|
python -m build packages/kaya-openapi
|
||||||
python -m build packages/kaya-cors
|
python -m build packages/kaya-cors
|
||||||
|
python -m build packages/kaya-forwarded
|
||||||
|
python -m build packages/kaya-otel
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from ._app import AbstractKayaApp, KayaApp
|
from ._app import AbstractKayaApp, KayaApp
|
||||||
from ._http_method import HttpMethod
|
from ._http_method import HttpMethod
|
||||||
from ._http_context import HttpContext, resolve_client
|
from ._http_context import HttpContext
|
||||||
from ._mixin import KayaMixin
|
from ._mixin import KayaMixin
|
||||||
from ._tree import Tree, PathIterator
|
from ._tree import Tree, PathIterator
|
||||||
from ._path_handler import PathHandler, Matches
|
from ._path_handler import PathHandler, Matches
|
||||||
@@ -13,7 +13,6 @@ __all__ = [
|
|||||||
'KayaApp',
|
'KayaApp',
|
||||||
'KayaMixin',
|
'KayaMixin',
|
||||||
'HttpContext',
|
'HttpContext',
|
||||||
'resolve_client',
|
|
||||||
'Tree',
|
'Tree',
|
||||||
'PathHandler',
|
'PathHandler',
|
||||||
'Matches',
|
'Matches',
|
||||||
|
|||||||
@@ -145,6 +145,9 @@ class KayaApp(AbstractKayaApp):
|
|||||||
await handler.handle_request(ctx, captured)
|
await handler.handle_request(ctx, captured)
|
||||||
else:
|
else:
|
||||||
await ctx.send_empty(404)
|
await ctx.send_empty(404)
|
||||||
|
except Exception as exc:
|
||||||
|
ctx.exception = exc
|
||||||
|
raise
|
||||||
finally:
|
finally:
|
||||||
for hook in reversed(self._after_request_hooks):
|
for hook in reversed(self._after_request_hooks):
|
||||||
await hook(ctx)
|
await hook(ctx)
|
||||||
@@ -161,6 +164,9 @@ class KayaApp(AbstractKayaApp):
|
|||||||
await handler.handle_request(ws, captured)
|
await handler.handle_request(ws, captured)
|
||||||
else:
|
else:
|
||||||
await ws.close(1000)
|
await ws.close(1000)
|
||||||
|
except Exception as exc:
|
||||||
|
ws.exception = exc
|
||||||
|
raise
|
||||||
finally:
|
finally:
|
||||||
for hook in reversed(self._after_websocket_hooks):
|
for hook in reversed(self._after_websocket_hooks):
|
||||||
await hook(ws)
|
await hook(ws)
|
||||||
@@ -173,6 +179,10 @@ class KayaApp(AbstractKayaApp):
|
|||||||
for mixin in self._mixins:
|
for mixin in self._mixins:
|
||||||
mixin.shutdown(loop)
|
mixin.shutdown(loop)
|
||||||
|
|
||||||
|
def route_template(self, url: str, method: HttpMethod = HttpMethod.GET) -> Optional[str]:
|
||||||
|
"""Return the registered route template that would handle ``url``."""
|
||||||
|
return self._tree.route_template(url, method)
|
||||||
|
|
||||||
def route(self,
|
def route(self,
|
||||||
paths: StrOrStrings,
|
paths: StrOrStrings,
|
||||||
methods: Optional[HttpMethod | Sequence[HttpMethod]] = None,
|
methods: Optional[HttpMethod | Sequence[HttpMethod]] = None,
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ from typing import (
|
|||||||
from pwo import Maybe
|
from pwo import Maybe
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from ._http_method import HttpMethod
|
from ._http_method import HttpMethod
|
||||||
from ._http_context import HttpContext, resolve_client
|
from ._http_context import HttpContext
|
||||||
from ._websocket import WebSocket, WebSocketMessage
|
from ._websocket import WebSocket, WebSocketMessage
|
||||||
from ._types import StrOrStrings
|
from ._types import StrOrStrings
|
||||||
from ._types.asgi import HTTPScope, WebSocketScope
|
from ._types.asgi import HTTPScope, WebSocketScope
|
||||||
@@ -84,9 +84,9 @@ class AsgiContext(HttpContext):
|
|||||||
self.query_string = scope['query_string'].decode()
|
self.query_string = scope['query_string'].decode()
|
||||||
self.method = HttpMethod(scope['method'])
|
self.method = HttpMethod(scope['method'])
|
||||||
self.scheme = scope.get('scheme', 'http')
|
self.scheme = scope.get('scheme', 'http')
|
||||||
self.headers = decode_headers(scope['headers'])
|
self.client = scope['client']
|
||||||
self.client = resolve_client(self.headers, scope['client'])
|
|
||||||
self.server = scope['server']
|
self.server = scope['server']
|
||||||
|
self.headers = decode_headers(scope['headers'])
|
||||||
self.request_body = request_body_iterator
|
self.request_body = request_body_iterator
|
||||||
self.session = (scope.get('state') or {}).get('kaya_session')
|
self.session = (scope.get('state') or {}).get('kaya_session')
|
||||||
|
|
||||||
@@ -169,9 +169,9 @@ class AsgiWebSocket(WebSocket):
|
|||||||
self.path = scope['path']
|
self.path = scope['path']
|
||||||
self.query_string = scope['query_string'].decode()
|
self.query_string = scope['query_string'].decode()
|
||||||
self.scheme = scope.get('scheme', 'ws')
|
self.scheme = scope.get('scheme', 'ws')
|
||||||
self.headers = decode_headers(scope['headers'])
|
self.client = scope['client']
|
||||||
self.client = resolve_client(self.headers, scope['client'])
|
|
||||||
self.server = scope['server']
|
self.server = scope['server']
|
||||||
|
self.headers = decode_headers(scope['headers'])
|
||||||
|
|
||||||
async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
message: Dict[str, Any] = {'type': 'websocket.accept'}
|
message: Dict[str, Any] = {'type': 'websocket.accept'}
|
||||||
|
|||||||
@@ -16,83 +16,6 @@ from ._http_method import HttpMethod
|
|||||||
from ._types.base import StrOrStrings
|
from ._types.base import StrOrStrings
|
||||||
|
|
||||||
|
|
||||||
def _split_host_port(value: str) -> Tuple[str, Optional[int]]:
|
|
||||||
if value.startswith('['):
|
|
||||||
# bracketed IPv6 address, optionally followed by :port
|
|
||||||
closing = value.find(']')
|
|
||||||
if closing == -1:
|
|
||||||
return value, None
|
|
||||||
host = value[1:closing]
|
|
||||||
rest = value[closing + 1:]
|
|
||||||
if rest.startswith(':'):
|
|
||||||
try:
|
|
||||||
return host, int(rest[1:])
|
|
||||||
except ValueError:
|
|
||||||
return host, None
|
|
||||||
return host, None
|
|
||||||
if value.count(':') == 1:
|
|
||||||
host, _, port_str = value.rpartition(':')
|
|
||||||
try:
|
|
||||||
return host, int(port_str)
|
|
||||||
except ValueError:
|
|
||||||
return value, None
|
|
||||||
return value, None
|
|
||||||
|
|
||||||
|
|
||||||
def _first_header_value(headers: Mapping[str, Sequence[str]], name: str) -> Optional[str]:
|
|
||||||
values = headers.get(name)
|
|
||||||
if not values:
|
|
||||||
return None
|
|
||||||
first = values[0].split(',')[0].strip()
|
|
||||||
return first or None
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_client(headers: Mapping[str, Sequence[str]],
|
|
||||||
client: Optional[Tuple[str, int]]) -> Optional[Tuple[str, int]]:
|
|
||||||
"""
|
|
||||||
Resolve the client (host, port) pair, honoring forwarded headers.
|
|
||||||
|
|
||||||
Resolution order:
|
|
||||||
1. the ``for=`` parameter of the first entry of the RFC 7239 ``Forwarded`` header
|
|
||||||
(a ``:port`` suffix, if present, also populates the port)
|
|
||||||
2. the first entry of ``X-Forwarded-For``
|
|
||||||
3. the first entry of ``X-Forwarded-Host``
|
|
||||||
4. the socket peer address (``client``), returned unchanged
|
|
||||||
|
|
||||||
In the ``X-Forwarded-*`` cases the port is taken from ``X-Forwarded-Port``
|
|
||||||
when present and valid, otherwise the socket port is kept.
|
|
||||||
"""
|
|
||||||
socket_port = client[1] if client is not None else 0
|
|
||||||
|
|
||||||
forwarded_values = headers.get('forwarded')
|
|
||||||
if forwarded_values:
|
|
||||||
for raw_value in forwarded_values:
|
|
||||||
first_entry = raw_value.split(',')[0]
|
|
||||||
for param in first_entry.split(';'):
|
|
||||||
key, sep, value = param.partition('=')
|
|
||||||
if sep and key.strip().lower() == 'for':
|
|
||||||
for_value = value.strip().strip('"')
|
|
||||||
if for_value and for_value.lower() != 'unknown':
|
|
||||||
forwarded_host, forwarded_port = _split_host_port(for_value)
|
|
||||||
if forwarded_host:
|
|
||||||
return (forwarded_host,
|
|
||||||
forwarded_port if forwarded_port is not None else socket_port)
|
|
||||||
|
|
||||||
for header_name in ('x-forwarded-for', 'x-forwarded-host'):
|
|
||||||
host = _first_header_value(headers, header_name)
|
|
||||||
if host is not None:
|
|
||||||
port = socket_port
|
|
||||||
x_forwarded_port = _first_header_value(headers, 'x-forwarded-port')
|
|
||||||
if x_forwarded_port is not None:
|
|
||||||
try:
|
|
||||||
port = int(x_forwarded_port)
|
|
||||||
except ValueError:
|
|
||||||
pass
|
|
||||||
return host, port
|
|
||||||
|
|
||||||
return client
|
|
||||||
|
|
||||||
|
|
||||||
class HttpContext(ABC):
|
class HttpContext(ABC):
|
||||||
pathsend: bool
|
pathsend: bool
|
||||||
receive: Callable[[], Awaitable[Any]]
|
receive: Callable[[], Awaitable[Any]]
|
||||||
@@ -106,6 +29,7 @@ class HttpContext(ABC):
|
|||||||
server: Optional[Tuple[str, Optional[int]]]
|
server: Optional[Tuple[str, Optional[int]]]
|
||||||
request_body: AsyncIterator[bytes]
|
request_body: AsyncIterator[bytes]
|
||||||
session: Optional[Any] = None
|
session: Optional[Any] = None
|
||||||
|
exception: Optional[BaseException] = None
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def stream_body(self,
|
async def stream_body(self,
|
||||||
|
|||||||
@@ -29,9 +29,22 @@ class KayaMixin(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
def setup(self, loop: AbstractEventLoop) -> None:
|
def setup(self, loop: AbstractEventLoop) -> None:
|
||||||
"""Called on lifespan startup (default: no-op)."""
|
"""Called on lifespan startup (default: no-op).
|
||||||
|
|
||||||
|
Under RSGI (granian) the loop is NOT yet running when this is
|
||||||
|
called: schedule work on the passed ``loop`` (e.g.
|
||||||
|
``loop.create_task``) instead of calling
|
||||||
|
``asyncio.get_running_loop()``, which raises ``RuntimeError``
|
||||||
|
there. Under ASGI the lifespan runs inside a coroutine, so a
|
||||||
|
running loop happens to be available — do not rely on that.
|
||||||
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def shutdown(self, loop: AbstractEventLoop) -> None:
|
def shutdown(self, loop: AbstractEventLoop) -> None:
|
||||||
"""Called on lifespan shutdown (default: no-op)."""
|
"""Called on lifespan shutdown (default: no-op).
|
||||||
|
|
||||||
|
As with :meth:`setup`, use the passed ``loop``; it may already be
|
||||||
|
stopped under RSGI, so ``asyncio.get_running_loop()`` is not
|
||||||
|
reliable here either.
|
||||||
|
"""
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -237,6 +237,37 @@ class Tree:
|
|||||||
# return (handler, unmatched)
|
# return (handler, unmatched)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def route_template(self, url: str, method: HttpMethod = HttpMethod.GET) -> Optional[str]:
|
||||||
|
"""Return the registered route template that would handle ``url``.
|
||||||
|
|
||||||
|
Static segments are returned verbatim, ``${name}``/``${name:int}``
|
||||||
|
parameter matchers are reconstructed from the matched nodes, and glob
|
||||||
|
matchers keep their pattern. Returns ``None`` when no route (including
|
||||||
|
recursive fallback routes) would handle the request.
|
||||||
|
"""
|
||||||
|
result = self.find_node((p for p in PathIterator(urlparse(url).path)), method)
|
||||||
|
if result is None:
|
||||||
|
return None
|
||||||
|
node, captured = result
|
||||||
|
if len(captured.unmatched_paths) > 0 and not any(handler.recursive for handler in node.handlers):
|
||||||
|
return None
|
||||||
|
|
||||||
|
parts: List[str] = []
|
||||||
|
current: Optional[Node | PathMatcher] = node
|
||||||
|
while current is not None:
|
||||||
|
if isinstance(current, Node):
|
||||||
|
key = current.key
|
||||||
|
if not isinstance(key, HttpMethod) and key != '/':
|
||||||
|
parts.append(str(key))
|
||||||
|
elif isinstance(current, IntMatcher):
|
||||||
|
parts.append('${%s:int}' % current.name)
|
||||||
|
elif isinstance(current, StrMatcher):
|
||||||
|
parts.append('${%s}' % current.name)
|
||||||
|
elif isinstance(current, GlobMatcher):
|
||||||
|
parts.append(current.pattern)
|
||||||
|
current = current.parent
|
||||||
|
return '/' + '/'.join(reversed(parts)) if parts else '/'
|
||||||
|
|
||||||
def parse(self, leaf: str, parent: Optional[Node | PathMatcher]) -> Node | PathMatcher:
|
def parse(self, leaf: str, parent: Optional[Node | PathMatcher]) -> Node | PathMatcher:
|
||||||
start = 0
|
start = 0
|
||||||
result = index_of_with_escape(leaf, '${', '\\', 0)
|
result = index_of_with_escape(leaf, '${', '\\', 0)
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ class WebSocket(ABC):
|
|||||||
client: Optional[Tuple[str, int]]
|
client: Optional[Tuple[str, int]]
|
||||||
server: Optional[Tuple[str, Optional[int]]]
|
server: Optional[Tuple[str, Optional[int]]]
|
||||||
session: Optional[Any] = None
|
session: Optional[Any] = None
|
||||||
|
exception: Optional[BaseException] = None
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
|||||||
@@ -62,11 +62,6 @@ class AsgiTest(unittest.TestCase):
|
|||||||
async def handle_request(ctx: HttpContext, _: List[str]) -> None:
|
async def handle_request(ctx: HttpContext, _: List[str]) -> None:
|
||||||
await ctx.stream_body(200, (chunk async for chunk in ctx.request_body))
|
await ctx.stream_body(200, (chunk async for chunk in ctx.request_body))
|
||||||
|
|
||||||
@self.app.GET('/client-ip')
|
|
||||||
async def handle_request(ctx: HttpContext) -> None:
|
|
||||||
host, port = ctx.client if ctx.client is not None else (None, None)
|
|
||||||
await ctx.send_str(200, json.dumps({'host': host, 'port': port}))
|
|
||||||
|
|
||||||
@async_test
|
@async_test
|
||||||
async def test_hello(self):
|
async def test_hello(self):
|
||||||
transport = httpx.ASGITransport(app=self.app)
|
transport = httpx.ASGITransport(app=self.app)
|
||||||
@@ -195,60 +190,6 @@ class AsgiTest(unittest.TestCase):
|
|||||||
'employee_id': 101325
|
'employee_id': 101325
|
||||||
}, response)
|
}, response)
|
||||||
|
|
||||||
@async_test
|
|
||||||
async def test_client_ip_forwarded(self):
|
|
||||||
transport = httpx.ASGITransport(app=self.app)
|
|
||||||
|
|
||||||
async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as client:
|
|
||||||
# socket peer, no forwarded headers
|
|
||||||
r = await client.get("/client-ip")
|
|
||||||
socket_client = json.loads(r.text)
|
|
||||||
self.assertEqual('127.0.0.1', socket_client['host'])
|
|
||||||
|
|
||||||
# RFC 7239 Forwarded header, with port
|
|
||||||
r = await client.get("/client-ip", headers={'Forwarded': 'for=203.0.113.5:1234'})
|
|
||||||
self.assertEqual({'host': '203.0.113.5', 'port': 1234}, json.loads(r.text))
|
|
||||||
|
|
||||||
# RFC 7239 Forwarded header, bracketed IPv6 with port
|
|
||||||
r = await client.get("/client-ip", headers={'Forwarded': 'for="[2001:db8::1]:4711"'})
|
|
||||||
self.assertEqual({'host': '2001:db8::1', 'port': 4711}, json.loads(r.text))
|
|
||||||
|
|
||||||
# RFC 7239 Forwarded header without port keeps the socket port
|
|
||||||
r = await client.get("/client-ip", headers={'Forwarded': 'for=203.0.113.5'})
|
|
||||||
self.assertEqual({'host': '203.0.113.5', 'port': socket_client['port']}, json.loads(r.text))
|
|
||||||
|
|
||||||
# Forwarded with for=unknown falls through to X-Forwarded-For
|
|
||||||
r = await client.get("/client-ip", headers={
|
|
||||||
'Forwarded': 'for=unknown',
|
|
||||||
'X-Forwarded-For': '198.51.100.7',
|
|
||||||
})
|
|
||||||
self.assertEqual('198.51.100.7', json.loads(r.text)['host'])
|
|
||||||
|
|
||||||
# Forwarded takes precedence over X-Forwarded-For
|
|
||||||
r = await client.get("/client-ip", headers={
|
|
||||||
'Forwarded': 'for=203.0.113.5',
|
|
||||||
'X-Forwarded-For': '198.51.100.7',
|
|
||||||
})
|
|
||||||
self.assertEqual('203.0.113.5', json.loads(r.text)['host'])
|
|
||||||
|
|
||||||
# X-Forwarded-For: first entry of the chain, port from X-Forwarded-Port
|
|
||||||
r = await client.get("/client-ip", headers={
|
|
||||||
'X-Forwarded-For': '203.0.113.5, 70.41.3.18',
|
|
||||||
'X-Forwarded-Port': '8443',
|
|
||||||
})
|
|
||||||
self.assertEqual({'host': '203.0.113.5', 'port': 8443}, json.loads(r.text))
|
|
||||||
|
|
||||||
# X-Forwarded-Host fallback
|
|
||||||
r = await client.get("/client-ip", headers={'X-Forwarded-Host': '198.51.100.7'})
|
|
||||||
self.assertEqual('198.51.100.7', json.loads(r.text)['host'])
|
|
||||||
|
|
||||||
# invalid X-Forwarded-Port is ignored, socket port is kept
|
|
||||||
r = await client.get("/client-ip", headers={
|
|
||||||
'X-Forwarded-For': '203.0.113.5',
|
|
||||||
'X-Forwarded-Port': 'not-a-port',
|
|
||||||
})
|
|
||||||
self.assertEqual({'host': '203.0.113.5', 'port': socket_client['port']}, json.loads(r.text))
|
|
||||||
|
|
||||||
@async_test
|
@async_test
|
||||||
async def test_nested_param_routes(self):
|
async def test_nested_param_routes(self):
|
||||||
app = KayaApp()
|
app = KayaApp()
|
||||||
@@ -309,3 +250,32 @@ class AsgiTest(unittest.TestCase):
|
|||||||
r = await client.put("/restaurants/42/menu")
|
r = await client.put("/restaurants/42/menu")
|
||||||
self.assertEqual(404, r.status_code)
|
self.assertEqual(404, r.status_code)
|
||||||
|
|
||||||
|
def test_route_template(self):
|
||||||
|
self.assertEqual('/employee/${employee_id}',
|
||||||
|
self.app.route_template('/employee/101325', HttpMethod.GET))
|
||||||
|
self.assertEqual('/square/${x:int}', self.app.route_template('/square/30', HttpMethod.GET))
|
||||||
|
self.assertIsNone(self.app.route_template('/unknown', HttpMethod.GET))
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_exception_exposed_to_after_hooks(self):
|
||||||
|
app = KayaApp()
|
||||||
|
seen = []
|
||||||
|
|
||||||
|
async def after_request(ctx: HttpContext) -> None:
|
||||||
|
seen.append(ctx.exception)
|
||||||
|
|
||||||
|
app.add_after_request_hook(after_request)
|
||||||
|
|
||||||
|
@app.GET('/raises')
|
||||||
|
async def raises(ctx: HttpContext) -> None:
|
||||||
|
raise RuntimeError('boom')
|
||||||
|
|
||||||
|
transport = httpx.ASGITransport(app=app)
|
||||||
|
async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as client:
|
||||||
|
with self.assertRaises(RuntimeError):
|
||||||
|
await client.get('/raises')
|
||||||
|
|
||||||
|
self.assertEqual(1, len(seen))
|
||||||
|
self.assertIsInstance(seen[0], RuntimeError)
|
||||||
|
self.assertEqual('boom', str(seen[0]))
|
||||||
|
|
||||||
|
|||||||
@@ -80,6 +80,27 @@ class TreeTest(unittest.TestCase):
|
|||||||
self.assertIs(Maybe.of(handler_num).map(self.handlers.__getitem__).or_none(),
|
self.assertIs(Maybe.of(handler_num).map(self.handlers.__getitem__).or_none(),
|
||||||
Maybe.of_nullable(res).map(lambda it: it[0]).or_none())
|
Maybe.of_nullable(res).map(lambda it: it[0]).or_none())
|
||||||
|
|
||||||
|
def test_route_template(self):
|
||||||
|
cases: Tuple[Tuple[str, HttpMethod, Optional[str]], ...] = (
|
||||||
|
('/home/something', HttpMethod.GET, '/home/something'),
|
||||||
|
('/home/something_else', HttpMethod.POST, '/home/something_else'),
|
||||||
|
('/home/README.md', HttpMethod.GET, '/home/*.md'),
|
||||||
|
('/home/something/ciao/blah/README.md', HttpMethod.GET, '/home/something/*/blah/*.md'),
|
||||||
|
('/home/bar/ciao/blah/README.md', HttpMethod.GET, '/home/bar/*'),
|
||||||
|
('/unknown', HttpMethod.GET, None),
|
||||||
|
)
|
||||||
|
for url, method, expected in cases:
|
||||||
|
with self.subTest(f'{method} {url}'):
|
||||||
|
self.assertEqual(expected, self.tree.route_template(url, method))
|
||||||
|
|
||||||
|
def test_route_template_with_params(self):
|
||||||
|
tree = Tree()
|
||||||
|
tree.add((p for p in ('foo', '${id:int}')), HttpMethod.PUT, self.handlers[0])
|
||||||
|
tree.add((p for p in ('foo', '${name}')), HttpMethod.GET, self.handlers[1])
|
||||||
|
self.assertEqual('/foo/${id:int}', tree.route_template('/foo/42', HttpMethod.PUT))
|
||||||
|
self.assertEqual('/foo/${name}', tree.route_template('/foo/bar', HttpMethod.GET))
|
||||||
|
self.assertIsNone(tree.route_template('/foo/bar', HttpMethod.DELETE))
|
||||||
|
|
||||||
def test_two_method_agnostic_matchers_raise(self):
|
def test_two_method_agnostic_matchers_raise(self):
|
||||||
tree = Tree()
|
tree = Tree()
|
||||||
tree.add((p for p in ('foo', '*')), None, self.handlers[0])
|
tree.add((p for p in ('foo', '*')), None, self.handlers[0])
|
||||||
|
|||||||
@@ -132,37 +132,6 @@ class WebSocketTest(unittest.TestCase):
|
|||||||
self.assertEqual(1, len(sent_messages))
|
self.assertEqual(1, len(sent_messages))
|
||||||
self.assertEqual({'type': 'websocket.accept'}, sent_messages[0])
|
self.assertEqual({'type': 'websocket.accept'}, sent_messages[0])
|
||||||
|
|
||||||
@async_test
|
|
||||||
async def test_client_forwarded_header(self):
|
|
||||||
async def send(message):
|
|
||||||
pass
|
|
||||||
|
|
||||||
async def receive():
|
|
||||||
return {'type': 'websocket.connect'}
|
|
||||||
|
|
||||||
scope = {
|
|
||||||
'type': 'websocket',
|
|
||||||
'path': '/echo',
|
|
||||||
'query_string': b'',
|
|
||||||
'scheme': 'ws',
|
|
||||||
'client': ('127.0.0.1', 12345),
|
|
||||||
'server': ('127.0.0.1', 80),
|
|
||||||
'headers': [(b'forwarded', b'for=203.0.113.5:1234')],
|
|
||||||
}
|
|
||||||
ws = AsgiWebSocket(scope, receive, send)
|
|
||||||
self.assertEqual(('203.0.113.5', 1234), ws.client)
|
|
||||||
|
|
||||||
scope['headers'] = [
|
|
||||||
(b'x-forwarded-for', b'203.0.113.5, 70.41.3.18'),
|
|
||||||
(b'x-forwarded-port', b'8443'),
|
|
||||||
]
|
|
||||||
ws = AsgiWebSocket(scope, receive, send)
|
|
||||||
self.assertEqual(('203.0.113.5', 8443), ws.client)
|
|
||||||
|
|
||||||
scope['headers'] = []
|
|
||||||
ws = AsgiWebSocket(scope, receive, send)
|
|
||||||
self.assertEqual(('127.0.0.1', 12345), ws.client)
|
|
||||||
|
|
||||||
@async_test
|
@async_test
|
||||||
async def test_websocket_scope_without_scheme(self):
|
async def test_websocket_scope_without_scheme(self):
|
||||||
# Daphne omits the optional `scheme` key from websocket scopes.
|
# Daphne omits the optional `scheme` key from websocket scopes.
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
# kaya-forwarded
|
||||||
|
|
||||||
|
Trusted-proxy forwarded header handling for the Kaya web framework.
|
||||||
|
|
||||||
|
Without this package, Kaya exposes the raw socket peer address as
|
||||||
|
`ctx.client` / `ws.client` and ignores `Forwarded` / `X-Forwarded-*` headers
|
||||||
|
entirely (they are client-controllable and trivially spoofable when the app is
|
||||||
|
directly exposed).
|
||||||
|
|
||||||
|
`ForwardedHeadersMixin` opts the application into honoring those headers, but
|
||||||
|
only when the direct socket peer is a trusted proxy, identified by a list of
|
||||||
|
trusted CIDRs/IPs.
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
|
||||||
|
```python
|
||||||
|
from kaya.core import KayaApp, HttpContext
|
||||||
|
from kaya.forwarded import ForwardedHeadersMixin
|
||||||
|
|
||||||
|
app = KayaApp(mixins=[
|
||||||
|
ForwardedHeadersMixin(trusted_proxies=['127.0.0.1', '::1', '10.0.0.0/8'])
|
||||||
|
])
|
||||||
|
|
||||||
|
@app.GET('/whoami')
|
||||||
|
async def whoami(ctx: HttpContext):
|
||||||
|
host, port = ctx.client
|
||||||
|
await ctx.send_str(200, f'{host}:{port}')
|
||||||
|
```
|
||||||
|
|
||||||
|
## How it works
|
||||||
|
|
||||||
|
When a request arrives:
|
||||||
|
|
||||||
|
1. If the socket peer IP does not belong to any trusted CIDR (or there is no
|
||||||
|
peer address), the mixin leaves the context untouched — `client` remains
|
||||||
|
the socket peer and all proxy headers are ignored.
|
||||||
|
2. Otherwise the client address is resolved from the headers, in order:
|
||||||
|
- `Forwarded` (RFC 7239): the `for=` entries are walked **from right to
|
||||||
|
left**, skipping entries that are themselves trusted proxies (and
|
||||||
|
`unknown`); the first untrusted entry is the client. This defeats
|
||||||
|
spoofing when the edge proxy *appends* to the header (e.g. nginx with
|
||||||
|
`$proxy_add_x_forwarded_for`), because attacker-supplied leftmost entries
|
||||||
|
are never selected. A `:port` in the selected `for=` value also
|
||||||
|
populates the port.
|
||||||
|
- `X-Forwarded-For`: same right-to-left trusted-proxy walk; the port comes
|
||||||
|
from `X-Forwarded-Port` when present and valid.
|
||||||
|
- `X-Forwarded-Host`: first entry; port from `X-Forwarded-Port` as above.
|
||||||
|
3. If none of the headers are present or usable, the socket peer is kept.
|
||||||
|
|
||||||
|
If every entry in the chain is a trusted proxy, the leftmost entry is used
|
||||||
|
(the whole chain is trusted, so the leftmost is the original client).
|
||||||
|
|
||||||
|
The resolved address is exposed by wrapping the request context /
|
||||||
|
websocket (the same pattern as `kaya-session`), so both ASGI and RSGI keep
|
||||||
|
working and `ctx.session` from other mixins is preserved.
|
||||||
|
|
||||||
|
## Note
|
||||||
|
|
||||||
|
Even with this mixin, the edge proxy should still strip or overwrite inbound
|
||||||
|
`Forwarded` / `X-Forwarded-*` headers from clients — the mixin protects the
|
||||||
|
application, the proxy protects the chain.
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
[build-system]
|
||||||
|
requires = ["setuptools>=61.0", "setuptools-scm>=8"]
|
||||||
|
build-backend = "setuptools.build_meta"
|
||||||
|
|
||||||
|
[project]
|
||||||
|
name = "kaya-forwarded"
|
||||||
|
dynamic = ["version"]
|
||||||
|
authors = [
|
||||||
|
{ name="Walter Oggioni", email="oggioni.walter@gmail.com" },
|
||||||
|
]
|
||||||
|
description = "Trusted-proxy forwarded header handling for the Kaya lightweight ASGI web framework"
|
||||||
|
readme = "README.md"
|
||||||
|
requires-python = ">=3.10"
|
||||||
|
license = "MIT"
|
||||||
|
classifiers = [
|
||||||
|
'Development Status :: 3 - Alpha',
|
||||||
|
'Topic :: Utilities',
|
||||||
|
'Intended Audience :: System Administrators',
|
||||||
|
'Intended Audience :: Developers',
|
||||||
|
'Environment :: Console',
|
||||||
|
'Programming Language :: Python :: 3',
|
||||||
|
]
|
||||||
|
|
||||||
|
dependencies = [
|
||||||
|
"kaya-core",
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
dev = [
|
||||||
|
"build", "mypy", "ipdb", "twine", "httpx", "httpx-ws", "kaya-rsgi"
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.urls]
|
||||||
|
"Homepage" = "https://github.com/woggioni/kaya"
|
||||||
|
"Bug Tracker" = "https://github.com/woggioni/kaya/issues"
|
||||||
|
|
||||||
|
[tool.setuptools.packages.find]
|
||||||
|
where = ["src"]
|
||||||
|
namespaces = true
|
||||||
|
|
||||||
|
[tool.mypy]
|
||||||
|
python_version = "3.12"
|
||||||
|
disallow_untyped_defs = true
|
||||||
|
show_error_codes = true
|
||||||
|
no_implicit_optional = true
|
||||||
|
warn_return_any = true
|
||||||
|
warn_unused_ignores = true
|
||||||
|
exclude = ["scripts", "docs", "test"]
|
||||||
|
strict = true
|
||||||
|
|
||||||
|
[tool.setuptools_scm]
|
||||||
|
root = "../.."
|
||||||
|
version_file = "src/kaya/forwarded/_version.py"
|
||||||
|
|
||||||
|
[tool.setuptools_scm.tag]
|
||||||
|
prefix = "release/"
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
from pkgutil import extend_path
|
||||||
|
|
||||||
|
__path__ = extend_path(__path__, __name__)
|
||||||
|
|
||||||
|
from ._mixin import ForwardedHeadersMixin
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
'ForwardedHeadersMixin',
|
||||||
|
]
|
||||||
@@ -0,0 +1,237 @@
|
|||||||
|
from ipaddress import ip_address, ip_network, IPv4Address, IPv4Network, IPv6Address, IPv6Network
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import (
|
||||||
|
Any,
|
||||||
|
AsyncGenerator,
|
||||||
|
List,
|
||||||
|
Mapping,
|
||||||
|
Optional,
|
||||||
|
Sequence,
|
||||||
|
Tuple,
|
||||||
|
Union,
|
||||||
|
)
|
||||||
|
|
||||||
|
from kaya.core import HttpContext, KayaApp, KayaMixin, WebSocket, WebSocketMessage
|
||||||
|
from kaya.core._types import StrOrStrings
|
||||||
|
|
||||||
|
_IPAddress = Union[IPv4Address, IPv6Address]
|
||||||
|
_IPNetwork = Union[IPv4Network, IPv6Network]
|
||||||
|
|
||||||
|
|
||||||
|
def _split_host_port(value: str) -> Tuple[str, Optional[int]]:
|
||||||
|
if value.startswith('['):
|
||||||
|
# bracketed IPv6 address, optionally followed by :port
|
||||||
|
closing = value.find(']')
|
||||||
|
if closing == -1:
|
||||||
|
return value, None
|
||||||
|
host = value[1:closing]
|
||||||
|
rest = value[closing + 1:]
|
||||||
|
if rest.startswith(':'):
|
||||||
|
try:
|
||||||
|
return host, int(rest[1:])
|
||||||
|
except ValueError:
|
||||||
|
return host, None
|
||||||
|
return host, None
|
||||||
|
if value.count(':') == 1:
|
||||||
|
host, _, port_str = value.rpartition(':')
|
||||||
|
try:
|
||||||
|
return host, int(port_str)
|
||||||
|
except ValueError:
|
||||||
|
return value, None
|
||||||
|
return value, None
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_forwarded_entries(headers: Mapping[str, Sequence[str]]) -> List[Tuple[str, Optional[int]]]:
|
||||||
|
"""Extract the (host, port) `for=` entries of the RFC 7239 `Forwarded` header,
|
||||||
|
flattened across all header values, in chain order (leftmost = original client).
|
||||||
|
"""
|
||||||
|
entries: List[Tuple[str, Optional[int]]] = []
|
||||||
|
for raw_value in headers.get('forwarded', ()):
|
||||||
|
for element in raw_value.split(','):
|
||||||
|
for param in element.split(';'):
|
||||||
|
key, sep, value = param.partition('=')
|
||||||
|
if sep and key.strip().lower() == 'for':
|
||||||
|
for_value = value.strip().strip('"')
|
||||||
|
if for_value and for_value.lower() != 'unknown':
|
||||||
|
entries.append(_split_host_port(for_value))
|
||||||
|
break
|
||||||
|
return entries
|
||||||
|
|
||||||
|
|
||||||
|
def _first_header_value(headers: Mapping[str, Sequence[str]], name: str) -> Optional[str]:
|
||||||
|
values = headers.get(name)
|
||||||
|
if not values:
|
||||||
|
return None
|
||||||
|
first = values[0].split(',')[0].strip()
|
||||||
|
return first or None
|
||||||
|
|
||||||
|
|
||||||
|
class _ForwardedHttpContext(HttpContext):
|
||||||
|
"""HttpContext wrapper that exposes the forwarded client address.
|
||||||
|
|
||||||
|
Everything except ``client`` is delegated to the wrapped context via
|
||||||
|
``__getattr__``, so it works with any concrete ``HttpContext`` (ASGI or
|
||||||
|
RSGI) and preserves attributes set by other mixins (e.g. ``session``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, ctx: HttpContext, client: Tuple[str, int]) -> None:
|
||||||
|
object.__setattr__(self, '_ctx', ctx)
|
||||||
|
object.__setattr__(self, 'session', ctx.session)
|
||||||
|
object.__setattr__(self, 'client', client)
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> Any:
|
||||||
|
if name == '_ctx':
|
||||||
|
raise AttributeError(name)
|
||||||
|
return getattr(self._ctx, name)
|
||||||
|
|
||||||
|
async def stream_body(self,
|
||||||
|
status: int,
|
||||||
|
body_generator: AsyncGenerator[bytes, None],
|
||||||
|
headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
await self._ctx.stream_body(status, body_generator, headers)
|
||||||
|
|
||||||
|
async def send_bytes(self, status: int, body: bytes, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
await self._ctx.send_bytes(status, body, headers)
|
||||||
|
|
||||||
|
async def send_str(self, status: int, body: str, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
await self._ctx.send_str(status, body, headers)
|
||||||
|
|
||||||
|
async def send_file(self, status: int, path: Path, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
await self._ctx.send_file(status, path, headers)
|
||||||
|
|
||||||
|
async def send_empty(self, status: int, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
await self._ctx.send_empty(status, headers)
|
||||||
|
|
||||||
|
|
||||||
|
class _ForwardedWebSocket(WebSocket):
|
||||||
|
"""WebSocket wrapper that exposes the forwarded client address.
|
||||||
|
|
||||||
|
Everything except ``client`` is delegated to the wrapped socket via
|
||||||
|
``__getattr__``, so it works with any concrete ``WebSocket`` (ASGI or
|
||||||
|
RSGI) and preserves attributes set by other mixins (e.g. ``session``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, ws: WebSocket, client: Tuple[str, int]) -> None:
|
||||||
|
object.__setattr__(self, '_ws', ws)
|
||||||
|
object.__setattr__(self, 'session', ws.session)
|
||||||
|
object.__setattr__(self, 'client', client)
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> Any:
|
||||||
|
if name == '_ws':
|
||||||
|
raise AttributeError(name)
|
||||||
|
return getattr(self._ws, name)
|
||||||
|
|
||||||
|
async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
await self._ws.accept(headers)
|
||||||
|
|
||||||
|
async def receive(self) -> WebSocketMessage:
|
||||||
|
return await self._ws.receive()
|
||||||
|
|
||||||
|
async def send_text(self, data: str) -> None:
|
||||||
|
await self._ws.send_text(data)
|
||||||
|
|
||||||
|
async def send_bytes(self, data: bytes) -> None:
|
||||||
|
await self._ws.send_bytes(data)
|
||||||
|
|
||||||
|
async def close(self, code: int = 1000) -> None:
|
||||||
|
await self._ws.close(code)
|
||||||
|
|
||||||
|
async def __anext__(self) -> WebSocketMessage:
|
||||||
|
return await self._ws.__anext__()
|
||||||
|
|
||||||
|
|
||||||
|
class ForwardedHeadersMixin(KayaMixin):
|
||||||
|
"""Kaya mixin that resolves the client address from proxy headers
|
||||||
|
(``Forwarded``, ``X-Forwarded-For``, ``X-Forwarded-Host``), but only when
|
||||||
|
the direct socket peer is a trusted proxy.
|
||||||
|
|
||||||
|
Without this mixin, Kaya exposes the raw socket peer as ``ctx.client`` /
|
||||||
|
``ws.client`` and ignores forwarded headers entirely. With the mixin
|
||||||
|
applied, forwarded headers are honored only if the socket peer IP belongs
|
||||||
|
to one of the ``trusted_proxies`` CIDRs; otherwise the context is left
|
||||||
|
untouched.
|
||||||
|
|
||||||
|
When the peer is trusted, the address chain is walked from right to left
|
||||||
|
and entries that are themselves trusted proxies are skipped, so a client
|
||||||
|
that prepends a spoofed entry cannot fool the resolution when the edge
|
||||||
|
proxy appends to the header (e.g. nginx with ``$proxy_add_x_forwarded_for``).
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
app = KayaApp(mixins=[
|
||||||
|
ForwardedHeadersMixin(trusted_proxies=['127.0.0.1', '::1', '10.0.0.0/8'])
|
||||||
|
])
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, trusted_proxies: Sequence[str] = ()) -> None:
|
||||||
|
self._trusted_networks: Tuple[_IPNetwork, ...] = tuple(
|
||||||
|
ip_network(cidr, strict=False) for cidr in trusted_proxies
|
||||||
|
)
|
||||||
|
|
||||||
|
def apply(self, app: KayaApp) -> None:
|
||||||
|
app.add_before_request_hook(self._before_request)
|
||||||
|
app.add_before_websocket_hook(self._before_websocket)
|
||||||
|
|
||||||
|
def _is_trusted(self, host: str) -> bool:
|
||||||
|
try:
|
||||||
|
addr: _IPAddress = ip_address(host)
|
||||||
|
except ValueError:
|
||||||
|
return False
|
||||||
|
return any(addr.version == network.version and addr in network
|
||||||
|
for network in self._trusted_networks)
|
||||||
|
|
||||||
|
def _select_forwarded_entry(self, entries: List[Tuple[str, Optional[int]]]) -> Optional[Tuple[str, Optional[int]]]:
|
||||||
|
"""Walk the chain right-to-left skipping trusted proxies; the first
|
||||||
|
untrusted entry is the client. If every entry is trusted, the leftmost
|
||||||
|
(original client) is returned.
|
||||||
|
"""
|
||||||
|
for entry in reversed(entries):
|
||||||
|
if not self._is_trusted(entry[0]):
|
||||||
|
return entry
|
||||||
|
return entries[0] if entries else None
|
||||||
|
|
||||||
|
def _forwarded_port(self, headers: Mapping[str, Sequence[str]], socket_port: int) -> int:
|
||||||
|
forwarded_port = _first_header_value(headers, 'x-forwarded-port')
|
||||||
|
if forwarded_port is not None:
|
||||||
|
try:
|
||||||
|
return int(forwarded_port)
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
|
return socket_port
|
||||||
|
|
||||||
|
def _resolve(self,
|
||||||
|
headers: Mapping[str, Sequence[str]],
|
||||||
|
client: Optional[Tuple[str, int]]) -> Optional[Tuple[str, int]]:
|
||||||
|
if client is None or not self._is_trusted(client[0]):
|
||||||
|
return client
|
||||||
|
socket_port = client[1]
|
||||||
|
|
||||||
|
selected = self._select_forwarded_entry(_parse_forwarded_entries(headers))
|
||||||
|
if selected is not None:
|
||||||
|
host, port = selected
|
||||||
|
return host, port if port is not None else socket_port
|
||||||
|
|
||||||
|
xff_values = headers.get('x-forwarded-for')
|
||||||
|
if xff_values:
|
||||||
|
xff_entries = [entry.strip() for value in xff_values for entry in value.split(',') if entry.strip()]
|
||||||
|
selected_host = self._select_forwarded_entry([(entry, None) for entry in xff_entries])
|
||||||
|
if selected_host is not None:
|
||||||
|
return selected_host[0], self._forwarded_port(headers, socket_port)
|
||||||
|
|
||||||
|
xfh = _first_header_value(headers, 'x-forwarded-host')
|
||||||
|
if xfh is not None:
|
||||||
|
return xfh, self._forwarded_port(headers, socket_port)
|
||||||
|
|
||||||
|
return client
|
||||||
|
|
||||||
|
async def _before_request(self, ctx: HttpContext) -> Optional[HttpContext]:
|
||||||
|
resolved = self._resolve(ctx.headers, ctx.client)
|
||||||
|
if resolved is None or resolved is ctx.client:
|
||||||
|
return None
|
||||||
|
return _ForwardedHttpContext(ctx, resolved)
|
||||||
|
|
||||||
|
async def _before_websocket(self, ws: WebSocket) -> Optional[WebSocket]:
|
||||||
|
resolved = self._resolve(ws.headers, ws.client)
|
||||||
|
if resolved is None or resolved is ws.client:
|
||||||
|
return None
|
||||||
|
return _ForwardedWebSocket(ws, resolved)
|
||||||
@@ -0,0 +1,244 @@
|
|||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import unittest
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
from pwo import async_test
|
||||||
|
|
||||||
|
from kaya.core import HttpContext, KayaApp
|
||||||
|
from kaya.core._asgi import AsgiWebSocket
|
||||||
|
from kaya.forwarded import ForwardedHeadersMixin
|
||||||
|
|
||||||
|
TRUSTED = ['127.0.0.1', '::1', '10.0.0.0/8']
|
||||||
|
|
||||||
|
|
||||||
|
def make_app(trusted_proxies=TRUSTED) -> KayaApp:
|
||||||
|
mixins = [ForwardedHeadersMixin(trusted_proxies=trusted_proxies)] if trusted_proxies is not None else []
|
||||||
|
app = KayaApp(mixins=mixins)
|
||||||
|
|
||||||
|
@app.GET('/client')
|
||||||
|
async def client(ctx: HttpContext) -> None:
|
||||||
|
host, port = ctx.client if ctx.client is not None else (None, None)
|
||||||
|
await ctx.send_str(200, json.dumps({'host': host, 'port': port}))
|
||||||
|
|
||||||
|
return app
|
||||||
|
|
||||||
|
|
||||||
|
async def request(app: KayaApp,
|
||||||
|
headers: Optional[dict[str, str]] = None,
|
||||||
|
client: tuple[str, int] = ('127.0.0.1', 123)) -> dict:
|
||||||
|
transport = httpx.ASGITransport(app=app, client=client)
|
||||||
|
async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as http_client:
|
||||||
|
r = await http_client.get('/client', headers=headers)
|
||||||
|
assert r.status_code == 200
|
||||||
|
return json.loads(r.text)
|
||||||
|
|
||||||
|
|
||||||
|
class TrustedPeerTest(unittest.TestCase):
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_no_headers_returns_socket_peer(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app)
|
||||||
|
self.assertEqual({'host': '127.0.0.1', 'port': 123}, result)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_forwarded_header_with_port(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={'Forwarded': 'for=203.0.113.5:1234'})
|
||||||
|
self.assertEqual({'host': '203.0.113.5', 'port': 1234}, result)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_forwarded_header_bracketed_ipv6(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={'Forwarded': 'for="[2001:db8::1]:4711"'})
|
||||||
|
self.assertEqual({'host': '2001:db8::1', 'port': 4711}, result)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_forwarded_header_without_port_keeps_socket_port(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={'Forwarded': 'for=203.0.113.5'})
|
||||||
|
self.assertEqual({'host': '203.0.113.5', 'port': 123}, result)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_forwarded_unknown_entry_skipped(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={'Forwarded': 'for=unknown, for=203.0.113.5'})
|
||||||
|
self.assertEqual('203.0.113.5', result['host'])
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_forwarded_rightmost_untrusted_wins(self):
|
||||||
|
# attacker-controlled leftmost entry is skipped: the rightmost
|
||||||
|
# untrusted entry (appended by the trusted edge proxy) is the client
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={'Forwarded': 'for=1.2.3.4, for=5.6.7.8, for=10.0.0.2'})
|
||||||
|
self.assertEqual('5.6.7.8', result['host'])
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_x_forwarded_for_spoofed_leftmost_entry_skipped(self):
|
||||||
|
# XFF = "<attacker-supplied>, <real client>" as appended by the proxy
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={'X-Forwarded-For': '1.2.3.4, 5.6.7.8'})
|
||||||
|
self.assertEqual('5.6.7.8', result['host'])
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_x_forwarded_for_all_trusted_chain_uses_leftmost(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={'X-Forwarded-For': '10.0.0.5, 10.0.0.2'})
|
||||||
|
self.assertEqual('10.0.0.5', result['host'])
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_x_forwarded_port(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={
|
||||||
|
'X-Forwarded-For': '203.0.113.5',
|
||||||
|
'X-Forwarded-Port': '8443',
|
||||||
|
})
|
||||||
|
self.assertEqual({'host': '203.0.113.5', 'port': 8443}, result)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_invalid_x_forwarded_port_ignored(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={
|
||||||
|
'X-Forwarded-For': '203.0.113.5',
|
||||||
|
'X-Forwarded-Port': 'not-a-port',
|
||||||
|
})
|
||||||
|
self.assertEqual({'host': '203.0.113.5', 'port': 123}, result)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_forwarded_takes_precedence_over_x_forwarded_for(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={
|
||||||
|
'Forwarded': 'for=203.0.113.5',
|
||||||
|
'X-Forwarded-For': '198.51.100.7',
|
||||||
|
})
|
||||||
|
self.assertEqual('203.0.113.5', result['host'])
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_x_forwarded_host_fallback(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app, headers={'X-Forwarded-Host': '198.51.100.7'})
|
||||||
|
self.assertEqual('198.51.100.7', result['host'])
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_ipv6_cidr_trust(self):
|
||||||
|
app = make_app(trusted_proxies=['2001:db8::/32'])
|
||||||
|
result = await request(app,
|
||||||
|
headers={'X-Forwarded-For': '203.0.113.5'},
|
||||||
|
client=('2001:db8::10', 9999))
|
||||||
|
self.assertEqual({'host': '203.0.113.5', 'port': 9999}, result)
|
||||||
|
|
||||||
|
|
||||||
|
class UntrustedPeerTest(unittest.TestCase):
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_untrusted_peer_ignores_forwarded_headers(self):
|
||||||
|
app = make_app()
|
||||||
|
result = await request(app,
|
||||||
|
headers={'Forwarded': 'for=1.2.3.4', 'X-Forwarded-For': '1.2.3.4'},
|
||||||
|
client=('203.0.113.99', 4567))
|
||||||
|
self.assertEqual({'host': '203.0.113.99', 'port': 4567}, result)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_empty_trusted_proxies_ignores_everything(self):
|
||||||
|
app = make_app(trusted_proxies=[])
|
||||||
|
result = await request(app, headers={'X-Forwarded-For': '1.2.3.4'})
|
||||||
|
self.assertEqual({'host': '127.0.0.1', 'port': 123}, result)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_invalid_cidr_fails_fast(self):
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
ForwardedHeadersMixin(trusted_proxies=['not-a-cidr'])
|
||||||
|
|
||||||
|
|
||||||
|
class OptOutTest(unittest.TestCase):
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_without_mixin_headers_are_ignored(self):
|
||||||
|
app = KayaApp()
|
||||||
|
|
||||||
|
@app.GET('/client')
|
||||||
|
async def client(ctx: HttpContext) -> None:
|
||||||
|
host, port = ctx.client if ctx.client is not None else (None, None)
|
||||||
|
await ctx.send_str(200, json.dumps({'host': host, 'port': port}))
|
||||||
|
|
||||||
|
result = await request(app, headers={'Forwarded': 'for=1.2.3.4', 'X-Forwarded-For': '1.2.3.4'})
|
||||||
|
self.assertEqual({'host': '127.0.0.1', 'port': 123}, result)
|
||||||
|
|
||||||
|
|
||||||
|
class WebSocketTest(unittest.TestCase):
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _make_ws(headers):
|
||||||
|
async def send(message):
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def receive():
|
||||||
|
return {'type': 'websocket.connect'}
|
||||||
|
|
||||||
|
scope = {
|
||||||
|
'type': 'websocket',
|
||||||
|
'path': '/ws',
|
||||||
|
'query_string': b'',
|
||||||
|
'scheme': 'ws',
|
||||||
|
'client': ('127.0.0.1', 12345),
|
||||||
|
'server': ('127.0.0.1', 80),
|
||||||
|
'headers': headers,
|
||||||
|
}
|
||||||
|
return AsgiWebSocket(scope, receive, send)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_websocket_trusted_peer(self):
|
||||||
|
mixin = ForwardedHeadersMixin(trusted_proxies=TRUSTED)
|
||||||
|
ws = self._make_ws([(b'x-forwarded-for', b'1.2.3.4, 5.6.7.8')])
|
||||||
|
wrapped = await mixin._before_websocket(ws)
|
||||||
|
assert wrapped is not None
|
||||||
|
self.assertEqual(('5.6.7.8', 12345), wrapped.client)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_websocket_untrusted_peer(self):
|
||||||
|
mixin = ForwardedHeadersMixin(trusted_proxies=['10.0.0.0/8'])
|
||||||
|
ws = self._make_ws([(b'x-forwarded-for', b'1.2.3.4')])
|
||||||
|
wrapped = await mixin._before_websocket(ws)
|
||||||
|
self.assertIsNone(wrapped)
|
||||||
|
self.assertEqual(('127.0.0.1', 12345), ws.client)
|
||||||
|
|
||||||
|
|
||||||
|
class RsgiTest(unittest.TestCase):
|
||||||
|
|
||||||
|
def test_rsgi_context(self):
|
||||||
|
from kaya.rsgi import RsgiContext
|
||||||
|
|
||||||
|
class FakeScope:
|
||||||
|
scheme = 'http'
|
||||||
|
method = 'GET'
|
||||||
|
path = '/'
|
||||||
|
query_string = ''
|
||||||
|
headers = {'x-forwarded-for': '1.2.3.4, 5.6.7.8'}
|
||||||
|
client = '127.0.0.1:12345'
|
||||||
|
server = '127.0.0.1:80'
|
||||||
|
|
||||||
|
mixin = ForwardedHeadersMixin(trusted_proxies=TRUSTED)
|
||||||
|
ctx = RsgiContext(FakeScope(), object()) # type: ignore[arg-type]
|
||||||
|
wrapped = asyncio.run(mixin._before_request(ctx))
|
||||||
|
assert wrapped is not None
|
||||||
|
self.assertEqual(('5.6.7.8', 12345), wrapped.client)
|
||||||
|
|
||||||
|
def test_rsgi_context_untrusted_peer(self):
|
||||||
|
from kaya.rsgi import RsgiContext
|
||||||
|
|
||||||
|
class FakeScope:
|
||||||
|
scheme = 'http'
|
||||||
|
method = 'GET'
|
||||||
|
path = '/'
|
||||||
|
query_string = ''
|
||||||
|
headers = {'x-forwarded-for': '1.2.3.4'}
|
||||||
|
client = '192.0.2.1:12345'
|
||||||
|
server = '127.0.0.1:80'
|
||||||
|
|
||||||
|
mixin = ForwardedHeadersMixin(trusted_proxies=TRUSTED)
|
||||||
|
ctx = RsgiContext(FakeScope(), object()) # type: ignore[arg-type]
|
||||||
|
wrapped = asyncio.run(mixin._before_request(ctx))
|
||||||
|
self.assertIsNone(wrapped)
|
||||||
|
self.assertEqual(('192.0.2.1', 12345), ctx.client)
|
||||||
@@ -0,0 +1,85 @@
|
|||||||
|
# kaya-otel
|
||||||
|
|
||||||
|
OpenTelemetry tracing and metrics for the Kaya web framework.
|
||||||
|
|
||||||
|
Provides `OTelMixin`, a `KayaMixin` that instruments HTTP requests and
|
||||||
|
WebSocket connections with OpenTelemetry spans and exports HTTP server
|
||||||
|
metrics, shipping everything to an OTLP/HTTP collector.
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
|
||||||
|
```python
|
||||||
|
from kaya.core import KayaApp, HttpContext
|
||||||
|
from kaya.otel import OTelMixin
|
||||||
|
|
||||||
|
app = KayaApp(mixins=[
|
||||||
|
OTelMixin(
|
||||||
|
service_name='my-service',
|
||||||
|
endpoint='http://localhost:4318',
|
||||||
|
)
|
||||||
|
])
|
||||||
|
|
||||||
|
@app.GET('/')
|
||||||
|
async def home(ctx: HttpContext):
|
||||||
|
await ctx.send_str(200, 'Hello World!')
|
||||||
|
```
|
||||||
|
|
||||||
|
## Parameters
|
||||||
|
|
||||||
|
- `service_name`: value of the `service.name` resource attribute.
|
||||||
|
- `endpoint`: base URL of an OTLP/HTTP collector; the signal paths
|
||||||
|
`/v1/traces` and `/v1/metrics` are appended. When `None`, the exporters
|
||||||
|
use their own defaults, including the standard
|
||||||
|
`OTEL_EXPORTER_OTLP_ENDPOINT` environment variable.
|
||||||
|
- `headers`: extra HTTP headers sent to the collector (e.g. authentication).
|
||||||
|
- `metric_export_interval_millis`: metric export interval
|
||||||
|
(default `60000`).
|
||||||
|
- `tracer_provider` / `meter_provider`: inject custom providers (e.g. with
|
||||||
|
in-memory exporters for tests) instead of the OTLP defaults. The mixin
|
||||||
|
only shuts down providers it created itself.
|
||||||
|
- `resource_attributes`: extra resource attributes merged with
|
||||||
|
`service.name` when the mixin creates the providers.
|
||||||
|
- `excluded_paths` / `excluded_path_regexes`: skip tracing and metrics for
|
||||||
|
matching paths (useful for health checks and metrics endpoints).
|
||||||
|
- `capture_request_headers` / `capture_response_headers`: opt-in HTTP header
|
||||||
|
capture as `http.request.header.<name>` / `http.response.header.<name>`
|
||||||
|
span attributes. Header names are normalized to lowercase with `-`
|
||||||
|
replaced by `_`; values are captured as string lists.
|
||||||
|
- `sanitize_headers`: headers captured as `REDACTED` (default:
|
||||||
|
`authorization`, `proxy-authorization`, `cookie`, `set-cookie`).
|
||||||
|
- `server_request_hook` / `server_response_hook` / `websocket_connect_hook` /
|
||||||
|
`websocket_close_hook`: optional synchronous callbacks invoked with the
|
||||||
|
span and the Kaya context/websocket at the corresponding lifecycle point.
|
||||||
|
Hook exceptions are logged and do not fail the request.
|
||||||
|
- `metrics_include_raw_path`: when `True`, metrics use raw `url.path`
|
||||||
|
attributes (legacy, potentially high cardinality). The default `False`
|
||||||
|
keeps metrics low-cardinality by using `http.route` when Kaya can resolve
|
||||||
|
a route template.
|
||||||
|
- `websocket_error_close_codes`: close codes that mark a websocket span as
|
||||||
|
failed. Defaults to protocol/application error codes such as `1002`,
|
||||||
|
`1003` and `1007`-`1011`.
|
||||||
|
|
||||||
|
## Behavior
|
||||||
|
|
||||||
|
- Every HTTP request gets a `SERVER` span named `<METHOD> <path>` with the
|
||||||
|
usual HTTP semantic attributes (`http.request.method`, `url.path`,
|
||||||
|
`client.address`, `server.address`, ...). When Kaya resolves a route
|
||||||
|
template, the span name is updated to `<METHOD> <route template>` and
|
||||||
|
`http.route` is set. The response status code is recorded as
|
||||||
|
`http.response.status_code` when the handler sends the response; 5xx
|
||||||
|
statuses mark the span as failed. Exceptions escaping the handler are
|
||||||
|
recorded on the span and also mark it as failed.
|
||||||
|
- Every WebSocket connection gets one span for its whole lifetime. The close
|
||||||
|
code is recorded as `kaya.websocket.close_code`; exceptions and configured
|
||||||
|
error close codes mark the span as failed.
|
||||||
|
- W3C `traceparent`/`tracestate` headers on incoming requests are honored,
|
||||||
|
so traces propagate from upstream services.
|
||||||
|
- Metrics:
|
||||||
|
- `http.server.request.duration` (histogram, seconds), with
|
||||||
|
`http.request.method`, `http.route` when known and
|
||||||
|
`http.response.status_code` attributes.
|
||||||
|
- `http.server.active_requests` (up-down counter), with
|
||||||
|
`http.request.method` attributes.
|
||||||
|
|
||||||
|
`OTelMixin` is a `KayaMixin`, so the app stays a `KayaApp` and both ASGI and
|
||||||
|
RSGI keep working.
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
[build-system]
|
||||||
|
requires = ["setuptools>=61.0", "setuptools-scm>=8"]
|
||||||
|
build-backend = "setuptools.build_meta"
|
||||||
|
|
||||||
|
[project]
|
||||||
|
name = "kaya-otel"
|
||||||
|
dynamic = ["version"]
|
||||||
|
authors = [
|
||||||
|
{ name="Walter Oggioni", email="oggioni.walter@gmail.com" },
|
||||||
|
]
|
||||||
|
description = "OpenTelemetry tracing and metrics for the Kaya lightweight ASGI web framework"
|
||||||
|
readme = "README.md"
|
||||||
|
requires-python = ">=3.10"
|
||||||
|
license = "MIT"
|
||||||
|
classifiers = [
|
||||||
|
'Development Status :: 3 - Alpha',
|
||||||
|
'Topic :: Utilities',
|
||||||
|
'Intended Audience :: System Administrators',
|
||||||
|
'Intended Audience :: Developers',
|
||||||
|
'Environment :: Console',
|
||||||
|
'Programming Language :: Python :: 3',
|
||||||
|
]
|
||||||
|
|
||||||
|
dependencies = [
|
||||||
|
"kaya-core",
|
||||||
|
"opentelemetry-sdk",
|
||||||
|
"opentelemetry-exporter-otlp-proto-http",
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
dev = [
|
||||||
|
"build", "mypy", "ipdb", "twine", "httpx", "httpx-ws", "kaya-rsgi"
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.urls]
|
||||||
|
"Homepage" = "https://github.com/woggioni/kaya"
|
||||||
|
"Bug Tracker" = "https://github.com/woggioni/kaya/issues"
|
||||||
|
|
||||||
|
[tool.setuptools.packages.find]
|
||||||
|
where = ["src"]
|
||||||
|
namespaces = true
|
||||||
|
|
||||||
|
[tool.mypy]
|
||||||
|
python_version = "3.12"
|
||||||
|
disallow_untyped_defs = true
|
||||||
|
show_error_codes = true
|
||||||
|
no_implicit_optional = true
|
||||||
|
warn_return_any = true
|
||||||
|
warn_unused_ignores = true
|
||||||
|
exclude = ["scripts", "docs", "test"]
|
||||||
|
strict = true
|
||||||
|
|
||||||
|
[tool.setuptools_scm]
|
||||||
|
root = "../.."
|
||||||
|
version_file = "src/kaya/otel/_version.py"
|
||||||
|
|
||||||
|
[tool.setuptools_scm.tag]
|
||||||
|
prefix = "release/"
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
from pkgutil import extend_path
|
||||||
|
|
||||||
|
__path__ = extend_path(__path__, __name__)
|
||||||
|
|
||||||
|
from ._mixin import OTelMixin
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
'OTelMixin',
|
||||||
|
]
|
||||||
@@ -0,0 +1,491 @@
|
|||||||
|
import re
|
||||||
|
import time
|
||||||
|
from asyncio import AbstractEventLoop
|
||||||
|
from logging import getLogger
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import (
|
||||||
|
Any,
|
||||||
|
AsyncGenerator,
|
||||||
|
Callable,
|
||||||
|
Dict,
|
||||||
|
Mapping,
|
||||||
|
Optional,
|
||||||
|
Sequence,
|
||||||
|
Tuple,
|
||||||
|
)
|
||||||
|
|
||||||
|
from kaya.core import HttpContext, KayaApp, KayaMixin, WebSocket, WebSocketMessage
|
||||||
|
from kaya.core._types import StrOrStrings
|
||||||
|
from opentelemetry import context as otel_context
|
||||||
|
from opentelemetry import trace
|
||||||
|
from opentelemetry.context import Context
|
||||||
|
from opentelemetry.exporter.otlp.proto.http.metric_exporter import OTLPMetricExporter
|
||||||
|
from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
|
||||||
|
from opentelemetry.metrics import Histogram, UpDownCounter
|
||||||
|
from opentelemetry.propagators.textmap import Getter
|
||||||
|
from opentelemetry.sdk.metrics import MeterProvider
|
||||||
|
from opentelemetry.sdk.metrics.export import PeriodicExportingMetricReader
|
||||||
|
from opentelemetry.sdk.resources import Resource
|
||||||
|
from opentelemetry.sdk.trace import TracerProvider
|
||||||
|
from opentelemetry.sdk.trace.export import BatchSpanProcessor
|
||||||
|
from opentelemetry.trace import Span, SpanKind, Status, StatusCode, Tracer
|
||||||
|
from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator
|
||||||
|
|
||||||
|
log = getLogger(__name__)
|
||||||
|
|
||||||
|
RequestHook = Callable[[Span, HttpContext], None]
|
||||||
|
WebSocketHook = Callable[[Span, WebSocket], None]
|
||||||
|
|
||||||
|
_DEFAULT_SANITIZE_HEADERS = (
|
||||||
|
'authorization',
|
||||||
|
'proxy-authorization',
|
||||||
|
'cookie',
|
||||||
|
'set-cookie',
|
||||||
|
)
|
||||||
|
_DEFAULT_WEBSOCKET_ERROR_CLOSE_CODES = frozenset((1002, 1003, 1007, 1008, 1009, 1010, 1011))
|
||||||
|
|
||||||
|
|
||||||
|
class _HeadersGetter(Getter[Mapping[str, Sequence[str]]]):
|
||||||
|
"""Extract propagation headers from a kaya request/websocket context.
|
||||||
|
|
||||||
|
Kaya header mappings have lowercase names with one or more values each,
|
||||||
|
which is all the W3C trace-context propagator needs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def get(self, carrier: Mapping[str, Sequence[str]], key: str) -> Optional[list[str]]:
|
||||||
|
values = carrier.get(key.lower())
|
||||||
|
if not values:
|
||||||
|
return None
|
||||||
|
return list(values)
|
||||||
|
|
||||||
|
def keys(self, carrier: Mapping[str, Sequence[str]]) -> list[str]:
|
||||||
|
return list(carrier.keys())
|
||||||
|
|
||||||
|
|
||||||
|
_PROPAGATOR = TraceContextTextMapPropagator()
|
||||||
|
_GETTER = _HeadersGetter()
|
||||||
|
|
||||||
|
|
||||||
|
def _server_attributes(ctx: HttpContext) -> Dict[str, Any]:
|
||||||
|
attributes: Dict[str, Any] = {
|
||||||
|
'http.request.method': str(ctx.method),
|
||||||
|
'url.scheme': ctx.scheme,
|
||||||
|
'url.path': ctx.path,
|
||||||
|
}
|
||||||
|
if ctx.query_string:
|
||||||
|
attributes['url.query'] = ctx.query_string
|
||||||
|
if ctx.client is not None:
|
||||||
|
attributes['client.address'] = ctx.client[0]
|
||||||
|
attributes['client.port'] = ctx.client[1]
|
||||||
|
if ctx.server is not None:
|
||||||
|
attributes['server.address'] = ctx.server[0]
|
||||||
|
if ctx.server[1] is not None:
|
||||||
|
attributes['server.port'] = ctx.server[1]
|
||||||
|
return attributes
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_header_name(name: str) -> str:
|
||||||
|
return name.lower().replace('-', '_')
|
||||||
|
|
||||||
|
|
||||||
|
def _header_values(value: StrOrStrings | Sequence[str]) -> list[str]:
|
||||||
|
if isinstance(value, str):
|
||||||
|
return [value]
|
||||||
|
return list(value)
|
||||||
|
|
||||||
|
|
||||||
|
class _SpanHolder:
|
||||||
|
"""Mutable per-request slot shared by the tracing hooks and wrappers."""
|
||||||
|
|
||||||
|
__slots__ = ('span', 'token', 'status_code', 'start_ns', 'active_attributes', 'close_code')
|
||||||
|
|
||||||
|
def __init__(self, span: Span, token: object, active_attributes: Dict[str, Any]) -> None:
|
||||||
|
self.span = span
|
||||||
|
self.token = token
|
||||||
|
self.status_code: Optional[int] = None
|
||||||
|
self.start_ns = time.perf_counter_ns()
|
||||||
|
self.active_attributes = active_attributes
|
||||||
|
self.close_code: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
|
class _TracedHttpContext(HttpContext):
|
||||||
|
"""HttpContext wrapper that records the response status code on the span.
|
||||||
|
|
||||||
|
Attributes not explicitly overridden are delegated to the wrapped context
|
||||||
|
via ``__getattr__``, so it works with any concrete ``HttpContext`` (ASGI
|
||||||
|
or RSGI) and preserves attributes set by other mixins (e.g. ``session``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, ctx: HttpContext, holder: _SpanHolder, mixin: 'OTelMixin') -> None:
|
||||||
|
object.__setattr__(self, '_ctx', ctx)
|
||||||
|
object.__setattr__(self, 'session', ctx.session)
|
||||||
|
object.__setattr__(self, '_holder', holder)
|
||||||
|
object.__setattr__(self, '_mixin', mixin)
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> Any:
|
||||||
|
if name == '_ctx':
|
||||||
|
raise AttributeError(name)
|
||||||
|
return getattr(self._ctx, name)
|
||||||
|
|
||||||
|
def _sent(self, status: int, headers: Optional[Mapping[str, StrOrStrings]]) -> None:
|
||||||
|
self._holder.status_code = status
|
||||||
|
self._holder.span.set_attribute('http.response.status_code', status)
|
||||||
|
self._mixin._capture_response_headers(self._holder.span, headers)
|
||||||
|
|
||||||
|
async def stream_body(self,
|
||||||
|
status: int,
|
||||||
|
body_generator: AsyncGenerator[bytes, None],
|
||||||
|
headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
self._sent(status, headers)
|
||||||
|
await self._ctx.stream_body(status, body_generator, headers)
|
||||||
|
|
||||||
|
async def send_bytes(self, status: int, body: bytes, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
self._sent(status, headers)
|
||||||
|
await self._ctx.send_bytes(status, body, headers)
|
||||||
|
|
||||||
|
async def send_str(self, status: int, body: str, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
self._sent(status, headers)
|
||||||
|
await self._ctx.send_str(status, body, headers)
|
||||||
|
|
||||||
|
async def send_file(self, status: int, path: Path, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
self._sent(status, headers)
|
||||||
|
await self._ctx.send_file(status, path, headers)
|
||||||
|
|
||||||
|
async def send_empty(self, status: int, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
self._sent(status, headers)
|
||||||
|
await self._ctx.send_empty(status, headers)
|
||||||
|
|
||||||
|
|
||||||
|
class _TracedWebSocket(WebSocket):
|
||||||
|
"""WebSocket wrapper that keeps the connection span reachable.
|
||||||
|
|
||||||
|
Everything is delegated to the wrapped socket via ``__getattr__``, so it
|
||||||
|
works with any concrete ``WebSocket`` (ASGI or RSGI) and preserves
|
||||||
|
attributes set by other mixins (e.g. ``session``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, ws: WebSocket, holder: _SpanHolder) -> None:
|
||||||
|
object.__setattr__(self, '_ws', ws)
|
||||||
|
object.__setattr__(self, 'session', ws.session)
|
||||||
|
object.__setattr__(self, '_holder', holder)
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> Any:
|
||||||
|
if name == '_ws':
|
||||||
|
raise AttributeError(name)
|
||||||
|
return getattr(self._ws, name)
|
||||||
|
|
||||||
|
async def accept(self, headers: Optional[Mapping[str, StrOrStrings]] = None) -> None:
|
||||||
|
await self._ws.accept(headers)
|
||||||
|
|
||||||
|
async def receive(self) -> WebSocketMessage:
|
||||||
|
return await self._ws.receive()
|
||||||
|
|
||||||
|
async def send_text(self, data: str) -> None:
|
||||||
|
await self._ws.send_text(data)
|
||||||
|
|
||||||
|
async def send_bytes(self, data: bytes) -> None:
|
||||||
|
await self._ws.send_bytes(data)
|
||||||
|
|
||||||
|
async def close(self, code: int = 1000) -> None:
|
||||||
|
self._holder.close_code = code
|
||||||
|
await self._ws.close(code)
|
||||||
|
|
||||||
|
async def __anext__(self) -> WebSocketMessage:
|
||||||
|
return await self._ws.__anext__()
|
||||||
|
|
||||||
|
|
||||||
|
class OTelMixin(KayaMixin):
|
||||||
|
"""Kaya mixin adding OpenTelemetry tracing and metrics.
|
||||||
|
|
||||||
|
HTTP requests get a ``SERVER`` span named ``<METHOD> <path>`` carrying the
|
||||||
|
usual HTTP semantic attributes; the response status code is recorded when
|
||||||
|
the handler sends the response and 5xx statuses mark the span as failed.
|
||||||
|
Exceptions escaping the handler are recorded and mark the span as failed.
|
||||||
|
WebSocket connections get one span for the whole connection lifetime.
|
||||||
|
W3C ``traceparent``/``tracestate`` headers on incoming requests are
|
||||||
|
honored, so traces propagate from upstream services.
|
||||||
|
|
||||||
|
Metrics exported on the meter provider:
|
||||||
|
|
||||||
|
- ``http.server.request.duration`` (histogram, seconds), with
|
||||||
|
``http.request.method``, ``http.route`` when known and
|
||||||
|
``http.response.status_code`` attributes;
|
||||||
|
- ``http.server.active_requests`` (up-down counter), with
|
||||||
|
``http.request.method`` attributes.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
app = KayaApp(mixins=[
|
||||||
|
OTelMixin(
|
||||||
|
service_name='my-service',
|
||||||
|
endpoint='http://localhost:4318',
|
||||||
|
)
|
||||||
|
])
|
||||||
|
|
||||||
|
``endpoint`` is the base URL of an OTLP/HTTP collector (the signal paths
|
||||||
|
``/v1/traces`` and ``/v1/metrics`` are appended); when ``None`` the
|
||||||
|
exporters fall back to their own defaults, including the standard
|
||||||
|
``OTEL_EXPORTER_OTLP_ENDPOINT`` environment variable. Custom
|
||||||
|
``tracer_provider``/``meter_provider`` instances (e.g. with in-memory
|
||||||
|
exporters for tests) can be injected instead of the OTLP defaults; the
|
||||||
|
mixin only shuts down providers it created itself.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
service_name: str,
|
||||||
|
endpoint: Optional[str] = None,
|
||||||
|
headers: Optional[Mapping[str, str]] = None,
|
||||||
|
metric_export_interval_millis: float = 60000,
|
||||||
|
tracer_provider: Optional[TracerProvider] = None,
|
||||||
|
meter_provider: Optional[MeterProvider] = None,
|
||||||
|
resource_attributes: Optional[Mapping[str, Any]] = None,
|
||||||
|
excluded_paths: Sequence[str] = (),
|
||||||
|
excluded_path_regexes: Sequence[str] = (),
|
||||||
|
capture_request_headers: Sequence[str] = (),
|
||||||
|
capture_response_headers: Sequence[str] = (),
|
||||||
|
sanitize_headers: Sequence[str] = _DEFAULT_SANITIZE_HEADERS,
|
||||||
|
server_request_hook: Optional[RequestHook] = None,
|
||||||
|
server_response_hook: Optional[RequestHook] = None,
|
||||||
|
websocket_connect_hook: Optional[WebSocketHook] = None,
|
||||||
|
websocket_close_hook: Optional[WebSocketHook] = None,
|
||||||
|
metrics_include_raw_path: bool = False,
|
||||||
|
websocket_error_close_codes: Optional[Sequence[int]] = None) -> None:
|
||||||
|
self._owns_tracer_provider = tracer_provider is None
|
||||||
|
self._owns_meter_provider = meter_provider is None
|
||||||
|
if tracer_provider is None:
|
||||||
|
tracer_provider = TracerProvider(resource=self._create_resource(service_name, resource_attributes))
|
||||||
|
span_exporter = OTLPSpanExporter(
|
||||||
|
endpoint=f'{endpoint}/v1/traces' if endpoint else None,
|
||||||
|
headers=dict(headers) if headers else None,
|
||||||
|
)
|
||||||
|
tracer_provider.add_span_processor(BatchSpanProcessor(span_exporter))
|
||||||
|
if meter_provider is None:
|
||||||
|
meter_provider = MeterProvider(
|
||||||
|
resource=self._create_resource(service_name, resource_attributes),
|
||||||
|
metric_readers=[
|
||||||
|
PeriodicExportingMetricReader(
|
||||||
|
OTLPMetricExporter(
|
||||||
|
endpoint=f'{endpoint}/v1/metrics' if endpoint else None,
|
||||||
|
headers=dict(headers) if headers else None,
|
||||||
|
),
|
||||||
|
export_interval_millis=metric_export_interval_millis,
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self._tracer_provider = tracer_provider
|
||||||
|
self._meter_provider = meter_provider
|
||||||
|
self._tracer: Tracer = tracer_provider.get_tracer('kaya.otel')
|
||||||
|
meter = meter_provider.get_meter('kaya.otel')
|
||||||
|
self._request_duration: Histogram = meter.create_histogram(
|
||||||
|
'http.server.request.duration',
|
||||||
|
unit='s',
|
||||||
|
description='Duration of HTTP server requests',
|
||||||
|
)
|
||||||
|
self._active_requests: UpDownCounter = meter.create_up_down_counter(
|
||||||
|
'http.server.active_requests',
|
||||||
|
unit='{request}',
|
||||||
|
description='Number of in-flight HTTP server requests',
|
||||||
|
)
|
||||||
|
self._app: Optional[KayaApp] = None
|
||||||
|
self._excluded_paths = frozenset(excluded_paths)
|
||||||
|
self._excluded_path_regexes = tuple(re.compile(regex) for regex in excluded_path_regexes)
|
||||||
|
self._request_headers_to_capture = tuple(capture_request_headers)
|
||||||
|
self._response_headers_to_capture = tuple(capture_response_headers)
|
||||||
|
self._sanitize_headers = frozenset(name.lower() for name in sanitize_headers)
|
||||||
|
self._server_request_hook = server_request_hook
|
||||||
|
self._server_response_hook = server_response_hook
|
||||||
|
self._websocket_connect_hook = websocket_connect_hook
|
||||||
|
self._websocket_close_hook = websocket_close_hook
|
||||||
|
self._metrics_include_raw_path = metrics_include_raw_path
|
||||||
|
self._websocket_error_close_codes = (
|
||||||
|
frozenset(websocket_error_close_codes)
|
||||||
|
if websocket_error_close_codes is not None
|
||||||
|
else _DEFAULT_WEBSOCKET_ERROR_CLOSE_CODES
|
||||||
|
)
|
||||||
|
# Live spans keyed by the id of the (wrapped) context/websocket the
|
||||||
|
# after hooks receive; entries are removed when the span ends.
|
||||||
|
self._spans: Dict[int, _SpanHolder] = {}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _create_resource(service_name: str,
|
||||||
|
resource_attributes: Optional[Mapping[str, Any]]) -> Resource:
|
||||||
|
attributes: Dict[str, Any] = {'service.name': service_name}
|
||||||
|
if resource_attributes:
|
||||||
|
attributes.update(resource_attributes)
|
||||||
|
return Resource.create(attributes)
|
||||||
|
|
||||||
|
def apply(self, app: KayaApp) -> None:
|
||||||
|
self._app = app
|
||||||
|
app.add_before_request_hook(self._before_request)
|
||||||
|
app.add_after_request_hook(self._after_request)
|
||||||
|
app.add_before_websocket_hook(self._before_websocket)
|
||||||
|
app.add_after_websocket_hook(self._after_websocket)
|
||||||
|
|
||||||
|
def _is_excluded(self, path: str) -> bool:
|
||||||
|
if path in self._excluded_paths:
|
||||||
|
return True
|
||||||
|
return any(regex.search(path) is not None for regex in self._excluded_path_regexes)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _call_hook(hook: Optional[Callable[[Span, Any], None]], span: Span, value: Any) -> None:
|
||||||
|
if hook is None:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
hook(span, value)
|
||||||
|
except Exception:
|
||||||
|
log.exception('OpenTelemetry span hook failed')
|
||||||
|
|
||||||
|
def _capture_headers(self,
|
||||||
|
span: Span,
|
||||||
|
headers: Mapping[str, StrOrStrings | Sequence[str]],
|
||||||
|
configured: Sequence[str],
|
||||||
|
prefix: str) -> None:
|
||||||
|
if not configured:
|
||||||
|
return
|
||||||
|
lower_headers = {name.lower(): value for name, value in headers.items()}
|
||||||
|
for name in configured:
|
||||||
|
key = name.lower()
|
||||||
|
if key not in lower_headers:
|
||||||
|
continue
|
||||||
|
values = _header_values(lower_headers[key])
|
||||||
|
if key in self._sanitize_headers:
|
||||||
|
values = ['REDACTED'] * len(values)
|
||||||
|
span.set_attribute(f'{prefix}.{_normalize_header_name(name)}', values)
|
||||||
|
|
||||||
|
def _capture_request_headers(self, span: Span, headers: Mapping[str, Sequence[str]]) -> None:
|
||||||
|
self._capture_headers(span, headers, self._request_headers_to_capture, 'http.request.header')
|
||||||
|
|
||||||
|
def _capture_response_headers(self,
|
||||||
|
span: Span,
|
||||||
|
headers: Optional[Mapping[str, StrOrStrings]]) -> None:
|
||||||
|
if headers is None:
|
||||||
|
return
|
||||||
|
self._capture_headers(span, headers, self._response_headers_to_capture, 'http.response.header')
|
||||||
|
|
||||||
|
def _start_span(self, name: str, attributes: Mapping[str, Any],
|
||||||
|
headers: Mapping[str, Sequence[str]]) -> _SpanHolder:
|
||||||
|
parent: Context = _PROPAGATOR.extract(carrier=headers, getter=_GETTER)
|
||||||
|
span = self._tracer.start_span(
|
||||||
|
name,
|
||||||
|
context=parent,
|
||||||
|
kind=SpanKind.SERVER,
|
||||||
|
attributes=dict(attributes),
|
||||||
|
)
|
||||||
|
self._capture_request_headers(span, headers)
|
||||||
|
token = otel_context.attach(trace.set_span_in_context(span, parent))
|
||||||
|
return _SpanHolder(span, token, {})
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _finish_span(holder: _SpanHolder, error: bool) -> None:
|
||||||
|
if error:
|
||||||
|
holder.span.set_status(Status(StatusCode.ERROR))
|
||||||
|
holder.span.end()
|
||||||
|
otel_context.detach(holder.token) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
def _pop_holder(self, obj: Any) -> Optional[_SpanHolder]:
|
||||||
|
"""Find and remove the span holder for a context/websocket.
|
||||||
|
|
||||||
|
The before hooks key holders by the id of the context they received,
|
||||||
|
but after hooks run on the outermost wrapper (later mixins may have
|
||||||
|
wrapped the context again), so the id may not match. In that case the
|
||||||
|
holder is reachable through the wrapper delegation chain as
|
||||||
|
``_holder`` and the stale dict entry is reaped by identity.
|
||||||
|
"""
|
||||||
|
holder = self._spans.pop(id(obj), None)
|
||||||
|
if holder is not None:
|
||||||
|
return holder
|
||||||
|
candidate = getattr(obj, '_holder', None)
|
||||||
|
if not isinstance(candidate, _SpanHolder):
|
||||||
|
return None
|
||||||
|
for key, value in list(self._spans.items()):
|
||||||
|
if value is candidate:
|
||||||
|
del self._spans[key]
|
||||||
|
return candidate
|
||||||
|
|
||||||
|
def _route_template(self, ctx: HttpContext) -> Optional[str]:
|
||||||
|
if self._app is None:
|
||||||
|
return None
|
||||||
|
return self._app.route_template(ctx.path, ctx.method)
|
||||||
|
|
||||||
|
async def _before_request(self, ctx: HttpContext) -> Optional[HttpContext]:
|
||||||
|
if self._is_excluded(ctx.path):
|
||||||
|
return None
|
||||||
|
holder = self._start_span(
|
||||||
|
f'{ctx.method} {ctx.path}',
|
||||||
|
_server_attributes(ctx),
|
||||||
|
ctx.headers,
|
||||||
|
)
|
||||||
|
holder.active_attributes = {'http.request.method': str(ctx.method)}
|
||||||
|
if self._metrics_include_raw_path:
|
||||||
|
holder.active_attributes['url.path'] = ctx.path
|
||||||
|
self._spans[id(ctx)] = holder
|
||||||
|
self._active_requests.add(1, holder.active_attributes)
|
||||||
|
self._call_hook(self._server_request_hook, holder.span, ctx)
|
||||||
|
return _TracedHttpContext(ctx, holder, self)
|
||||||
|
|
||||||
|
async def _after_request(self, ctx: HttpContext) -> None:
|
||||||
|
holder = self._pop_holder(ctx)
|
||||||
|
if holder is None:
|
||||||
|
return
|
||||||
|
status_code = holder.status_code
|
||||||
|
exception = getattr(ctx, 'exception', None)
|
||||||
|
error = status_code is not None and status_code >= 500
|
||||||
|
if isinstance(exception, BaseException):
|
||||||
|
holder.span.record_exception(exception)
|
||||||
|
error = True
|
||||||
|
route_template = self._route_template(ctx)
|
||||||
|
if route_template is not None:
|
||||||
|
holder.span.update_name(f'{ctx.method} {route_template}')
|
||||||
|
holder.span.set_attribute('http.route', route_template)
|
||||||
|
self._call_hook(self._server_response_hook, holder.span, ctx)
|
||||||
|
self._finish_span(holder, error)
|
||||||
|
self._active_requests.add(-1, holder.active_attributes)
|
||||||
|
metric_attributes: Dict[str, Any] = {
|
||||||
|
'http.request.method': str(ctx.method),
|
||||||
|
}
|
||||||
|
if self._metrics_include_raw_path:
|
||||||
|
metric_attributes['url.path'] = ctx.path
|
||||||
|
elif route_template is not None:
|
||||||
|
metric_attributes['http.route'] = route_template
|
||||||
|
if status_code is not None:
|
||||||
|
metric_attributes['http.response.status_code'] = status_code
|
||||||
|
duration = (time.perf_counter_ns() - holder.start_ns) / 1e9
|
||||||
|
self._request_duration.record(duration, metric_attributes)
|
||||||
|
|
||||||
|
async def _before_websocket(self, ws: WebSocket) -> Optional[WebSocket]:
|
||||||
|
if self._is_excluded(ws.path):
|
||||||
|
return None
|
||||||
|
holder = self._start_span(
|
||||||
|
f'WS {ws.path}',
|
||||||
|
{
|
||||||
|
'url.scheme': ws.scheme,
|
||||||
|
'url.path': ws.path,
|
||||||
|
},
|
||||||
|
ws.headers,
|
||||||
|
)
|
||||||
|
self._spans[id(ws)] = holder
|
||||||
|
self._call_hook(self._websocket_connect_hook, holder.span, ws)
|
||||||
|
return _TracedWebSocket(ws, holder)
|
||||||
|
|
||||||
|
async def _after_websocket(self, ws: WebSocket) -> None:
|
||||||
|
holder = self._pop_holder(ws)
|
||||||
|
if holder is None:
|
||||||
|
return
|
||||||
|
exception = getattr(ws, 'exception', None)
|
||||||
|
error = isinstance(exception, BaseException)
|
||||||
|
if isinstance(exception, BaseException):
|
||||||
|
holder.span.record_exception(exception)
|
||||||
|
close_code = holder.close_code
|
||||||
|
if close_code is not None:
|
||||||
|
holder.span.set_attribute('kaya.websocket.close_code', close_code)
|
||||||
|
error = error or close_code in self._websocket_error_close_codes
|
||||||
|
self._call_hook(self._websocket_close_hook, holder.span, ws)
|
||||||
|
self._finish_span(holder, error)
|
||||||
|
|
||||||
|
def shutdown(self, loop: AbstractEventLoop) -> None:
|
||||||
|
# Only providers created by this mixin are shut down; injected ones
|
||||||
|
# stay under the caller's control.
|
||||||
|
if self._owns_tracer_provider:
|
||||||
|
self._tracer_provider.shutdown()
|
||||||
|
if self._owns_meter_provider:
|
||||||
|
self._meter_provider.shutdown()
|
||||||
@@ -0,0 +1,394 @@
|
|||||||
|
import json
|
||||||
|
import unittest
|
||||||
|
from typing import Any, Optional, Tuple
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
from pwo import async_test
|
||||||
|
|
||||||
|
from kaya.core import HttpContext, KayaApp
|
||||||
|
from kaya.core._asgi import AsgiWebSocket
|
||||||
|
from kaya.otel import OTelMixin
|
||||||
|
from opentelemetry.sdk.metrics import MeterProvider
|
||||||
|
from opentelemetry.sdk.metrics.export import InMemoryMetricReader
|
||||||
|
from opentelemetry.sdk.resources import Resource
|
||||||
|
from opentelemetry.sdk.trace import TracerProvider
|
||||||
|
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
|
||||||
|
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
|
||||||
|
from opentelemetry.trace import StatusCode
|
||||||
|
|
||||||
|
TRACEPARENT = '00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01'
|
||||||
|
|
||||||
|
|
||||||
|
def make_mixin(**kwargs) -> Tuple[OTelMixin, InMemorySpanExporter, InMemoryMetricReader]:
|
||||||
|
resource = Resource.create({'service.name': 'test-service'})
|
||||||
|
span_exporter = InMemorySpanExporter()
|
||||||
|
tracer_provider = TracerProvider(resource=resource)
|
||||||
|
tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter))
|
||||||
|
metric_reader = InMemoryMetricReader()
|
||||||
|
meter_provider = MeterProvider(resource=resource, metric_readers=[metric_reader])
|
||||||
|
mixin = OTelMixin(
|
||||||
|
service_name='test-service',
|
||||||
|
tracer_provider=tracer_provider,
|
||||||
|
meter_provider=meter_provider,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
return mixin, span_exporter, metric_reader
|
||||||
|
|
||||||
|
|
||||||
|
def make_app(mixin: OTelMixin) -> KayaApp:
|
||||||
|
app = KayaApp(mixins=[mixin])
|
||||||
|
|
||||||
|
@app.GET('/hello')
|
||||||
|
async def hello(ctx: HttpContext) -> None:
|
||||||
|
await ctx.send_str(200, json.dumps({'ok': True}))
|
||||||
|
|
||||||
|
@app.GET('/boom')
|
||||||
|
async def boom(ctx: HttpContext) -> None:
|
||||||
|
await ctx.send_str(500, 'boom')
|
||||||
|
|
||||||
|
return app
|
||||||
|
|
||||||
|
|
||||||
|
async def request(app: KayaApp, path: str = '/hello', headers: Optional[dict[str, str]] = None) -> httpx.Response:
|
||||||
|
transport = httpx.ASGITransport(app=app, client=('127.0.0.1', 123))
|
||||||
|
async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as http_client:
|
||||||
|
return await http_client.get(path, headers=headers)
|
||||||
|
|
||||||
|
|
||||||
|
class HttpTracingTest(unittest.TestCase):
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_request_produces_server_span(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
app = make_app(mixin)
|
||||||
|
r = await request(app)
|
||||||
|
self.assertEqual(200, r.status_code)
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
span = spans[0]
|
||||||
|
self.assertEqual('GET /hello', span.name)
|
||||||
|
assert span.attributes is not None
|
||||||
|
self.assertEqual('GET', span.attributes['http.request.method'])
|
||||||
|
self.assertEqual('/hello', span.attributes['url.path'])
|
||||||
|
self.assertEqual(200, span.attributes['http.response.status_code'])
|
||||||
|
self.assertEqual(StatusCode.UNSET, span.status.status_code)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_5xx_marks_span_as_error(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
app = make_app(mixin)
|
||||||
|
r = await request(app, '/boom')
|
||||||
|
self.assertEqual(500, r.status_code)
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
self.assertEqual(500, spans[0].attributes['http.response.status_code']) # type: ignore[index]
|
||||||
|
self.assertEqual(StatusCode.ERROR, spans[0].status.status_code)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_traceparent_header_propagates(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
app = make_app(mixin)
|
||||||
|
await request(app, headers={'traceparent': TRACEPARENT})
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
span = spans[0]
|
||||||
|
self.assertEqual(0x4bf92f3577b34da6a3ce929d0e0e4736, span.context.trace_id) # type: ignore[union-attr]
|
||||||
|
assert span.parent is not None
|
||||||
|
self.assertEqual(0x00f067aa0ba902b7, span.parent.span_id)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_unmatched_route_still_traced(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
app = make_app(mixin)
|
||||||
|
r = await request(app, '/nowhere')
|
||||||
|
self.assertEqual(404, r.status_code)
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
self.assertEqual(404, spans[0].attributes['http.response.status_code']) # type: ignore[index]
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_no_leaked_spans_after_requests(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
app = make_app(mixin)
|
||||||
|
await request(app)
|
||||||
|
await request(app, '/boom')
|
||||||
|
self.assertEqual(0, len(mixin._spans))
|
||||||
|
self.assertEqual(2, len(span_exporter.get_finished_spans()))
|
||||||
|
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_exception_marks_span_as_error(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
app = KayaApp(mixins=[mixin])
|
||||||
|
|
||||||
|
@app.GET('/raises')
|
||||||
|
async def raises(ctx: HttpContext) -> None:
|
||||||
|
raise RuntimeError('boom')
|
||||||
|
|
||||||
|
transport = httpx.ASGITransport(app=app, client=('127.0.0.1', 123))
|
||||||
|
async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as http_client:
|
||||||
|
with self.assertRaises(RuntimeError):
|
||||||
|
await http_client.get('/raises')
|
||||||
|
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
span = spans[0]
|
||||||
|
self.assertEqual(StatusCode.ERROR, span.status.status_code)
|
||||||
|
self.assertTrue(any(event.name == 'exception' for event in span.events))
|
||||||
|
self.assertEqual(0, len(mixin._spans))
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_excluded_path_skips_tracing_and_metrics(self):
|
||||||
|
mixin, span_exporter, metric_reader = make_mixin(excluded_paths=('/health',))
|
||||||
|
app = KayaApp(mixins=[mixin])
|
||||||
|
|
||||||
|
@app.GET('/health')
|
||||||
|
async def health(ctx: HttpContext) -> None:
|
||||||
|
await ctx.send_str(200, 'ok')
|
||||||
|
|
||||||
|
@app.GET('/hello')
|
||||||
|
async def hello(ctx: HttpContext) -> None:
|
||||||
|
await ctx.send_str(200, 'hello')
|
||||||
|
|
||||||
|
await request(app, '/health')
|
||||||
|
await request(app, '/hello')
|
||||||
|
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
self.assertEqual('GET /hello', spans[0].name)
|
||||||
|
|
||||||
|
metrics = MetricsTest._metric_names(metric_reader)
|
||||||
|
duration_points = list(metrics['http.server.request.duration'].data.data_points)
|
||||||
|
self.assertEqual(1, len(duration_points))
|
||||||
|
self.assertEqual(1, duration_points[0].count)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_route_template_used_for_span_name_and_metrics(self):
|
||||||
|
mixin, span_exporter, metric_reader = make_mixin()
|
||||||
|
app = KayaApp(mixins=[mixin])
|
||||||
|
|
||||||
|
@app.GET('/items/${item_id:int}')
|
||||||
|
async def item(ctx: HttpContext, item_id: int) -> None:
|
||||||
|
await ctx.send_str(200, str(item_id))
|
||||||
|
|
||||||
|
r = await request(app, '/items/123')
|
||||||
|
self.assertEqual(200, r.status_code)
|
||||||
|
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
span = spans[0]
|
||||||
|
self.assertEqual('GET /items/${item_id:int}', span.name)
|
||||||
|
assert span.attributes is not None
|
||||||
|
self.assertEqual('/items/${item_id:int}', span.attributes['http.route'])
|
||||||
|
self.assertEqual('/items/123', span.attributes['url.path'])
|
||||||
|
|
||||||
|
metrics = MetricsTest._metric_names(metric_reader)
|
||||||
|
duration_points = list(metrics['http.server.request.duration'].data.data_points)
|
||||||
|
self.assertEqual(1, len(duration_points))
|
||||||
|
duration_attributes = dict(duration_points[0].attributes)
|
||||||
|
self.assertEqual('GET', duration_attributes['http.request.method'])
|
||||||
|
self.assertEqual('/items/${item_id:int}', duration_attributes['http.route'])
|
||||||
|
self.assertEqual(200, duration_attributes['http.response.status_code'])
|
||||||
|
self.assertNotIn('url.path', duration_attributes)
|
||||||
|
|
||||||
|
active_points = list(metrics['http.server.active_requests'].data.data_points)
|
||||||
|
self.assertEqual(1, len(active_points))
|
||||||
|
self.assertEqual({'http.request.method': 'GET'}, dict(active_points[0].attributes))
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_header_capture_and_sanitization(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin(
|
||||||
|
capture_request_headers=('X-Request-Id', 'Authorization'),
|
||||||
|
capture_response_headers=('X-Response-Id', 'Set-Cookie'),
|
||||||
|
)
|
||||||
|
app = KayaApp(mixins=[mixin])
|
||||||
|
|
||||||
|
@app.GET('/headers')
|
||||||
|
async def headers(ctx: HttpContext) -> None:
|
||||||
|
await ctx.send_str(200, 'ok', {
|
||||||
|
'X-Response-Id': 'res-1',
|
||||||
|
'Set-Cookie': 'sid=secret',
|
||||||
|
})
|
||||||
|
|
||||||
|
r = await request(app, '/headers', headers={
|
||||||
|
'X-Request-Id': 'req-1',
|
||||||
|
'Authorization': 'Bearer secret',
|
||||||
|
})
|
||||||
|
self.assertEqual(200, r.status_code)
|
||||||
|
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
attributes = spans[0].attributes
|
||||||
|
assert attributes is not None
|
||||||
|
self.assertEqual(['req-1'], list(attributes['http.request.header.x_request_id']))
|
||||||
|
self.assertEqual(['REDACTED'], list(attributes['http.request.header.authorization']))
|
||||||
|
self.assertEqual(['res-1'], list(attributes['http.response.header.x_response_id']))
|
||||||
|
self.assertEqual(['REDACTED'], list(attributes['http.response.header.set_cookie']))
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_http_hooks_are_called(self):
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def request_hook(span: Any, ctx: HttpContext) -> None:
|
||||||
|
calls.append(('request', ctx.path))
|
||||||
|
span.set_attribute('test.request_hook', True)
|
||||||
|
|
||||||
|
def response_hook(span: Any, ctx: HttpContext) -> None:
|
||||||
|
calls.append(('response', ctx.path))
|
||||||
|
span.set_attribute('test.response_hook', True)
|
||||||
|
|
||||||
|
mixin, span_exporter, _ = make_mixin(
|
||||||
|
server_request_hook=request_hook,
|
||||||
|
server_response_hook=response_hook,
|
||||||
|
)
|
||||||
|
app = make_app(mixin)
|
||||||
|
await request(app)
|
||||||
|
|
||||||
|
self.assertEqual([('request', '/hello'), ('response', '/hello')], calls)
|
||||||
|
attributes = span_exporter.get_finished_spans()[0].attributes
|
||||||
|
assert attributes is not None
|
||||||
|
self.assertTrue(attributes['test.request_hook'])
|
||||||
|
self.assertTrue(attributes['test.response_hook'])
|
||||||
|
|
||||||
|
|
||||||
|
class MetricsTest(unittest.TestCase):
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _metric_names(metric_reader: InMemoryMetricReader) -> dict:
|
||||||
|
data = metric_reader.get_metrics_data()
|
||||||
|
return {m.name: m for rm in data.resource_metrics for m in rm.scope_metrics[0].metrics
|
||||||
|
if rm.scope_metrics}
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_duration_histogram_and_active_requests(self):
|
||||||
|
mixin, _, metric_reader = make_mixin()
|
||||||
|
app = make_app(mixin)
|
||||||
|
await request(app)
|
||||||
|
metrics = self._metric_names(metric_reader)
|
||||||
|
self.assertIn('http.server.request.duration', metrics)
|
||||||
|
self.assertIn('http.server.active_requests', metrics)
|
||||||
|
duration_points = list(metrics['http.server.request.duration'].data.data_points)
|
||||||
|
self.assertEqual(1, len(duration_points))
|
||||||
|
self.assertEqual(1, duration_points[0].count)
|
||||||
|
self.assertGreaterEqual(duration_points[0].sum, 0)
|
||||||
|
active_points = list(metrics['http.server.active_requests'].data.data_points)
|
||||||
|
self.assertEqual(1, len(active_points))
|
||||||
|
self.assertEqual(0, active_points[0].value)
|
||||||
|
|
||||||
|
|
||||||
|
class WebSocketTracingTest(unittest.TestCase):
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _make_ws(headers=()):
|
||||||
|
async def send(message):
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def receive():
|
||||||
|
return {'type': 'websocket.connect'}
|
||||||
|
|
||||||
|
scope = {
|
||||||
|
'type': 'websocket',
|
||||||
|
'path': '/ws/games/abc',
|
||||||
|
'query_string': b'',
|
||||||
|
'scheme': 'ws',
|
||||||
|
'client': ('127.0.0.1', 12345),
|
||||||
|
'server': ('127.0.0.1', 80),
|
||||||
|
'headers': headers,
|
||||||
|
}
|
||||||
|
return AsgiWebSocket(scope, receive, send)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_websocket_lifecycle_produces_span(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
ws = self._make_ws()
|
||||||
|
wrapped = await mixin._before_websocket(ws)
|
||||||
|
self.assertIsNotNone(wrapped)
|
||||||
|
self.assertEqual(1, len(mixin._spans))
|
||||||
|
await mixin._after_websocket(wrapped) # type: ignore[arg-type]
|
||||||
|
self.assertEqual(0, len(mixin._spans))
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
self.assertEqual('WS /ws/games/abc', spans[0].name)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_websocket_traceparent_propagates(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
ws = self._make_ws(headers=[(b'traceparent', TRACEPARENT.encode())])
|
||||||
|
wrapped = await mixin._before_websocket(ws)
|
||||||
|
await mixin._after_websocket(wrapped) # type: ignore[arg-type]
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
self.assertEqual(0x4bf92f3577b34da6a3ce929d0e0e4736, spans[0].context.trace_id) # type: ignore[union-attr]
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_wrapped_websocket_delegates(self):
|
||||||
|
mixin, _, _ = make_mixin()
|
||||||
|
ws = self._make_ws()
|
||||||
|
wrapped = await mixin._before_websocket(ws)
|
||||||
|
assert wrapped is not None
|
||||||
|
self.assertEqual('/ws/games/abc', wrapped.path)
|
||||||
|
self.assertEqual(('127.0.0.1', 12345), wrapped.client)
|
||||||
|
await mixin._after_websocket(wrapped)
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_websocket_error_close_code_marks_span_error(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
ws = self._make_ws()
|
||||||
|
wrapped = await mixin._before_websocket(ws)
|
||||||
|
assert wrapped is not None
|
||||||
|
await wrapped.close(1011)
|
||||||
|
await mixin._after_websocket(wrapped)
|
||||||
|
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
self.assertEqual(StatusCode.ERROR, spans[0].status.status_code)
|
||||||
|
assert spans[0].attributes is not None
|
||||||
|
self.assertEqual(1011, spans[0].attributes['kaya.websocket.close_code'])
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_websocket_exception_marks_span_error(self):
|
||||||
|
mixin, span_exporter, _ = make_mixin()
|
||||||
|
ws = self._make_ws()
|
||||||
|
wrapped = await mixin._before_websocket(ws)
|
||||||
|
assert wrapped is not None
|
||||||
|
wrapped.exception = RuntimeError('ws boom')
|
||||||
|
await mixin._after_websocket(wrapped)
|
||||||
|
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
self.assertEqual(StatusCode.ERROR, spans[0].status.status_code)
|
||||||
|
self.assertTrue(any(event.name == 'exception' for event in spans[0].events))
|
||||||
|
|
||||||
|
@async_test
|
||||||
|
async def test_websocket_hooks_are_called(self):
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def connect_hook(span: Any, ws: Any) -> None:
|
||||||
|
calls.append(('connect', ws.path))
|
||||||
|
span.set_attribute('test.websocket_connect_hook', True)
|
||||||
|
|
||||||
|
def close_hook(span: Any, ws: Any) -> None:
|
||||||
|
calls.append(('close', ws.path))
|
||||||
|
span.set_attribute('test.websocket_close_hook', True)
|
||||||
|
|
||||||
|
mixin, span_exporter, _ = make_mixin(
|
||||||
|
websocket_connect_hook=connect_hook,
|
||||||
|
websocket_close_hook=close_hook,
|
||||||
|
)
|
||||||
|
ws = self._make_ws()
|
||||||
|
wrapped = await mixin._before_websocket(ws)
|
||||||
|
assert wrapped is not None
|
||||||
|
await wrapped.close(1000)
|
||||||
|
await mixin._after_websocket(wrapped)
|
||||||
|
|
||||||
|
self.assertEqual([('connect', '/ws/games/abc'), ('close', '/ws/games/abc')], calls)
|
||||||
|
spans = span_exporter.get_finished_spans()
|
||||||
|
self.assertEqual(1, len(spans))
|
||||||
|
self.assertEqual(StatusCode.UNSET, spans[0].status.status_code)
|
||||||
|
assert spans[0].attributes is not None
|
||||||
|
self.assertTrue(spans[0].attributes['test.websocket_connect_hook'])
|
||||||
|
self.assertTrue(spans[0].attributes['test.websocket_close_hook'])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
unittest.main()
|
||||||
@@ -23,7 +23,7 @@ from granian._granian import ( # type: ignore[attr-defined]
|
|||||||
)
|
)
|
||||||
from pwo import Maybe
|
from pwo import Maybe
|
||||||
|
|
||||||
from kaya.core import AbstractKayaApp, HttpContext, HttpMethod, WebSocket, WebSocketMessage, resolve_client
|
from kaya.core import AbstractKayaApp, HttpContext, HttpMethod, WebSocket, WebSocketMessage
|
||||||
from kaya.core._types import StrOrStrings
|
from kaya.core._types import StrOrStrings
|
||||||
|
|
||||||
|
|
||||||
@@ -51,8 +51,7 @@ class RsgiContext(HttpContext):
|
|||||||
|
|
||||||
fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc)
|
fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc)
|
||||||
self.headers = reduce(fun, scope.headers.items(), {})
|
self.headers = reduce(fun, scope.headers.items(), {})
|
||||||
self.client = resolve_client(self.headers,
|
self.client = (Maybe.of(scope.client.split(':'))
|
||||||
Maybe.of(scope.client.split(':'))
|
|
||||||
.map(lambda it: (it[0], int(it[1])))
|
.map(lambda it: (it[0], int(it[1])))
|
||||||
.or_else_throw(RuntimeError))
|
.or_else_throw(RuntimeError))
|
||||||
self.server = (Maybe.of(scope.server.split(':'))
|
self.server = (Maybe.of(scope.server.split(':'))
|
||||||
@@ -135,8 +134,7 @@ class RsgiWebSocket(WebSocket):
|
|||||||
|
|
||||||
fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc)
|
fun = cast(Callable[[Mapping[str, Sequence[str]], tuple[str, str]], Mapping[str, Sequence[str]]], acc)
|
||||||
self.headers = reduce(fun, scope.headers.items(), {})
|
self.headers = reduce(fun, scope.headers.items(), {})
|
||||||
self.client = resolve_client(self.headers,
|
self.client = (Maybe.of(scope.client.split(':'))
|
||||||
Maybe.of(scope.client.split(':'))
|
|
||||||
.map(lambda it: (it[0], int(it[1])))
|
.map(lambda it: (it[0], int(it[1])))
|
||||||
.or_else_throw(RuntimeError))
|
.or_else_throw(RuntimeError))
|
||||||
self.server = (Maybe.of(scope.server.split(':'))
|
self.server = (Maybe.of(scope.server.split(':'))
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from kaya.rsgi import RsgiContext, RsgiWebSocket
|
from kaya.rsgi import RsgiWebSocket
|
||||||
|
|
||||||
|
|
||||||
class RsgiWebSocketTest(unittest.TestCase):
|
class RsgiWebSocketTest(unittest.TestCase):
|
||||||
@@ -20,40 +20,3 @@ class RsgiWebSocketTest(unittest.TestCase):
|
|||||||
RsgiWebSocket(FakeScope(), FakeProtocol()) # type: ignore[arg-type]
|
RsgiWebSocket(FakeScope(), FakeProtocol()) # type: ignore[arg-type]
|
||||||
|
|
||||||
self.assertIn('Granian was not configured for websockets', str(ctx.exception))
|
self.assertIn('Granian was not configured for websockets', str(ctx.exception))
|
||||||
|
|
||||||
|
|
||||||
class RsgiContextTest(unittest.TestCase):
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _make_context(headers):
|
|
||||||
class FakeScope:
|
|
||||||
scheme = 'http'
|
|
||||||
method = 'GET'
|
|
||||||
path = '/'
|
|
||||||
query_string = ''
|
|
||||||
client = '127.0.0.1:12345'
|
|
||||||
server = '127.0.0.1:80'
|
|
||||||
|
|
||||||
def __init__(self, headers):
|
|
||||||
self.headers = headers
|
|
||||||
|
|
||||||
return RsgiContext(FakeScope(headers), object()) # type: ignore[arg-type]
|
|
||||||
|
|
||||||
def test_forwarded_header(self):
|
|
||||||
ctx = self._make_context({'forwarded': 'for=203.0.113.5:1234'})
|
|
||||||
self.assertEqual(('203.0.113.5', 1234), ctx.client)
|
|
||||||
|
|
||||||
def test_x_forwarded_headers(self):
|
|
||||||
ctx = self._make_context({
|
|
||||||
'x-forwarded-for': '203.0.113.5, 70.41.3.18',
|
|
||||||
'x-forwarded-port': '8443',
|
|
||||||
})
|
|
||||||
self.assertEqual(('203.0.113.5', 8443), ctx.client)
|
|
||||||
|
|
||||||
def test_x_forwarded_host_fallback(self):
|
|
||||||
ctx = self._make_context({'x-forwarded-host': '198.51.100.7'})
|
|
||||||
self.assertEqual(('198.51.100.7', 12345), ctx.client)
|
|
||||||
|
|
||||||
def test_socket_peer_fallback(self):
|
|
||||||
ctx = self._make_context({})
|
|
||||||
self.assertEqual(('127.0.0.1', 12345), ctx.client)
|
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ kaya-session-memcache @ file:./packages/kaya-session-memcache
|
|||||||
kaya-oidc @ file:./packages/kaya-oidc
|
kaya-oidc @ file:./packages/kaya-oidc
|
||||||
kaya-openapi @ file:./packages/kaya-openapi
|
kaya-openapi @ file:./packages/kaya-openapi
|
||||||
kaya-cors @ file:./packages/kaya-cors
|
kaya-cors @ file:./packages/kaya-cors
|
||||||
|
kaya-forwarded @ file:./packages/kaya-forwarded
|
||||||
|
kaya-otel @ file:./packages/kaya-otel
|
||||||
build
|
build
|
||||||
fakeredis
|
fakeredis
|
||||||
mypy
|
mypy
|
||||||
|
|||||||
@@ -44,6 +44,8 @@ executing==2.2.1
|
|||||||
# via stack-data
|
# via stack-data
|
||||||
fakeredis==2.36.2
|
fakeredis==2.36.2
|
||||||
# via -r requirements-dev.in
|
# via -r requirements-dev.in
|
||||||
|
googleapis-common-protos==1.75.3
|
||||||
|
# via opentelemetry-exporter-otlp-proto-http
|
||||||
granian==2.7.9
|
granian==2.7.9
|
||||||
# via kaya-rsgi
|
# via kaya-rsgi
|
||||||
h11==0.16.0
|
h11==0.16.0
|
||||||
@@ -90,16 +92,22 @@ file:./packages/kaya-core
|
|||||||
# via
|
# via
|
||||||
# -r requirements-dev.in
|
# -r requirements-dev.in
|
||||||
# kaya-cors
|
# kaya-cors
|
||||||
|
# kaya-forwarded
|
||||||
# kaya-oidc
|
# kaya-oidc
|
||||||
# kaya-openapi
|
# kaya-openapi
|
||||||
|
# kaya-otel
|
||||||
# kaya-rsgi
|
# kaya-rsgi
|
||||||
# kaya-session
|
# kaya-session
|
||||||
file:./packages/kaya-cors
|
file:./packages/kaya-cors
|
||||||
# via -r requirements-dev.in
|
# via -r requirements-dev.in
|
||||||
|
file:./packages/kaya-forwarded
|
||||||
|
# via -r requirements-dev.in
|
||||||
file:./packages/kaya-oidc
|
file:./packages/kaya-oidc
|
||||||
# via -r requirements-dev.in
|
# via -r requirements-dev.in
|
||||||
file:./packages/kaya-openapi
|
file:./packages/kaya-openapi
|
||||||
# via -r requirements-dev.in
|
# via -r requirements-dev.in
|
||||||
|
file:./packages/kaya-otel
|
||||||
|
# via -r requirements-dev.in
|
||||||
file:./packages/kaya-rsgi
|
file:./packages/kaya-rsgi
|
||||||
# via -r requirements-dev.in
|
# via -r requirements-dev.in
|
||||||
file:./packages/kaya-session
|
file:./packages/kaya-session
|
||||||
@@ -132,6 +140,25 @@ mypy-extensions==1.1.0
|
|||||||
# via mypy
|
# via mypy
|
||||||
nh3==0.3.6
|
nh3==0.3.6
|
||||||
# via readme-renderer
|
# via readme-renderer
|
||||||
|
opentelemetry-api==1.44.0
|
||||||
|
# via
|
||||||
|
# opentelemetry-exporter-otlp-proto-http
|
||||||
|
# opentelemetry-sdk
|
||||||
|
# opentelemetry-semantic-conventions
|
||||||
|
opentelemetry-exporter-otlp-proto-common==1.44.0
|
||||||
|
# via opentelemetry-exporter-otlp-proto-http
|
||||||
|
opentelemetry-exporter-otlp-proto-http==1.44.0
|
||||||
|
# via kaya-otel
|
||||||
|
opentelemetry-proto==1.44.0
|
||||||
|
# via
|
||||||
|
# opentelemetry-exporter-otlp-proto-common
|
||||||
|
# opentelemetry-exporter-otlp-proto-http
|
||||||
|
opentelemetry-sdk==1.44.0
|
||||||
|
# via
|
||||||
|
# kaya-otel
|
||||||
|
# opentelemetry-exporter-otlp-proto-http
|
||||||
|
opentelemetry-semantic-conventions==0.65b0
|
||||||
|
# via opentelemetry-sdk
|
||||||
packaging==26.2
|
packaging==26.2
|
||||||
# via
|
# via
|
||||||
# build
|
# build
|
||||||
@@ -144,6 +171,10 @@ pexpect==4.9.0
|
|||||||
# via ipython
|
# via ipython
|
||||||
prompt-toolkit==3.0.52
|
prompt-toolkit==3.0.52
|
||||||
# via ipython
|
# via ipython
|
||||||
|
protobuf==7.36.2
|
||||||
|
# via
|
||||||
|
# googleapis-common-protos
|
||||||
|
# opentelemetry-proto
|
||||||
psutil==7.2.2
|
psutil==7.2.2
|
||||||
# via ipython
|
# via ipython
|
||||||
ptyprocess==0.7.0
|
ptyprocess==0.7.0
|
||||||
@@ -175,6 +206,7 @@ redis==8.0.1
|
|||||||
# kaya-session-redis
|
# kaya-session-redis
|
||||||
requests==2.34.2
|
requests==2.34.2
|
||||||
# via
|
# via
|
||||||
|
# opentelemetry-exporter-otlp-proto-http
|
||||||
# requests-toolbelt
|
# requests-toolbelt
|
||||||
# twine
|
# twine
|
||||||
requests-toolbelt==1.0.0
|
requests-toolbelt==1.0.0
|
||||||
@@ -199,6 +231,10 @@ typing-extensions==4.16.0
|
|||||||
# via
|
# via
|
||||||
# kaya-core
|
# kaya-core
|
||||||
# mypy
|
# mypy
|
||||||
|
# opentelemetry-api
|
||||||
|
# opentelemetry-exporter-otlp-proto-http
|
||||||
|
# opentelemetry-sdk
|
||||||
|
# opentelemetry-semantic-conventions
|
||||||
# pwo
|
# pwo
|
||||||
urllib3==2.7.0
|
urllib3==2.7.0
|
||||||
# via
|
# via
|
||||||
|
|||||||
Reference in New Issue
Block a user