CoolFace
Modelpublic

hymenjj/llama-cpp-python-prebuilt

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
errors.py213 linesDownload Raw Back to server
1from __future__ import annotations2 3import sys4import traceback5import time6from re import compile, Match, Pattern7from typing import Callable, Coroutine, Optional, Tuple, Union, Dict8from typing_extensions import TypedDict9 10 11from fastapi import (12    Request,13    Response,14    HTTPException,15)16from fastapi.responses import JSONResponse17from fastapi.routing import APIRoute18 19from llama_cpp.server.types import (20    CreateCompletionRequest,21    CreateEmbeddingRequest,22    CreateChatCompletionRequest,23)24 25 26class ErrorResponse(TypedDict):27    """OpenAI style error response"""28 29    message: str30    type: str31    param: Optional[str]32    code: Optional[str]33 34 35class ErrorResponseFormatters:36    """Collection of formatters for error responses.37 38    Args:39        request (Union[CreateCompletionRequest, CreateChatCompletionRequest]):40            Request body41        match (Match[str]): Match object from regex pattern42 43    Returns:44        Tuple[int, ErrorResponse]: Status code and error response45    """46 47    @staticmethod48    def context_length_exceeded(49        request: Union["CreateCompletionRequest", "CreateChatCompletionRequest"],50        match,  # type: Match[str] # type: ignore51    ) -> Tuple[int, ErrorResponse]:52        """Formatter for context length exceeded error"""53 54        context_window = int(match.group(2))55        prompt_tokens = int(match.group(1))56        completion_tokens = request.max_tokens57        if hasattr(request, "messages"):58            # Chat completion59            message = (60                "This model's maximum context length is {} tokens. "61                "However, you requested {} tokens "62                "({} in the messages, {} in the completion). "63                "Please reduce the length of the messages or completion."64            )65        else:66            # Text completion67            message = (68                "This model's maximum context length is {} tokens, "69                "however you requested {} tokens "70                "({} in your prompt; {} for the completion). "71                "Please reduce your prompt; or completion length."72            )73        return 400, ErrorResponse(74            message=message.format(75                context_window,76                (completion_tokens or 0) + prompt_tokens,77                prompt_tokens,78                completion_tokens,79            ),  # type: ignore80            type="invalid_request_error",81            param="messages",82            code="context_length_exceeded",83        )84 85    @staticmethod86    def model_not_found(87        request: Union["CreateCompletionRequest", "CreateChatCompletionRequest"],88        match,  # type: Match[str] # type: ignore89    ) -> Tuple[int, ErrorResponse]:90        """Formatter for model_not_found error"""91 92        model_path = str(match.group(1))93        message = f"The model `{model_path}` does not exist"94        return 400, ErrorResponse(95            message=message,96            type="invalid_request_error",97            param=None,98            code="model_not_found",99        )100 101 102class RouteErrorHandler(APIRoute):103    """Custom APIRoute that handles application errors and exceptions"""104 105    # key: regex pattern for original error message from llama_cpp106    # value: formatter function107    pattern_and_formatters: Dict[108        "Pattern[str]",109        Callable[110            [111                Union["CreateCompletionRequest", "CreateChatCompletionRequest"],112                "Match[str]",113            ],114            Tuple[int, ErrorResponse],115        ],116    ] = {117        compile(118            r"Requested tokens \((\d+)\) exceed context window of (\d+)"119        ): ErrorResponseFormatters.context_length_exceeded,120        compile(121            r"Model path does not exist: (.+)"122        ): ErrorResponseFormatters.model_not_found,123    }124 125    def error_message_wrapper(126        self,127        error: Exception,128        body: Optional[129            Union[130                "CreateChatCompletionRequest",131                "CreateCompletionRequest",132                "CreateEmbeddingRequest",133            ]134        ] = None,135    ) -> Tuple[int, ErrorResponse]:136        """Wraps error message in OpenAI style error response"""137        if body is not None and isinstance(138            body,139            (140                CreateCompletionRequest,141                CreateChatCompletionRequest,142            ),143        ):144            # When text completion or chat completion145            for pattern, callback in self.pattern_and_formatters.items():146                match = pattern.search(str(error))147                if match is not None:148                    return callback(body, match)149 150        # Only print the trace on unexpected exceptions151        print(f"Exception: {str(error)}", file=sys.stderr)152        traceback.print_exc(file=sys.stderr)153 154        # Wrap other errors as internal server error155        return 500, ErrorResponse(156            message=str(error),157            type="internal_server_error",158            param=None,159            code=None,160        )161 162    def get_route_handler(163        self,164    ) -> Callable[[Request], Coroutine[None, None, Response]]:165        """Defines custom route handler that catches exceptions and formats166        in OpenAI style error response"""167 168        original_route_handler = super().get_route_handler()169 170        async def custom_route_handler(request: Request) -> Response:171            try:172                start_sec = time.perf_counter()173                response = await original_route_handler(request)174                elapsed_time_ms = int((time.perf_counter() - start_sec) * 1000)175                response.headers["openai-processing-ms"] = f"{elapsed_time_ms}"176                return response177            except HTTPException as unauthorized:178                # api key check failed179                raise unauthorized180            except Exception as exc:181                json_body = await request.json()182                try:183                    if "messages" in json_body:184                        # Chat completion185                        body: Optional[186                            Union[187                                CreateChatCompletionRequest,188                                CreateCompletionRequest,189                                CreateEmbeddingRequest,190                            ]191                        ] = CreateChatCompletionRequest(**json_body)192                    elif "prompt" in json_body:193                        # Text completion194                        body = CreateCompletionRequest(**json_body)195                    else:196                        # Embedding197                        body = CreateEmbeddingRequest(**json_body)198                except Exception:199                    # Invalid request body200                    body = None201 202                # Get proper error message from the exception203                (204                    status_code,205                    error_message,206                ) = self.error_message_wrapper(error=exc, body=body)207                return JSONResponse(208                    {"error": error_message},209                    status_code=status_code,210                )211 212        return custom_route_handler213