_headers.py 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148
  1. from __future__ import annotations
  2. from collections.abc import (
  3. ItemsView,
  4. Iterator,
  5. KeysView,
  6. Mapping,
  7. MutableMapping,
  8. Sequence,
  9. ValuesView,
  10. )
  11. from typing import cast
  12. class Headers(MutableMapping[str, str]):
  13. """Container of HTTP headers.
  14. This class behaves like a dictionary with case-insensitive keys and
  15. string values. It additionally can be used to store multiple values for
  16. the same key by using the `add` method, and retrieve a view including
  17. duplicates using `allitems`.
  18. """
  19. _store: dict[str, str]
  20. _extra: dict[str, list[str]] | None
  21. def __init__(
  22. self, items: Mapping[str, str] | Sequence[tuple[str, str]] = ()
  23. ) -> None:
  24. self._extra = None
  25. if isinstance(items, Mapping):
  26. # ty does not preserve type parameters when narrowing a union via
  27. # `isinstance(x, Mapping)`, so we re-assert `Mapping[str, str]`.
  28. # See https://github.com/astral-sh/ty/issues/456.
  29. items = cast("Mapping[str, str]", items)
  30. self._store = {k.lower(): v for k, v in items.items()}
  31. else:
  32. self._store = {}
  33. for k, v in items:
  34. self.add(k, v)
  35. def __getitem__(self, key: str) -> str:
  36. key = key.lower()
  37. return self._store[key]
  38. def __setitem__(self, key: str, value: str) -> None:
  39. key = key.lower()
  40. self._store[key] = value
  41. if self._extra is not None:
  42. self._extra.pop(key, None)
  43. def __delitem__(self, key: str) -> None:
  44. key = key.lower()
  45. del self._store[key]
  46. # If value wasn't in store, it's not in extra
  47. # so no need to worry about exception.
  48. if self._extra is not None:
  49. self._extra.pop(key, None)
  50. def __iter__(self) -> Iterator[str]:
  51. return iter(self._store)
  52. def __len__(self) -> int:
  53. return len(self._store)
  54. def __repr__(self) -> str:
  55. return repr(list(self.allitems()))
  56. def add(self, key: str, value: str) -> None:
  57. """Add a header, appending to existing values without overwriting.
  58. To overwrite an existing value, use `self[key] = value` instead.
  59. """
  60. key = key.lower()
  61. if key in self._store:
  62. if self._extra is None:
  63. self._extra = {}
  64. self._extra.setdefault(key, []).append(value)
  65. return
  66. self._store[key] = value
  67. def clear(self) -> None:
  68. """Clear all headers."""
  69. self._store.clear()
  70. self._extra = None
  71. def getall(self, key: str) -> Sequence[str]:
  72. """Get all values for a header key, including duplicates."""
  73. key = key.lower()
  74. if key not in self._store:
  75. return ()
  76. res = [self._store[key]]
  77. if self._extra is not None:
  78. res.extend(self._extra.get(key, ()))
  79. return res
  80. def allitems(self) -> ItemsView[str, str]:
  81. """Return an iterable view of all header items, including duplicates."""
  82. if self._extra is None:
  83. return self._store.items()
  84. return _AllItemsView(self._store, self._extra)
  85. # Commonly used functions that delegate to _store for performance vs the base class
  86. # implementations that rely on dunder methods.
  87. def items(self) -> ItemsView[str, str]:
  88. """Return an iterable view of the headers, without duplicates."""
  89. return self._store.items()
  90. def keys(self) -> KeysView[str]:
  91. """Return an iterable view of the header keys."""
  92. return self._store.keys()
  93. def values(self) -> ValuesView[str]:
  94. """Return an iterable view of the header values, without duplicates."""
  95. return self._store.values()
  96. def __contains__(self, key: object) -> bool:
  97. if not isinstance(key, str):
  98. return False
  99. key = key.lower()
  100. return key in self._store
  101. class _AllItemsView(ItemsView[str, str]):
  102. """An iterable view of all header items, including duplicates."""
  103. _store: dict[str, str]
  104. _extra: dict[str, list[str]] | None
  105. def __init__(
  106. self, store: dict[str, str], extra: dict[str, list[str]] | None
  107. ) -> None:
  108. self._store = store
  109. self._extra = extra
  110. def __iter__(self) -> Iterator[tuple[str, str]]:
  111. for key, v in self._store.items():
  112. yield (key, v)
  113. if self._extra:
  114. for vv in self._extra.get(key, ()):
  115. yield (key, vv)
  116. def __len__(self) -> int:
  117. size = len(self._store)
  118. if self._extra:
  119. size += sum(len(v) for v in self._extra.values())
  120. return size