patx/micropie
Refactor _asgi_app_http to parse body before middleware and defer session loading
Commit f290d32 · patx · 2025-06-22T14:22:01-04:00
Refactor _asgi_app_http to parse body before middleware and defer session loading Move request body parsing before middleware execution to ensure request.body_params is populated for middleware like CSRFMiddleware. Remove early session loading to avoid conflicts with SignedSessionMiddleware, deferring session handling to middleware. Consolidate session saving to occur once after middleware and handler execution.
Comments
No comments yet.
Diff
diff --git a/MicroPie.py b/MicroPie.py
index dd5299f..e409eef 100644
--- a/MicroPie.py
+++ b/MicroPie.py
@@ -217,43 +217,12 @@ class App:
response_body: Any = ""
extra_headers: List[Tuple[str, str]] = []
try:
- # Middleware: before request
- for mw in self.middlewares:
- if result := await mw.before_request(request):
- status_code, response_body, extra_headers = (
- result["status_code"],
- result["body"],
- result.get("headers", []),
- )
- await self._send_response(send, status_code, response_body, extra_headers)
- return
-
- # Parse path and find handler
- path: str = scope["path"].lstrip("/")
- parts: List[str] = path.split("/") if path else []
- # Check if request._route_handler has been set by middleware
- if hasattr(request, "_route_handler"):
- func_name: str = request._route_handler
- else:
- func_name: str = parts[0] if parts else "index"
- if func_name.startswith("_"):
- await self._send_response(send, 404, "404 Not Found")
- return
-
- # Respect path_params set in middleware
- if not request.path_params:
- request.path_params = parts[1:] if len(parts) > 1 else []
- handler = getattr(self, func_name, None) or getattr(self, "index", None)
- if not handler:
- await self._send_response(send, 404, "404 Not Found")
- return
-
- # Parse request details
+ # Parse request details (query params, cookies, session)
request.query_params = parse_qs(scope.get("query_string", b"").decode("utf-8", "ignore"))
cookies = self._parse_cookies(request.headers.get("cookie", ""))
request.session = await self.session_backend.load(cookies.get("session_id", "")) or {}
- # Parse body parameters.
+ # Parse body parameters for POST, PUT, PATCH requests
if request.method in ("POST", "PUT", "PATCH"):
body_data = bytearray()
while True:
@@ -280,9 +249,41 @@ class App:
await self._send_response(send, 400, "400 Bad Request: Missing boundary")
return
else:
+ # Default to application/x-www-form-urlencoded
request.body_params = parse_qs(body_data.decode("utf-8", "ignore"))
- # Build function arguments from path, query, body, files, and session values.
+ # Now run middleware before_request with populated request data
+ for mw in self.middlewares:
+ if result := await mw.before_request(request):
+ status_code, response_body, extra_headers = (
+ result["status_code"],
+ result["body"],
+ result.get("headers", []),
+ )
+ await self._send_response(send, status_code, response_body, extra_headers)
+ return
+
+ # Parse path and find handler
+ path: str = scope["path"].lstrip("/")
+ parts: List[str] = path.split("/") if path else []
+ # Check if request._route_handler has been set by middleware
+ if hasattr(request, "_route_handler"):
+ func_name: str = request._route_handler
+ else:
+ func_name: str = parts[0] if parts else "index"
+ if func_name.startswith("_"):
+ await self._send_response(send, 404, "404 Not Found")
+ return
+
+ # Respect path_params set in middleware
+ if not request.path_params:
+ request.path_params = parts[1:] if len(parts) > 1 else []
+ handler = getattr(self, func_name, None) or getattr(self, "index", None)
+ if not handler:
+ await self._send_response(send, 404, "404 Not Found")
+ return
+
+ # Build function arguments from path, query, body, files, and session values
sig = inspect.signature(handler)
func_args: List[Any] = []
path_params_copy = request.path_params[:] # Create a copy to avoid modifying original
diff --git a/examples/middleware/csrf.py b/examples/middleware/csrf.py
new file mode 100644
index 0000000..c2aebf8
--- /dev/null
+++ b/examples/middleware/csrf.py
@@ -0,0 +1,70 @@
+from typing import Optional, Dict, List, Tuple, Any
+from html import escape
+import uuid
+from itsdangerous import URLSafeTimedSerializer, BadSignature
+from MicroPie import App, HttpMiddleware, Request
+
+
+class CSRFMiddleware(HttpMiddleware):
+ """Middleware for CSRF protection using itsdangerous-signed tokens."""
+ def __init__(self, app: App, secret_key: str, max_age: int = 8 * 3600):
+ self.app = app # Store the App instance
+ self.serializer = URLSafeTimedSerializer(secret_key, salt="csrf-token")
+ self.max_age = max_age
+
+ async def before_request(self, request: Request) -> Optional[Dict]:
+ """Verify CSRF token for POST/PUT/PATCH requests and generate a new token if needed."""
+ # Extract session ID from cookies or generate a new one
+ session_id = request.headers.get("cookie", "").split("session_id=")[-1].split(";")[0] if "session_id=" in request.headers.get("cookie", "") else str(uuid.uuid4())
+
+ if request.method in ("POST", "PUT", "PATCH"):
+ print(f"Request body_params: {request.body_params}") # Debugging
+ submitted_token = request.body_params.get("csrf_token", [""])[0]
+ print(f"Submitted CSRF token: {submitted_token}") # Debugging
+ if not submitted_token:
+ return {"status_code": 403, "body": "Missing CSRF token"}
+ try:
+ # Verify the submitted token's signature
+ self.serializer.loads(submitted_token, max_age=self.max_age)
+ except BadSignature:
+ return {"status_code": 403, "body": "Invalid or expired CSRF token"}
+
+ # Generate a new CSRF token if one doesn't exist in the session
+ if "csrf_token" not in request.session:
+ csrf_token = str(uuid.uuid4())
+ signed_token = self.serializer.dumps(csrf_token)
+ request.session["csrf_token"] = signed_token
+ # Save the session
+ print(f"Saving session with CSRF token: {signed_token}") # Debugging
+ await self.app.session_backend.save(session_id, request.session, self.max_age)
+
+ return None
+
+ async def after_request(
+ self, request: Request, status_code: int, response_body: Any, extra_headers: List[Tuple[str, str]]
+ ) -> Optional[Dict]:
+ """Include CSRF token in response headers for client-side use."""
+ if request.session.get("csrf_token"):
+ extra_headers.append(("X-CSRF-Token", request.session["csrf_token"]))
+ return None
+
+
+class Root(App):
+
+ async def index(self):
+ csrf_token = self.request.session.get("csrf_token", "")
+ print(f"Rendering form with CSRF token: {csrf_token}")
+ return f"""<form method="POST" action="/submit">
+ <input type="hidden" name="csrf_token" value="{escape(csrf_token)}">
+ <input type="text" name="name">
+ <button type="submit">Submit</button>
+ </form>"""
+
+ async def submit(self):
+ if self.request.method == "POST":
+ name = self.request.body_params.get("name", ["World"])[0]
+ return f"Hello {name}"
+
+
+app = Root()
+app.middlewares.append(CSRFMiddleware(app=app, secret_key="my-secret-key"))
diff --git a/examples/middleware/sessions.py b/examples/middleware/sessions.py
new file mode 100644
index 0000000..8030d86
--- /dev/null
+++ b/examples/middleware/sessions.py
@@ -0,0 +1,59 @@
+from html import escape
+import os
+import uuid
+from typing import Optional, Dict, List, Tuple, Any
+from itsdangerous import URLSafeTimedSerializer, BadSignature
+from MicroPie import App, HttpMiddleware, Request, SESSION_TIMEOUT
+
+
+class SignedSessionMiddleware(HttpMiddleware):
+ """Middleware to sign and verify session cookies using itsdangerous."""
+ def __init__(self, app: App, secret_key: str, max_age: int = SESSION_TIMEOUT):
+ self.app = app # Store the App instance
+ self.serializer = URLSafeTimedSerializer(secret_key)
+ self.max_age = max_age
+
+ async def before_request(self, request: Request) -> Optional[Dict]:
+ """Verify the session_id cookie before processing the request."""
+ cookies = self.app._parse_cookies(request.headers.get("cookie", ""))
+ session_id = cookies.get("session_id", "")
+ try:
+ verified_id = self.serializer.loads(session_id, max_age=self.max_age)
+ request.session = await self.app.session_backend.load(verified_id) or {}
+ except BadSignature:
+ request.session = {} # Invalid or expired session_id
+ return None
+
+ async def after_request(
+ self, request: Request, status_code: int, response_body: Any, extra_headers: List[Tuple[str, str]]
+ ) -> Optional[Dict]:
+ """Sign and set the session_id cookie after processing the request."""
+ if request.session:
+ cookies = self.app._parse_cookies(request.headers.get("cookie", ""))
+ session_id = cookies.get("session_id", "") or str(uuid.uuid4())
+ try:
+ session_id = self.serializer.loads(session_id, max_age=self.max_age)
+ except BadSignature:
+ session_id = str(uuid.uuid4())
+ signed_session_id = self.serializer.dumps(session_id)
+ current_session = await self.app.session_backend.load(session_id) or {}
+ if current_session != request.session:
+ await self.app.session_backend.save(session_id, request.session, SESSION_TIMEOUT)
+ if not cookies.get("session_id"):
+ extra_headers.append(("Set-Cookie", f"session_id={signed_session_id}; Path=/; SameSite=Lax; HttpOnly; Secure;"))
+ return None
+
+
+class Root(App):
+
+ async def index(self):
+ if "visits" not in self.request.session:
+ self.request.session["visits"] = 1
+ else:
+ self.request.session["visits"] += 1
+ visits = self.request.session["visits"]
+ return f"You have visited {escape(str(visits))} times."
+
+
+app = Root()
+app.middlewares.append(SignedSessionMiddleware(app=app, secret_key="my-secret-key"))