| 131 | |
| 132 | |
| 133 | class ModelFallback(Middleware): |
| 134 | def __init__(self, fallback_model: str) -> None: |
| 135 | self.fallback_model = fallback_model |
| 136 | |
| 137 | def _fallback(self, request: APIRequest) -> APIRequest: |
| 138 | body = request.json |
| 139 | assert isinstance(body, dict) |
| 140 | return request.copy(body={**cast("dict[str, object]", body), "model": self.fallback_model}) |
| 141 | |
| 142 | @override |
| 143 | def handle(self, request: APIRequest, call_next: CallNext) -> Any: |
| 144 | response = call_next(request) |
| 145 | if response.status_code == 529: |
| 146 | return call_next(self._fallback(request)) |
| 147 | return response |
| 148 | |
| 149 | @override |
| 150 | async def handle_async(self, request: APIRequest, call_next: AsyncCallNext) -> Any: |
| 151 | response = await call_next(request) |
| 152 | if response.status_code == 529: |
| 153 | return await call_next(self._fallback(request)) |
| 154 | return response |
| 155 | |
| 156 | |
| 157 | class ShrinkMaxTokensOnTooLarge(Middleware): |
no outgoing calls