# Generated by oapi-gen 0.1.0. DO NOT EDIT.
# Source SHA256: 7b3381eb8339af725253f04fe07b232aa9922b9cd9a5216c51fec7c10670ab71
"""Small HTTP helpers copied into generated Starlette packages."""

from __future__ import annotations

import base64
import binascii
from collections.abc import AsyncIterator
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from typing import Any, Literal, NamedTuple, overload

import msgspec
from starlette.datastructures import FormData, UploadFile
from starlette.requests import Request
from starlette.responses import JSONResponse
from starlette.types import Message


class RequestError(Exception):
    def __init__(self, message: str, location: tuple[str, ...]) -> None:
        self.message = message
        self.location = location


def error_response(error: RequestError) -> JSONResponse:
    return JSONResponse(
        {"detail": [{"loc": error.location, "msg": error.message, "type": "value_error"}]},
        status_code=422,
    )


def parameter(
    value: Any,
    target: Any,
    location: tuple[str, ...],
    *,
    required: bool = True,
    default: Any = None,
) -> Any:
    if value is None or value == []:
        if required:
            raise RequestError("Field required", location)
        value = default
    try:
        return msgspec.convert(value, type=target, strict=False)
    except msgspec.ValidationError as error:
        raise RequestError(str(error), location) from error


def boolean(value: Any) -> Any:
    if isinstance(value, str):
        lowered = value.lower()
        if lowered in {"true", "1", "on", "yes", "t", "y"}:
            return True
        if lowered in {"false", "0", "off", "no", "f", "n"}:
            return False
    return value


def comma_values(value: str | list[str] | None) -> list[str] | None:
    if value is None or value == []:
        return None
    if isinstance(value, str):
        return value.split(",")
    return [item for part in value for item in part.split(",")]


@overload
def json_body[T](
    data: bytes,
    decoder: msgspec.json.Decoder[T],
    *,
    required: Literal[True],
    property_counts: dict[str, Any] | None = None,
) -> T: ...


@overload
def json_body[T](
    data: bytes,
    decoder: msgspec.json.Decoder[T],
    *,
    required: Literal[False],
    property_counts: dict[str, Any] | None = None,
) -> T | None: ...


def json_body[T](
    data: bytes,
    decoder: msgspec.json.Decoder[T],
    *,
    required: bool,
    property_counts: dict[str, Any] | None = None,
) -> T | None:
    if not data:
        if required:
            raise RequestError("Field required", ("body",))
        return None
    try:
        result = decoder.decode(data)
        if property_counts is not None:
            check_property_counts(msgspec.json.decode(data), **property_counts)
        return result
    except msgspec.DecodeError as error:
        raise RequestError(str(error), ("body",)) from error


def check_property_counts(
    value: Any,
    schema: dict[str, Any],
    references: dict[str, Any],
    path: str = "$",
    seen_refs: frozenset[str] = frozenset(),
) -> None:
    """Check wire keys before Struct decoding discards extras or inserts defaults."""
    if (ref := schema.get("$ref")) is not None and ref not in seen_refs:
        check_property_counts(value, references[ref], references, path, seen_refs | {ref})
    for member in schema.get("allOf", []):
        check_property_counts(value, member, references, path, seen_refs)
    for keyword in ("anyOf", "oneOf"):
        if keyword not in schema:
            continue
        errors = []
        tag = schema.get("discriminator", {}).get("propertyName")
        for member in schema[keyword]:
            if not _property_branch_matches(value, member, references, tag):
                continue
            try:
                check_property_counts(value, member, references, path, seen_refs)
            except msgspec.ValidationError as error:
                errors.append(error)
            else:
                break
        else:
            if errors:
                raise errors[0]
    if isinstance(value, dict):
        for keyword, outside, description in (
            ("minProperties", len(value) < schema.get("minProperties", 0), "at least"),
            ("maxProperties", len(value) > schema.get("maxProperties", len(value)), "at most"),
        ):
            if outside:
                raise msgspec.ValidationError(
                    f"Expected object with {description} {schema[keyword]} properties "
                    f"({keyword}) - at `{path}`"
                )
        properties = schema.get("properties", {})
        for name, child in value.items():
            child_schema = properties.get(name, schema.get("additionalProperties"))
            if isinstance(child_schema, dict):
                check_property_counts(child, child_schema, references, f"{path}[{name!r}]")
    elif isinstance(value, list) and "items" in schema:
        for index, child in enumerate(value):
            check_property_counts(child, schema["items"], references, f"{path}[{index}]")


def _property_branch_matches(
    value: Any,
    schema: dict[str, Any],
    references: dict[str, Any],
    tag: str | None = None,
    seen_refs: frozenset[str] = frozenset(),
) -> bool:
    # Codecs already validate unions and require tags for multiple object models.
    # Select the matching JSON type/tag so another variant's bounds cannot leak in.
    if (
        (ref := schema.get("$ref")) is not None
        and ref not in seen_refs
        and not _property_branch_matches(value, references[ref], references, tag, seen_refs | {ref})
    ):
        return False
    if value is None and schema.get("nullable") is True:
        return True
    if "const" in schema and value != schema["const"]:
        return False
    if "enum" in schema and value not in schema["enum"]:
        return False
    types = schema.get("type")
    if types is not None:
        types = types if isinstance(types, list) else [types]
        value_type = {
            dict: "object",
            list: "array",
            str: "string",
            bool: "boolean",
            int: "integer",
            float: "number",
            type(None): "null",
        }.get(type(value))
        if value_type not in types and not (value_type == "integer" and "number" in types):
            return False
    if (
        isinstance(value, dict)
        and tag in value
        and tag in schema.get("properties", {})
        and not _property_branch_matches(value[tag], schema["properties"][tag], references)
    ):
        return False
    for keyword in ("anyOf", "oneOf"):
        if keyword in schema and not any(
            _property_branch_matches(
                value,
                member,
                references,
                schema.get("discriminator", {}).get("propertyName", tag),
                seen_refs,
            )
            for member in schema[keyword]
        ):
            return False
    return all(
        _property_branch_matches(value, member, references, tag, seen_refs)
        for member in schema.get("allOf", [])
    )


def check_json_content_type(value: str | None) -> None:
    media_type = (value or "").partition(";")[0].strip().lower()
    if media_type != "application/json" and not (
        media_type.startswith("application/") and media_type.endswith("+json")
    ):
        raise RequestError("Expected an application/json content type", ("body",))


@overload
def multipart_body(
    request: Request, *, required: Literal[True]
) -> AbstractAsyncContextManager[FormData]: ...


@overload
def multipart_body(
    request: Request, *, required: Literal[False]
) -> AbstractAsyncContextManager[FormData | None]: ...


@asynccontextmanager
async def multipart_body(request: Request, *, required: bool) -> AsyncIterator[FormData | None]:
    content_type = request.headers.get("content-type", "").partition(";")[0].strip().lower()
    if content_type != "multipart/form-data":
        if await anext(request.stream()):
            raise RequestError("Expected a multipart/form-data content type", ("body",))
        if required:
            raise RequestError("Field required", ("body",))
        yield None
        return

    has_body = False
    stream = request.stream()

    async def receive() -> Message:
        nonlocal has_body
        chunk = await anext(stream)
        has_body = has_body or bool(chunk)
        return {"type": "http.request", "body": chunk, "more_body": bool(chunk)}

    # Track presence while Starlette streams the form; do not buffer uploaded files.
    async with Request(request.scope, receive=receive).form() as form:
        if not has_body and required:
            raise RequestError("Field required", ("body",))
        yield form if has_body else None


def upload(
    value: Any,
    name: str,
    *,
    required: bool,
    array: bool,
    min_length: int = 0,
    max_length: int | None = None,
) -> Any:
    if value is None or value == []:
        if required:
            raise RequestError("Field required", ("body", name))
        return None
    values = value if array else [value]
    if any(not isinstance(item, UploadFile) for item in values):
        raise RequestError("Expected UploadFile", ("body", name))
    if len(values) < min_length or (max_length is not None and len(values) > max_length):
        raise RequestError("File count does not satisfy minItems/maxItems", ("body", name))
    return value


def encode_header(value: Any) -> str:
    if isinstance(value, bool):
        return "true" if value else "false"
    if isinstance(value, list):
        return ",".join(encode_header(item) for item in value)
    if isinstance(value, dict):
        return ",".join(encode_header(part) for pair in value.items() for part in pair)
    return str(value)


class BasicCredentials(NamedTuple):
    username: str
    password: str


class BearerCredentials(NamedTuple):
    credentials: str


def basic_auth(value: str | None) -> BasicCredentials | None:
    scheme, _, credentials = (value or "").partition(" ")
    if scheme.lower() != "basic":
        return None
    try:
        decoded = base64.b64decode(credentials, validate=True).decode("ascii")
        username, separator, password = decoded.partition(":")
        if not separator:
            raise ValueError("Missing separator")
    except (ValueError, UnicodeDecodeError, binascii.Error):
        return None
    return BasicCredentials(username, password)


def bearer_auth(value: str | None) -> BearerCredentials | None:
    scheme, _, credentials = (value or "").partition(" ")
    if scheme.lower() != "bearer" or not credentials:
        return None
    return BearerCredentials(credentials)
