# Generated by oapi-gen 0.1.9. DO NOT EDIT.
# Source SHA256: 3bdca241e5ef74b0620c69faa4a6e49f171c0542f987ae89044b07b922113dc4
"""Streaming response helpers copied only into packages that use itemSchema."""

from __future__ import annotations

from collections.abc import AsyncGenerator, AsyncIterable, Callable, Mapping
from contextlib import AsyncExitStack
from typing import Any, NamedTuple

import anyio
import msgspec
from starlette.requests import ClientDisconnect
from starlette.responses import StreamingResponse
from starlette.types import Message, Receive, Scope, Send

from ._runtime import check_property_counts

_encoder = msgspec.json.Encoder()
_HEARTBEAT_SECONDS = 15.0


class EventField(NamedTuple):
    name: str
    required: bool
    decoder: msgspec.json.Decoder
    wire_decoder: msgspec.json.Decoder
    json_encoded: bool
    property_counts: dict[str, Any] | None


def encode_json_item(
    item: Any,
    *,
    decoder: msgspec.json.Decoder,
    checked: bool,
    media_type: str,
    property_counts: dict[str, Any] | None = None,
) -> bytes:
    if checked:
        item = msgspec.convert(msgspec.to_builtins(item), type=decoder.type)
    encoded = _encoder.encode(item)
    if checked and property_counts is not None:
        check_property_counts(msgspec.json.decode(encoded), **property_counts)
    prefix = b"\x1e" if media_type == "application/json-seq" else b""
    return prefix + encoded + b"\n"


def encode_sse_item(
    item: Any,
    *,
    event_type: type,
    fields: tuple[EventField, ...],
    checked: bool,
    property_counts: dict[str, Any] | None = None,
) -> bytes:
    if not isinstance(item, event_type):
        raise TypeError(f"SSE stream requires {event_type.__qualname__} events")
    values: dict[str, Any] = {}
    for field in fields:
        value = getattr(item, field.name)
        if value is None and not field.required:
            continue
        if field.name == "data" and not field.json_encoded and isinstance(value, str):
            # Schemas apply to the parsed event, whose line endings are LF.
            value = value.replace("\r\n", "\n").replace("\r", "\n")
        if checked:
            value = msgspec.convert(msgspec.to_builtins(value), type=field.decoder.type)
            if field.property_counts is not None:
                check_property_counts(
                    msgspec.json.decode(_encoder.encode(value)), **field.property_counts
                )
        value = (
            _encoder.encode(value).decode("utf-8")
            if field.json_encoded
            else msgspec.to_builtins(value)
        )
        if checked and field.json_encoded:
            # Validate string constraints against the serialized JSON, too.
            field.wire_decoder.decode(_encoder.encode(value))
        values[field.name] = value
    if checked and property_counts is not None:
        check_property_counts(values, **property_counts)

    lines: list[str] = []
    for name in ("event", "id", "retry", "data"):
        if name not in values:
            continue
        value = values[name]
        # Framing rules also apply when schema validation is disabled.
        if name == "retry":
            if type(value) is not int or value < 0:
                raise ValueError("SSE retry must be a non-negative integer")
        elif not isinstance(value, str):
            raise TypeError(f"SSE {name} must serialize to a string")
        elif name == "data":
            lines.extend(
                "data: " + line
                for line in value.replace("\r\n", "\n").replace("\r", "\n").split("\n")
            )
            continue
        elif "\r" in value or "\n" in value or (name == "id" and "\x00" in value):
            raise ValueError(f"SSE {name} contains a forbidden character")
        lines.append(f"{name}: {value}")
    return ("\n".join(lines) + "\n\n").encode("utf-8")


class StreamResponse(StreamingResponse):
    """Close the source and request resources on completion, error or disconnect."""

    def __init__(
        self,
        body: AsyncIterable[Any],
        *,
        encode: Callable[[Any], bytes],
        status_code: int,
        media_type: str,
        headers: Mapping[str, str] | None = None,
        resources: AsyncExitStack | None = None,
    ) -> None:
        if not isinstance(body, AsyncIterable):
            raise TypeError("streaming response body must be an AsyncIterable")
        self._source = aiter(body)
        self._encode = encode
        self._sse = media_type == "text/event-stream"
        self._encoded = self._encoded_items()
        super().__init__(
            self._encoded, status_code=status_code, media_type=media_type, headers=headers
        )
        if "content-length" in self.headers:
            raise ValueError("streaming responses cannot declare Content-Length")
        if self._sse:
            self.headers.setdefault("cache-control", "no-cache")
        self._resources = resources.pop_all() if resources is not None else AsyncExitStack()
        close = getattr(self._source, "aclose", None)
        if close is not None:
            self._resources.push_async_callback(close)
        self._resources.push_async_callback(self._encoded.aclose)

    async def _encoded_items(self) -> AsyncGenerator[bytes, None]:
        async for item in self._source:
            yield self._encode(item)

    async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
        started = anyio.Event()
        send_lock = anyio.Lock()

        async def locked_send(message: Message) -> None:
            async with send_lock:
                try:
                    await send(message)
                except OSError as error:
                    raise ClientDisconnect() from error
                if message["type"] == "http.response.start":
                    started.set()

        async def heartbeat() -> None:
            await started.wait()
            while True:
                await anyio.sleep(_HEARTBEAT_SECONDS)
                await locked_send(
                    {"type": "http.response.body", "body": b": keep-alive\n\n", "more_body": True}
                )

        try:
            async with anyio.create_task_group() as group:

                async def stream() -> None:
                    await self.stream_response(locked_send)
                    group.cancel_scope.cancel()

                group.start_soon(stream)
                if self._sse:
                    group.start_soon(heartbeat)
                await self.listen_for_disconnect(receive)
                group.cancel_scope.cancel()
        except BaseExceptionGroup as error:
            if len(error.exceptions) == 1:
                raise error.exceptions[0] from None
            raise
        finally:
            with anyio.CancelScope(shield=True):
                await self._resources.aclose()
