fix: Update max-line-length in .flake8 and refactor routing and server code for improved readability and functionality

This commit is contained in:
Илья Глазунов
2025-09-02 14:23:01 +03:00
parent 84cd1c974f
commit 6b157d7626
3 changed files with 113 additions and 93 deletions
+64 -47
View File
@@ -4,8 +4,8 @@ import time
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import Response, PlainTextResponse
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.routing import Route
from starlette.types import ASGIApp, Receive, Scope, Send
from pathlib import Path
from typing import Optional, Dict, Any
@@ -17,22 +17,28 @@ from . import __version__
logger = get_logger(__name__)
class PyServeMiddleware(BaseHTTPMiddleware):
def __init__(self, app, extension_manager: ExtensionManager):
super().__init__(app)
class PyServeMiddleware:
def __init__(self, app: ASGIApp, extension_manager: ExtensionManager):
self.app = app
self.extension_manager = extension_manager
self.access_logger = get_logger('pyserve.access')
async def dispatch(self, request: Request, call_next):
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] != "http":
await self.app(scope, receive, send)
return
start_time = time.time()
request = Request(scope, receive)
response = await self.extension_manager.process_request(request)
if response is None:
response = await call_next(request)
await self.app(scope, receive, send)
return
response = await self.extension_manager.process_response(request, response)
response.headers["Server"] = f"pyserve/{__version__}"
client_ip = request.client.host if request.client else "unknown"
method = request.method
path = str(request.url.path)
@@ -41,13 +47,13 @@ class PyServeMiddleware(BaseHTTPMiddleware):
path += f"?{query}"
status_code = response.status_code
process_time = round((time.time() - start_time) * 1000, 2)
self.access_logger.info(f"{client_ip} - {method} {path} - {status_code} - {process_time}ms")
return response
await response(scope, receive, send)
class PyServeServer:
class PyServeServer:
def __init__(self, config: Config):
self.config = config
self.extension_manager = ExtensionManager()
@@ -55,34 +61,45 @@ class PyServeServer:
self._setup_logging()
self._load_extensions()
self._create_app()
def _setup_logging(self) -> None:
self.config.setup_logging()
logger.info("PyServe сервер инициализирован")
logger.info("PyServe server initialized")
def _load_extensions(self) -> None:
for ext_config in self.config.extensions:
self.extension_manager.load_extension(
ext_config.type,
ext_config.type,
ext_config.config
)
def _create_app(self) -> None:
routes = [
Route("/health", self._health_check, methods=["GET"]),
Route("/metrics", self._metrics, methods=["GET"]),
Route("/{path:path}", self._catch_all, methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"]),
Route(
"/{path:path}",
self._catch_all,
methods=[
"GET",
"POST",
"PUT",
"DELETE",
"PATCH",
"OPTIONS"
]
),
]
self.app = Starlette(routes=routes)
self.app.add_middleware(PyServeMiddleware, extension_manager=self.extension_manager)
async def _health_check(self, request: Request) -> Response:
return PlainTextResponse("OK", status_code=200)
async def _metrics(self, request: Request) -> Response:
metrics = {}
for extension in self.extension_manager.extensions:
if hasattr(extension, 'get_metrics'):
try:
@@ -96,22 +113,22 @@ class PyServeServer:
json.dumps(metrics, ensure_ascii=False, indent=2),
media_type="application/json"
)
async def _catch_all(self, request: Request) -> Response:
return PlainTextResponse("404 Not Found", status_code=404)
def _create_ssl_context(self) -> Optional[ssl.SSLContext]:
if not self.config.ssl.enabled:
return None
if not Path(self.config.ssl.cert_file).exists():
logger.error(f"SSL certificate not found: {self.config.ssl.cert_file}")
return None
if not Path(self.config.ssl.key_file).exists():
logger.error(f"SSL key not found: {self.config.ssl.key_file}")
return None
try:
context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH)
context.load_cert_chain(
@@ -123,17 +140,16 @@ class PyServeServer:
except Exception as e:
logger.error(f"Error creating SSL context: {e}")
return None
def run(self) -> None:
if not self.config.validate():
logger.error("Configuration is invalid, server cannot be started")
return
self._ensure_directories()
ssl_context = self._create_ssl_context()
uvicorn_config = {
"app": self.app,
uvicorn_config: Dict[str, Any] = {
"host": self.config.server.host,
"port": self.config.server.port,
"log_level": "critical",
@@ -141,7 +157,7 @@ class PyServeServer:
"use_colors": False,
"server_header": False,
}
if ssl_context:
uvicorn_config.update({
"ssl_keyfile": self.config.ssl.key_file,
@@ -154,21 +170,22 @@ class PyServeServer:
logger.info(f"Starting PyServe server at {protocol}://{self.config.server.host}:{self.config.server.port}")
try:
uvicorn.run(**uvicorn_config)
assert self.app is not None, "App not initialized"
uvicorn.run(self.app, **uvicorn_config)
except KeyboardInterrupt:
logger.info("Received shutdown signal")
except Exception as e:
logger.error(f"Error starting server: {e}")
finally:
self.shutdown()
async def run_async(self) -> None:
if not self.config.validate():
logger.error("Configuration is invalid, server cannot be started")
return
self._ensure_directories()
config = uvicorn.Config(
app=self.app, # type: ignore
host=self.config.server.host,
@@ -177,24 +194,24 @@ class PyServeServer:
access_log=False,
use_colors=False,
)
server = uvicorn.Server(config)
try:
await server.serve()
finally:
self.shutdown()
def _ensure_directories(self) -> None:
directories = [
self.config.http.static_dir,
self.config.http.templates_dir,
]
log_dir = Path(self.config.logging.log_file).parent
if log_dir != Path("."):
directories.append(str(log_dir))
for directory in directories:
Path(directory).mkdir(parents=True, exist_ok=True)
logger.debug(f"Created/checked directory: {directory}")
@@ -202,7 +219,7 @@ class PyServeServer:
def shutdown(self) -> None:
logger.info("Shutting down PyServe server")
self.extension_manager.cleanup()
from .logging_utils import shutdown_logging
shutdown_logging()
@@ -210,10 +227,10 @@ class PyServeServer:
def add_extension(self, extension_type: str, config: Dict[str, Any]) -> None:
self.extension_manager.load_extension(extension_type, config)
def get_metrics(self) -> Dict[str, Any]:
metrics = {"server_status": "running"}
for extension in self.extension_manager.extensions:
if hasattr(extension, 'get_metrics'):
try: