from __future__ import annotations
import dataclasses
import json
from dataclasses import replace
from typing import Any, Callable, Generic, Iterator, Optional, TypeVar
from net import Net
from net_sdk.stream import EventStream, SubscribeOpts, TypedEventStream
from net_sdk.types import Receipt
T = TypeVar("T")
CHANNEL_TAG_KEY = "_channel"
MAX_CHANNEL_NAME_LEN = 255
_ALLOWED_CHANNEL_CHARS = frozenset(
"abcdefghijklmnopqrstuvwxyz" "0123456789" "-_./"
)
class ChannelNameError(ValueError):
def validate_channel_name(name: str) -> str:
if not isinstance(name, str):
raise ChannelNameError(
f"channel name must be a str, got {type(name).__name__}"
)
if not name:
raise ChannelNameError("channel name must not be empty")
encoded_len = len(name.encode("utf-8"))
if encoded_len > MAX_CHANNEL_NAME_LEN:
raise ChannelNameError(
f"channel name too long: {encoded_len} bytes "
f"(max {MAX_CHANNEL_NAME_LEN})"
)
if name.startswith("/") or name.endswith("/"):
raise ChannelNameError("channel name must not start or end with '/'")
if "//" in name:
raise ChannelNameError("channel name must not contain '//'")
for ch in name:
if "A" <= ch <= "Z":
raise ChannelNameError(
f"uppercase character {ch!r} not allowed — "
"channel names are lowercase only"
)
if ch not in _ALLOWED_CHANNEL_CHARS:
raise ChannelNameError(f"invalid character {ch!r} in channel name")
for seg in name.split("/"):
if seg in (".", ".."):
raise ChannelNameError(f"path segment {seg!r} is reserved")
return name
_JSON_SCALARS = (str, int, float, bool, type(None))
def _to_dict(event: Any) -> dict:
if hasattr(event, "model_dump"):
return event.model_dump()
elif isinstance(event, dict):
return dict(event)
elif dataclasses.is_dataclass(event) and not isinstance(event, type):
return dataclasses.asdict(event)
elif hasattr(event, "__dict__"):
return dict(event.__dict__)
elif isinstance(event, _JSON_SCALARS) or isinstance(event, (list, tuple)):
return {"_value": list(event) if isinstance(event, tuple) else event}
elif hasattr(type(event), "__slots__"):
slots: dict[str, Any] = {}
for cls in type(event).__mro__:
for slot in getattr(cls, "__slots__", ()):
if slot not in slots and hasattr(event, slot):
slots[slot] = getattr(event, slot)
return slots
else:
raise TypeError(
f"cannot serialize {type(event).__name__} to a channel event: "
"pass a dict, a dataclass, a Pydantic model, or an object with "
"__dict__/__slots__"
)
def _strip_channel_tag(raw: str) -> str:
if f'"{CHANNEL_TAG_KEY}"' not in raw:
return raw
data = json.loads(raw)
if not isinstance(data, dict) or CHANNEL_TAG_KEY not in data:
return raw
data.pop(CHANNEL_TAG_KEY)
return json.dumps(data)
class TypedChannel(Generic[T]):
def __init__(
self,
bus: Net,
name: str,
model: Optional[type] = None,
parse: Optional[Callable[[str], T]] = None,
) -> None:
self._bus = bus
self._name = validate_channel_name(name)
self._model = model
self._parse = parse
self._filter = json.dumps({"path": CHANNEL_TAG_KEY, "value": name})
@property
def name(self) -> str:
return self._name
def publish(self, event: T) -> Receipt:
data = _to_dict(event)
data[CHANNEL_TAG_KEY] = self._name
result = self._bus.ingest_raw(json.dumps(data))
return Receipt(shard_id=result.shard_id, timestamp=result.timestamp)
def publish_batch(self, events: list[T]) -> int:
payloads = []
for event in events:
data = _to_dict(event)
data[CHANNEL_TAG_KEY] = self._name
payloads.append(json.dumps(data))
return self._bus.ingest_raw_batch(payloads)
def subscribe(self, opts: Optional[SubscribeOpts] = None) -> TypedEventStream[T]:
merged = SubscribeOpts() if opts is None else replace(opts)
if merged.filter is None:
merged.filter = self._filter
if self._parse is not None:
user_parse = self._parse
def parse_fn(raw: str) -> T:
return user_parse(_strip_channel_tag(raw))
elif self._model is not None:
model = self._model
def parse_fn(raw: str) -> T:
data = json.loads(raw)
data.pop(CHANNEL_TAG_KEY, None)
return model(**data) else:
def parse_fn(raw: str) -> T:
data = json.loads(raw)
data.pop(CHANNEL_TAG_KEY, None)
return data
return TypedEventStream(self._bus, parse_fn, merged)
def subscribe_raw(self, opts: Optional[SubscribeOpts] = None) -> EventStream:
merged = SubscribeOpts() if opts is None else replace(opts)
if merged.filter is None:
merged.filter = self._filter
return EventStream(self._bus, merged)