hymenjj/llama-cpp-python-prebuilt
0
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 