future-framework 1.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- future/__init__.py +3 -0
- future/application.py +656 -0
- future/authentication/Auth0Authentication.py +9 -0
- future/authentication/Authentication.py +11 -0
- future/authentication/AzureADAuthentication.py +9 -0
- future/authentication/BasicAuthentication.py +9 -0
- future/authentication/KerberosAuthentication.py +9 -0
- future/authentication/KeycloakAuthentication.py +9 -0
- future/authentication/OAuth2Authentication.py +9 -0
- future/authentication/OpenIdConnectAuthentication.py +9 -0
- future/authentication/SAMLAuthentication.py +9 -0
- future/cli/__init__.py +1 -0
- future/cli/main.py +608 -0
- future/cli/stubs.py +114 -0
- future/controllers/__init__.py +4 -0
- future/controllers/base.py +8 -0
- future/controllers/builtins.py +41 -0
- future/controllers/graphql.py +23 -0
- future/controllers/openapi.py +224 -0
- future/databases/Clickhouse.py +176 -0
- future/databases/Connections.py +20 -0
- future/databases/Database.py +66 -0
- future/databases/Elasticsearch.py +146 -0
- future/databases/MongoDB.py +177 -0
- future/databases/MySQL.py +215 -0
- future/databases/Postgres.py +214 -0
- future/databases/Redis.py +146 -0
- future/databases/SQLite.py +219 -0
- future/exceptions.py +27 -0
- future/graphql/__init__.py +1 -0
- future/graphql/schema.py +148 -0
- future/lifespan.py +70 -0
- future/logger.py +53 -0
- future/middleware/Middleware.py +201 -0
- future/middleware/SessionMiddleware.py +77 -0
- future/middleware/__init__.py +21 -0
- future/migrations/Blueprint.py +54 -0
- future/migrations/Column.py +33 -0
- future/migrations/Migration.py +11 -0
- future/migrations/MigrationGenerator.py +137 -0
- future/migrations/Migrator.py +72 -0
- future/migrations/Schema.py +21 -0
- future/models/__init__.py +1 -0
- future/models/model.py +224 -0
- future/openapi.py +233 -0
- future/plugins/ElasticsearchPlugin.py +496 -0
- future/plugins/__init__.py +9 -0
- future/request.py +143 -0
- future/response.py +201 -0
- future/routing.py +342 -0
- future/seeds/SeedGenerator.py +95 -0
- future/seeds/SeedRunner.py +42 -0
- future/seeds/Seeder.py +9 -0
- future/settings.py +94 -0
- future/tasks/__init__.py +14 -0
- future/tasks/scheduler.py +267 -0
- future/testing/__init__.py +1 -0
- future/testing/client.py +135 -0
- future/types.py +47 -0
- future_framework-1.1.0.dist-info/METADATA +68 -0
- future_framework-1.1.0.dist-info/RECORD +64 -0
- future_framework-1.1.0.dist-info/WHEEL +4 -0
- future_framework-1.1.0.dist-info/entry_points.txt +3 -0
- future_framework-1.1.0.dist-info/licenses/LICENSE +21 -0
future/graphql/schema.py
ADDED
|
@@ -0,0 +1,148 @@
|
|
|
1
|
+
from textwrap import dedent
|
|
2
|
+
|
|
3
|
+
import strawberry
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
# Sample data
|
|
7
|
+
db_users = [
|
|
8
|
+
{"id": "1", "name": "Alice", "email": "alice@example.com"},
|
|
9
|
+
{"id": "2", "name": "Bob", "email": "bob@example.com"},
|
|
10
|
+
{"id": "3", "name": "Charlie", "email": "charlie@example.com"},
|
|
11
|
+
]
|
|
12
|
+
|
|
13
|
+
db_posts = [
|
|
14
|
+
{"id": "101", "title": "GraphQL vs REST", "content": "GraphQL is amazing!", "author_id": "1"},
|
|
15
|
+
{"id": "102", "title": "Strawberry Rocks", "content": "Strawberry is great for Python!", "author_id": "2"},
|
|
16
|
+
]
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
# GraphQL queries
|
|
20
|
+
queries = {
|
|
21
|
+
"GetEverything": dedent("""
|
|
22
|
+
query GetEverything {
|
|
23
|
+
users {
|
|
24
|
+
id
|
|
25
|
+
name
|
|
26
|
+
email
|
|
27
|
+
}
|
|
28
|
+
posts {
|
|
29
|
+
id
|
|
30
|
+
title
|
|
31
|
+
content
|
|
32
|
+
author {
|
|
33
|
+
id
|
|
34
|
+
name
|
|
35
|
+
email
|
|
36
|
+
}
|
|
37
|
+
}
|
|
38
|
+
}
|
|
39
|
+
"""), # type: ignore[reportGeneralTypeIssues]
|
|
40
|
+
"GetAllUsers": dedent("""
|
|
41
|
+
query GetAllUsers {
|
|
42
|
+
users {
|
|
43
|
+
id
|
|
44
|
+
name
|
|
45
|
+
email
|
|
46
|
+
}
|
|
47
|
+
}
|
|
48
|
+
"""), # type: ignore[reportGeneralTypeIssues]
|
|
49
|
+
"GetAllPosts": dedent("""
|
|
50
|
+
query GetAllPosts {
|
|
51
|
+
posts {
|
|
52
|
+
id
|
|
53
|
+
title
|
|
54
|
+
content
|
|
55
|
+
author {
|
|
56
|
+
id
|
|
57
|
+
name
|
|
58
|
+
email
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
}
|
|
62
|
+
"""), # type: ignore[reportGeneralTypeIssues]
|
|
63
|
+
"GetPostsWithAuthors": dedent("""
|
|
64
|
+
query GetPostsWithAuthors {
|
|
65
|
+
posts {
|
|
66
|
+
id
|
|
67
|
+
title
|
|
68
|
+
content
|
|
69
|
+
author {
|
|
70
|
+
name
|
|
71
|
+
email
|
|
72
|
+
}
|
|
73
|
+
}
|
|
74
|
+
}
|
|
75
|
+
"""), # type: ignore[reportGeneralTypeIssues]
|
|
76
|
+
"GetTitlesAndAuthors": dedent("""
|
|
77
|
+
query GetTitlesAndAuthors {
|
|
78
|
+
posts {
|
|
79
|
+
title
|
|
80
|
+
author {
|
|
81
|
+
name
|
|
82
|
+
}
|
|
83
|
+
}
|
|
84
|
+
}
|
|
85
|
+
"""), # type: ignore[reportGeneralTypeIssues]
|
|
86
|
+
"GetEmails": dedent("""
|
|
87
|
+
query GetEmails {
|
|
88
|
+
users {
|
|
89
|
+
email
|
|
90
|
+
}
|
|
91
|
+
}
|
|
92
|
+
"""), # type: ignore[reportGeneralTypeIssues]
|
|
93
|
+
}
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
# Define GraphQL types
|
|
97
|
+
class User:
|
|
98
|
+
id: str
|
|
99
|
+
name: str
|
|
100
|
+
email: str
|
|
101
|
+
|
|
102
|
+
def __init__(self, id: str, name: str, email: str):
|
|
103
|
+
self.id = id
|
|
104
|
+
self.name = name
|
|
105
|
+
self.email = email
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
class Post:
|
|
109
|
+
id: str
|
|
110
|
+
title: str
|
|
111
|
+
content: str
|
|
112
|
+
author: "User"
|
|
113
|
+
|
|
114
|
+
def __init__(self, id: str, title: str, content: str, author: "User"):
|
|
115
|
+
self.id = id
|
|
116
|
+
self.title = title
|
|
117
|
+
self.content = content
|
|
118
|
+
self.author = author
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
# Create Strawberry types
|
|
122
|
+
UserType = strawberry.type(User)
|
|
123
|
+
PostType = strawberry.type(Post)
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
# Define resolvers
|
|
127
|
+
async def resolve_users() -> list[UserType]: # type: ignore[misc,valid-type]
|
|
128
|
+
return [UserType(id=u["id"], name=u["name"], email=u["email"]) for u in db_users]
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
async def resolve_posts() -> list[PostType]: # type: ignore[misc,valid-type]
|
|
132
|
+
posts = []
|
|
133
|
+
for p in db_posts:
|
|
134
|
+
author = next(u for u in db_users if u["id"] == p["author_id"])
|
|
135
|
+
user = UserType(id=author["id"], name=author["name"], email=author["email"])
|
|
136
|
+
post = PostType(id=p["id"], title=p["title"], content=p["content"], author=user)
|
|
137
|
+
posts.append(post)
|
|
138
|
+
return posts
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
# Define Query type
|
|
142
|
+
class Query:
|
|
143
|
+
users = strawberry.field(resolver=resolve_users)
|
|
144
|
+
posts = strawberry.field(resolver=resolve_posts)
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
# NB:!!!! schema cannot be instantiated in the controller. it breaks then.
|
|
148
|
+
schema = strawberry.Schema(query=strawberry.type(Query))
|
future/lifespan.py
ADDED
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
from future.logger import log
|
|
6
|
+
from future.tasks.scheduler import CronScheduler, Task
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class Lifespan:
|
|
10
|
+
"""ASGI application lifespan: startup / shutdown, with optional cron tasks."""
|
|
11
|
+
|
|
12
|
+
def __init__(self, startup_tasks: list[Task] | None = None, shutdown_tasks: list[Task] | None = None, cron_tasks: list[Task] | None = None) -> None:
|
|
13
|
+
self.app = None
|
|
14
|
+
self.startup_tasks = startup_tasks or []
|
|
15
|
+
self.shutdown_tasks = shutdown_tasks or []
|
|
16
|
+
self.cron_tasks = cron_tasks or []
|
|
17
|
+
self.scheduler = CronScheduler()
|
|
18
|
+
self.db = None
|
|
19
|
+
self.s3_client = None
|
|
20
|
+
self.redis_client = None
|
|
21
|
+
self.settings = None
|
|
22
|
+
|
|
23
|
+
async def __aenter__(self) -> dict[str, Any]:
|
|
24
|
+
"""Application startup."""
|
|
25
|
+
async with asyncio.timeout(30):
|
|
26
|
+
log.info("Starting application...")
|
|
27
|
+
await self._run_startup_tasks()
|
|
28
|
+
await self.scheduler.start()
|
|
29
|
+
await self._register_cron_jobs()
|
|
30
|
+
log.info("Application startup complete")
|
|
31
|
+
return {
|
|
32
|
+
"db": self.db,
|
|
33
|
+
"s3_client": self.s3_client,
|
|
34
|
+
"redis_client": self.redis_client,
|
|
35
|
+
"settings": self.settings,
|
|
36
|
+
}
|
|
37
|
+
|
|
38
|
+
async def __aexit__(self, exc_type: type | None, exc_value: Exception | None, tb: Any) -> bool | None:
|
|
39
|
+
"""Application shutdown."""
|
|
40
|
+
async with asyncio.timeout(30):
|
|
41
|
+
log.info("Shutting down application...")
|
|
42
|
+
await self.scheduler.stop()
|
|
43
|
+
await self._run_shutdown_tasks()
|
|
44
|
+
log.info("Shutdown complete")
|
|
45
|
+
return None
|
|
46
|
+
|
|
47
|
+
async def _run_startup_tasks(self) -> None:
|
|
48
|
+
for task in self.startup_tasks:
|
|
49
|
+
if task.func is not None:
|
|
50
|
+
if asyncio.iscoroutinefunction(task.func):
|
|
51
|
+
await task.func(*task.args, **task.kwargs)
|
|
52
|
+
else:
|
|
53
|
+
loop = asyncio.get_event_loop()
|
|
54
|
+
await loop.run_in_executor(None, task.func, *task.args, **task.kwargs)
|
|
55
|
+
|
|
56
|
+
async def _run_shutdown_tasks(self) -> None:
|
|
57
|
+
for task in self.shutdown_tasks:
|
|
58
|
+
if task.func is not None:
|
|
59
|
+
if asyncio.iscoroutinefunction(task.func):
|
|
60
|
+
await task.func(*task.args, **task.kwargs)
|
|
61
|
+
else:
|
|
62
|
+
loop = asyncio.get_event_loop()
|
|
63
|
+
await loop.run_in_executor(None, task.func, *task.args, **task.kwargs)
|
|
64
|
+
|
|
65
|
+
async def _register_cron_jobs(self) -> None:
|
|
66
|
+
"""Register cron jobs with the scheduler."""
|
|
67
|
+
for task in self.cron_tasks:
|
|
68
|
+
self.scheduler.add_task(task)
|
|
69
|
+
|
|
70
|
+
log.info(f"Registered {len(self.scheduler.tasks)} cron jobs")
|
future/logger.py
ADDED
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
import atexit
|
|
2
|
+
import logging
|
|
3
|
+
import logging.config
|
|
4
|
+
import logging.handlers
|
|
5
|
+
|
|
6
|
+
from future.settings import LOGGING_CONFIG
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
# ANSI color codes
|
|
10
|
+
class ColoredFormatter(logging.Formatter):
|
|
11
|
+
"""Custom formatter with colors like uvicorn"""
|
|
12
|
+
|
|
13
|
+
COLORS = {
|
|
14
|
+
"DEBUG": "\033[36m", # Cyan
|
|
15
|
+
"INFO": "\033[32m", # Green
|
|
16
|
+
"WARNING": "\033[33m", # Yellow
|
|
17
|
+
"ERROR": "\033[31m", # Red
|
|
18
|
+
"CRITICAL": "\033[35m", # Magenta
|
|
19
|
+
}
|
|
20
|
+
RESET = "\033[0m"
|
|
21
|
+
|
|
22
|
+
def format(self, record: logging.LogRecord) -> str:
|
|
23
|
+
# Add color to levelname
|
|
24
|
+
levelname = record.levelname
|
|
25
|
+
if levelname in self.COLORS:
|
|
26
|
+
# Calculate padding based on original levelname length BEFORE adding colors
|
|
27
|
+
padding = " " * (9 - len(levelname)) # 9 is max level length (CRITICAL) + 1
|
|
28
|
+
colored_levelname = f"{self.COLORS[levelname]}{levelname}{self.RESET}"
|
|
29
|
+
return f"{colored_levelname}:{padding}{record.getMessage()}"
|
|
30
|
+
return super().format(record)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
# Configure logging using settings
|
|
34
|
+
logging.config.dictConfig(LOGGING_CONFIG)
|
|
35
|
+
|
|
36
|
+
log = logging.getLogger("future")
|
|
37
|
+
formatter = ColoredFormatter("%(message)s")
|
|
38
|
+
|
|
39
|
+
# QueueHandler itself does not format; the listener's stdout handler does.
|
|
40
|
+
# Handler level NOTSET so Future's log.setLevel(APP_DEBUG) alone controls filtering.
|
|
41
|
+
queue_handler = logging.getHandlerByName("queue_handler")
|
|
42
|
+
if queue_handler is not None and isinstance(queue_handler, logging.handlers.QueueHandler):
|
|
43
|
+
listener = getattr(queue_handler, "listener", None)
|
|
44
|
+
if listener is not None:
|
|
45
|
+
for handler in listener.handlers:
|
|
46
|
+
handler.setFormatter(formatter)
|
|
47
|
+
handler.setLevel(logging.NOTSET)
|
|
48
|
+
listener.start()
|
|
49
|
+
atexit.register(listener.stop)
|
|
50
|
+
else:
|
|
51
|
+
for handler in log.handlers:
|
|
52
|
+
handler.setFormatter(formatter)
|
|
53
|
+
handler.setLevel(logging.NOTSET)
|
|
@@ -0,0 +1,201 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
|
|
3
|
+
from future.request import Request
|
|
4
|
+
from future.response import Response
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
# README - a note on the use of middlewares:
|
|
8
|
+
# if a middleware does not RETURN or hit any EXCEPTIONS, it means it passes (all checks ok)
|
|
9
|
+
# if it however RETURNS or hits an EXCEPTION, the middleware check will deny further processing of the request.
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class Middleware:
|
|
13
|
+
name: Optional[str] = None
|
|
14
|
+
apply: bool = True
|
|
15
|
+
priority: int = 0
|
|
16
|
+
|
|
17
|
+
def __init__(self, request: Request, response: Response) -> None:
|
|
18
|
+
self.request = request
|
|
19
|
+
self.response = response
|
|
20
|
+
|
|
21
|
+
async def before(self) -> Optional[Response]:
|
|
22
|
+
return None
|
|
23
|
+
|
|
24
|
+
async def after(self) -> Optional[Response]:
|
|
25
|
+
return None
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class TestMiddlewareRequest(Middleware):
|
|
29
|
+
name = "testRequestMiddleware"
|
|
30
|
+
|
|
31
|
+
async def before(self) -> Optional[Response]:
|
|
32
|
+
if self.request.headers.get("x-interrupt") == "1":
|
|
33
|
+
return self.response.text("Request intercepted!")
|
|
34
|
+
return None
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class TestMiddlewareResponse(Middleware):
|
|
38
|
+
name = "testResponseMiddleware"
|
|
39
|
+
|
|
40
|
+
async def after(self) -> Optional[Response]:
|
|
41
|
+
if self.request.headers.get("x-interrupt-response") == "1":
|
|
42
|
+
return self.response.text("Response intercepted!")
|
|
43
|
+
return None
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class ResponseCodeConfuser(Middleware):
|
|
47
|
+
name = "response code confuser"
|
|
48
|
+
|
|
49
|
+
async def after(self) -> Optional[Response]:
|
|
50
|
+
import random
|
|
51
|
+
|
|
52
|
+
response_codes = [
|
|
53
|
+
100, 101, 102, 103, 200, 201, 202, 203, 204, 205, 206, 207, 208, 226,
|
|
54
|
+
300, 301, 302, 303, 304, 305, 306, 307, 308,
|
|
55
|
+
400, 401, 402, 403, 404, 405, 406, 407, 408, 409, 410, 411, 412, 413, 414,
|
|
56
|
+
415, 416, 417, 418, 421, 422, 423, 424, 425, 426, 428, 429, 431, 451,
|
|
57
|
+
500, 501, 502, 503, 504, 505, 506, 507, 508, 510, 511,
|
|
58
|
+
]
|
|
59
|
+
if "curl" in self.request.headers.get("user-agent", ""):
|
|
60
|
+
return self.response.empty(status=random.choice(response_codes))
|
|
61
|
+
return None
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
class RateLimitMiddleware(Middleware):
|
|
65
|
+
name = "RateLimitMiddleware"
|
|
66
|
+
limit = 60
|
|
67
|
+
window_seconds = 60
|
|
68
|
+
_hits: dict[str, list[float]] = {}
|
|
69
|
+
|
|
70
|
+
async def before(self) -> Optional[Response]:
|
|
71
|
+
import time
|
|
72
|
+
client = (self.request.scope.get("client") or ("unknown", 0))[0]
|
|
73
|
+
now = time.time()
|
|
74
|
+
bucket = self._hits.setdefault(client, [])
|
|
75
|
+
cutoff = now - self.window_seconds
|
|
76
|
+
self._hits[client] = [stamp for stamp in bucket if stamp >= cutoff]
|
|
77
|
+
if len(self._hits[client]) >= self.limit:
|
|
78
|
+
return self.response.json({"error": "Too Many Requests", "status_code": 429}, status=429)
|
|
79
|
+
self._hits[client].append(now)
|
|
80
|
+
return None
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
class CSRFMiddleware(Middleware):
|
|
84
|
+
name = "CSRFMiddleware"
|
|
85
|
+
cookie_name = "csrf_token"
|
|
86
|
+
header_name = "x-csrf-token"
|
|
87
|
+
safe_methods = {"GET", "HEAD", "OPTIONS"}
|
|
88
|
+
|
|
89
|
+
async def before(self) -> Optional[Response]:
|
|
90
|
+
import secrets
|
|
91
|
+
token = self.request.cookies.get(self.cookie_name)
|
|
92
|
+
if not token:
|
|
93
|
+
token = secrets.token_urlsafe(32)
|
|
94
|
+
self.request.context["csrf_token_new"] = token
|
|
95
|
+
self.request.context["csrf_token"] = token
|
|
96
|
+
if self.request.method in self.safe_methods:
|
|
97
|
+
return None
|
|
98
|
+
provided = self.request.headers.get(self.header_name) or self.request.headers.get("x-xsrf-token")
|
|
99
|
+
if not provided:
|
|
100
|
+
form = await self.request.form()
|
|
101
|
+
value = form.get("csrf_token")
|
|
102
|
+
provided = value if isinstance(value, str) else None
|
|
103
|
+
if not provided or not secrets.compare_digest(str(provided), str(token)):
|
|
104
|
+
return self.response.json({"error": "CSRF token missing or invalid", "status_code": 403}, status=403)
|
|
105
|
+
return None
|
|
106
|
+
|
|
107
|
+
async def after(self) -> Optional[Response]:
|
|
108
|
+
token = self.request.context.get("csrf_token_new") or self.request.context.get("csrf_token")
|
|
109
|
+
if token:
|
|
110
|
+
self.response.set_cookie(self.cookie_name, token, httponly=False, samesite="Strict", secure=(self.request.scheme == "https"))
|
|
111
|
+
return None
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
class WebServerConfuser(Middleware):
|
|
115
|
+
pass
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
class BruteforcePrevention(Middleware):
|
|
119
|
+
pass
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
class SQLiConfuser(Middleware):
|
|
123
|
+
pass
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
class HTAccessConfuser(Middleware):
|
|
127
|
+
pass
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
class HeaderConfuser(Middleware):
|
|
131
|
+
pass
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
class GZipMiddleware(Middleware):
|
|
135
|
+
name = "GZipMiddleware"
|
|
136
|
+
|
|
137
|
+
async def after(self) -> Optional[Response]:
|
|
138
|
+
accept = self.request.headers.get("accept-encoding", "")
|
|
139
|
+
if "gzip" not in accept.lower():
|
|
140
|
+
return None
|
|
141
|
+
if not self.response.body:
|
|
142
|
+
return None
|
|
143
|
+
for key, _value in self.response.headers:
|
|
144
|
+
if key.lower() == b"content-encoding":
|
|
145
|
+
return None
|
|
146
|
+
import gzip
|
|
147
|
+
compressed = gzip.compress(self.response.body)
|
|
148
|
+
self.response.body = compressed
|
|
149
|
+
self.response.headers = [pair for pair in self.response.headers if pair[0].lower() != b"content-length"]
|
|
150
|
+
self.response.headers.append([b"content-encoding", b"gzip"])
|
|
151
|
+
self.response.headers.append([b"content-length", str(len(compressed)).encode("utf-8")])
|
|
152
|
+
return None
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
class CORSMiddleware(Middleware):
|
|
156
|
+
name = "CORSMiddleware"
|
|
157
|
+
allow_origin = "*"
|
|
158
|
+
allow_methods = "GET, POST, PUT, PATCH, DELETE, OPTIONS"
|
|
159
|
+
allow_headers = "origin, content-type, accept, authorization, x-xsrf-token, x-request-id"
|
|
160
|
+
allow_credentials = "false"
|
|
161
|
+
|
|
162
|
+
def _apply_headers(self) -> None:
|
|
163
|
+
self.response.headers.append([b"access-control-allow-origin", self.allow_origin.encode("utf-8")])
|
|
164
|
+
self.response.headers.append([b"access-control-allow-methods", self.allow_methods.encode("utf-8")])
|
|
165
|
+
self.response.headers.append([b"access-control-allow-headers", self.allow_headers.encode("utf-8")])
|
|
166
|
+
if self.allow_origin != "*":
|
|
167
|
+
self.response.headers.append([b"access-control-allow-credentials", self.allow_credentials.encode("utf-8")])
|
|
168
|
+
|
|
169
|
+
async def before(self) -> Optional[Response]:
|
|
170
|
+
if self.request.method == "OPTIONS":
|
|
171
|
+
self.response.empty(status=204)
|
|
172
|
+
self._apply_headers()
|
|
173
|
+
return self.response
|
|
174
|
+
return None
|
|
175
|
+
|
|
176
|
+
async def after(self) -> Optional[Response]:
|
|
177
|
+
self._apply_headers()
|
|
178
|
+
return None
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
class ScopeValidationMiddleware(Middleware):
|
|
182
|
+
name = "scopeValidation"
|
|
183
|
+
priority = 10
|
|
184
|
+
|
|
185
|
+
async def before(self) -> Optional[Response]:
|
|
186
|
+
if self.request.path in ["/health", "/ping"]:
|
|
187
|
+
return None
|
|
188
|
+
route = getattr(self.request, "route", None)
|
|
189
|
+
if not route:
|
|
190
|
+
return None
|
|
191
|
+
required_scopes = getattr(route, "scopes", [])
|
|
192
|
+
if not required_scopes:
|
|
193
|
+
return None
|
|
194
|
+
user_id = self.request.context.get("user_id")
|
|
195
|
+
if not user_id:
|
|
196
|
+
return self.response.text("Unauthorized - no user ID found", status=401)
|
|
197
|
+
user_scopes = ["read:posts", "write:posts", "admin", "read:public", "read:api", "user", "debug"]
|
|
198
|
+
for required_scope in required_scopes:
|
|
199
|
+
if required_scope not in user_scopes:
|
|
200
|
+
return self.response.text("Insufficient permissions", status=403)
|
|
201
|
+
return None
|
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
import base64
|
|
2
|
+
import hashlib
|
|
3
|
+
import hmac
|
|
4
|
+
import json
|
|
5
|
+
|
|
6
|
+
from typing import Optional
|
|
7
|
+
|
|
8
|
+
from future.middleware.Middleware import Middleware
|
|
9
|
+
from future.response import Response
|
|
10
|
+
from future.settings import APP_KEY
|
|
11
|
+
|
|
12
|
+
SESSION_COOKIE_NAME = "session"
|
|
13
|
+
SESSION_COOKIE_PATH = "/"
|
|
14
|
+
SESSION_COOKIE_HTTPONLY = True
|
|
15
|
+
SESSION_COOKIE_SAMESITE = "Lax"
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class SessionMiddleware(Middleware):
|
|
19
|
+
name = "session"
|
|
20
|
+
|
|
21
|
+
def _secret(self) -> str:
|
|
22
|
+
app = self.request.scope.get("app")
|
|
23
|
+
if app is not None and getattr(app, "config", None):
|
|
24
|
+
return str(app.config.get("APP_KEY") or APP_KEY)
|
|
25
|
+
return str(APP_KEY)
|
|
26
|
+
|
|
27
|
+
def _sign(self, payload: bytes) -> str:
|
|
28
|
+
digest = hmac.new(self._secret().encode("utf-8"), payload, hashlib.sha256).hexdigest()
|
|
29
|
+
return base64.urlsafe_b64encode(payload).decode("ascii") + "." + digest
|
|
30
|
+
|
|
31
|
+
def _unsign(self, raw: str) -> dict:
|
|
32
|
+
body_b64, digest = raw.rsplit(".", 1)
|
|
33
|
+
payload = base64.urlsafe_b64decode(body_b64.encode("ascii"))
|
|
34
|
+
expected = hmac.new(self._secret().encode("utf-8"), payload, hashlib.sha256).hexdigest()
|
|
35
|
+
if not hmac.compare_digest(digest, expected):
|
|
36
|
+
raise ValueError("invalid session signature")
|
|
37
|
+
data = json.loads(payload.decode("utf-8"))
|
|
38
|
+
if not isinstance(data, dict):
|
|
39
|
+
raise ValueError("invalid session payload")
|
|
40
|
+
return data
|
|
41
|
+
|
|
42
|
+
async def before(self) -> Optional[Response]:
|
|
43
|
+
raw = self.request.cookies.get(SESSION_COOKIE_NAME)
|
|
44
|
+
if not raw:
|
|
45
|
+
return None
|
|
46
|
+
try:
|
|
47
|
+
if "." in raw:
|
|
48
|
+
self.request.session = self._unsign(raw)
|
|
49
|
+
else:
|
|
50
|
+
# Legacy unsigned base64 JSON sessions
|
|
51
|
+
payload = base64.urlsafe_b64decode(raw.encode("ascii"))
|
|
52
|
+
data = json.loads(payload.decode("utf-8"))
|
|
53
|
+
self.request.session = data if isinstance(data, dict) else {}
|
|
54
|
+
except Exception:
|
|
55
|
+
self.request.session = {}
|
|
56
|
+
return None
|
|
57
|
+
|
|
58
|
+
async def after(self) -> Optional[Response]:
|
|
59
|
+
if not self.request.session:
|
|
60
|
+
self.response.delete_cookie(SESSION_COOKIE_NAME, path=SESSION_COOKIE_PATH)
|
|
61
|
+
return None
|
|
62
|
+
payload = json.dumps(
|
|
63
|
+
self.request.session,
|
|
64
|
+
default=str,
|
|
65
|
+
ensure_ascii=False,
|
|
66
|
+
separators=(",", ":"),
|
|
67
|
+
).encode("utf-8")
|
|
68
|
+
value = self._sign(payload)
|
|
69
|
+
self.response.set_cookie(
|
|
70
|
+
SESSION_COOKIE_NAME,
|
|
71
|
+
value,
|
|
72
|
+
path=SESSION_COOKIE_PATH,
|
|
73
|
+
httponly=SESSION_COOKIE_HTTPONLY,
|
|
74
|
+
samesite=SESSION_COOKIE_SAMESITE,
|
|
75
|
+
secure=(self.request.scheme == "https"),
|
|
76
|
+
)
|
|
77
|
+
return None
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
from future.middleware.Middleware import (
|
|
2
|
+
CORSMiddleware,
|
|
3
|
+
CSRFMiddleware,
|
|
4
|
+
GZipMiddleware,
|
|
5
|
+
Middleware,
|
|
6
|
+
RateLimitMiddleware,
|
|
7
|
+
ScopeValidationMiddleware,
|
|
8
|
+
TestMiddlewareRequest,
|
|
9
|
+
TestMiddlewareResponse,
|
|
10
|
+
)
|
|
11
|
+
from future.middleware.SessionMiddleware import (
|
|
12
|
+
SESSION_COOKIE_HTTPONLY,
|
|
13
|
+
SESSION_COOKIE_NAME,
|
|
14
|
+
SESSION_COOKIE_PATH,
|
|
15
|
+
SESSION_COOKIE_SAMESITE,
|
|
16
|
+
SessionMiddleware,
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
# Back-compat aliases during migration
|
|
20
|
+
SessionRequestMiddleware = SessionMiddleware
|
|
21
|
+
SessionResponseMiddleware = SessionMiddleware
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
# Schema blueprint — collects portable columns; applied on context exit.
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
from .Column import Column
|
|
5
|
+
from future.databases.Connections import Connections
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class Blueprint:
|
|
9
|
+
def __init__(self, name, connection_name, action="create"):
|
|
10
|
+
self.name = name
|
|
11
|
+
self.connection_name = connection_name
|
|
12
|
+
self.action = action
|
|
13
|
+
self.columns = []
|
|
14
|
+
|
|
15
|
+
def __enter__(self):
|
|
16
|
+
return self
|
|
17
|
+
|
|
18
|
+
def __exit__(self, exc_type, exc, tb):
|
|
19
|
+
if exc_type is not None:
|
|
20
|
+
return False
|
|
21
|
+
connection = Connections().get_connection(self.connection_name)
|
|
22
|
+
if self.action == "create":
|
|
23
|
+
connection.schema_create(self)
|
|
24
|
+
return False
|
|
25
|
+
|
|
26
|
+
def _add(self, column):
|
|
27
|
+
self.columns.append(column)
|
|
28
|
+
return column
|
|
29
|
+
|
|
30
|
+
def id(self):
|
|
31
|
+
return self._add(Column("id", "string", length=255)).primary()
|
|
32
|
+
|
|
33
|
+
def string(self, name, length=255):
|
|
34
|
+
return self._add(Column(name, "string", length=length))
|
|
35
|
+
|
|
36
|
+
def text(self, name):
|
|
37
|
+
return self._add(Column(name, "text"))
|
|
38
|
+
|
|
39
|
+
def integer(self, name):
|
|
40
|
+
return self._add(Column(name, "integer"))
|
|
41
|
+
|
|
42
|
+
def float(self, name):
|
|
43
|
+
return self._add(Column(name, "float"))
|
|
44
|
+
|
|
45
|
+
def boolean(self, name):
|
|
46
|
+
return self._add(Column(name, "boolean"))
|
|
47
|
+
|
|
48
|
+
def datetime(self, name):
|
|
49
|
+
return self._add(Column(name, "datetime"))
|
|
50
|
+
|
|
51
|
+
def timestamps(self):
|
|
52
|
+
self.datetime("created_at").nullable()
|
|
53
|
+
self.datetime("updated_at").nullable()
|
|
54
|
+
return self
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
# Portable column definition for Schema blueprints (Orator/Masonite-style).
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class Column:
|
|
5
|
+
def __init__(self, name, type, length=None):
|
|
6
|
+
self.name = name
|
|
7
|
+
self.type = type
|
|
8
|
+
self.length = length
|
|
9
|
+
self.is_nullable = False
|
|
10
|
+
self.is_unique = False
|
|
11
|
+
self.is_primary = False
|
|
12
|
+
self.is_index = False
|
|
13
|
+
self.default_value = None
|
|
14
|
+
|
|
15
|
+
def nullable(self):
|
|
16
|
+
self.is_nullable = True
|
|
17
|
+
return self
|
|
18
|
+
|
|
19
|
+
def unique(self):
|
|
20
|
+
self.is_unique = True
|
|
21
|
+
return self
|
|
22
|
+
|
|
23
|
+
def primary(self):
|
|
24
|
+
self.is_primary = True
|
|
25
|
+
return self
|
|
26
|
+
|
|
27
|
+
def index(self):
|
|
28
|
+
self.is_index = True
|
|
29
|
+
return self
|
|
30
|
+
|
|
31
|
+
def default(self, value):
|
|
32
|
+
self.default_value = value
|
|
33
|
+
return self
|