This commit is contained in:
@@ -0,0 +1,246 @@
|
||||
import asyncio
|
||||
import unittest
|
||||
from typing import Any, Optional
|
||||
|
||||
import httpx
|
||||
from pwo import async_test
|
||||
|
||||
from kaya.core import HttpContext, KayaApp
|
||||
from kaya.cors import CorsMixin
|
||||
|
||||
|
||||
def make_app(**cors_kwargs: Any) -> KayaApp:
|
||||
app = KayaApp(mixins=[CorsMixin(**cors_kwargs)])
|
||||
|
||||
@app.GET('/hello')
|
||||
async def hello(ctx: HttpContext) -> None:
|
||||
await ctx.send_str(200, 'Hello World!')
|
||||
|
||||
return app
|
||||
|
||||
|
||||
async def request(app: KayaApp,
|
||||
method: str,
|
||||
path: str = '/hello',
|
||||
headers: Optional[dict[str, str]] = None) -> httpx.Response:
|
||||
transport = httpx.ASGITransport(app=app)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://127.0.0.1:80") as client:
|
||||
return await client.request(method, path, headers=headers)
|
||||
|
||||
|
||||
class CorsSimpleRequestTest(unittest.TestCase):
|
||||
|
||||
@async_test
|
||||
async def test_allowed_origin(self):
|
||||
app = make_app(allow_origins=['https://example.com'])
|
||||
r = await request(app, 'GET', headers={'Origin': 'https://example.com'})
|
||||
self.assertEqual(200, r.status_code)
|
||||
self.assertEqual('Hello World!', r.text)
|
||||
self.assertEqual('https://example.com', r.headers.get('Access-Control-Allow-Origin'))
|
||||
self.assertEqual('Origin', r.headers.get('Vary'))
|
||||
|
||||
@async_test
|
||||
async def test_disallowed_origin(self):
|
||||
app = make_app(allow_origins=['https://example.com'])
|
||||
r = await request(app, 'GET', headers={'Origin': 'https://evil.com'})
|
||||
self.assertEqual(200, r.status_code)
|
||||
self.assertNotIn('Access-Control-Allow-Origin', r.headers)
|
||||
|
||||
@async_test
|
||||
async def test_wildcard_origin(self):
|
||||
app = make_app(allow_origins=['*'])
|
||||
r = await request(app, 'GET', headers={'Origin': 'https://anything.example.com'})
|
||||
self.assertEqual('*', r.headers.get('Access-Control-Allow-Origin'))
|
||||
|
||||
@async_test
|
||||
async def test_wildcard_origin_with_credentials_echoes_origin(self):
|
||||
app = make_app(allow_origins=['*'], allow_credentials=True)
|
||||
r = await request(app, 'GET', headers={'Origin': 'https://example.com'})
|
||||
self.assertEqual('https://example.com', r.headers.get('Access-Control-Allow-Origin'))
|
||||
self.assertEqual('true', r.headers.get('Access-Control-Allow-Credentials'))
|
||||
self.assertEqual('Origin', r.headers.get('Vary'))
|
||||
|
||||
@async_test
|
||||
async def test_origin_regex(self):
|
||||
app = make_app(allow_origin_regex=r'https://.*\.example\.com')
|
||||
r = await request(app, 'GET', headers={'Origin': 'https://api.example.com'})
|
||||
self.assertEqual('https://api.example.com', r.headers.get('Access-Control-Allow-Origin'))
|
||||
r = await request(app, 'GET', headers={'Origin': 'https://example.com.evil.org'})
|
||||
self.assertNotIn('Access-Control-Allow-Origin', r.headers)
|
||||
|
||||
@async_test
|
||||
async def test_no_origin_header(self):
|
||||
app = make_app(allow_origins=['*'])
|
||||
r = await request(app, 'GET')
|
||||
self.assertEqual(200, r.status_code)
|
||||
self.assertNotIn('Access-Control-Allow-Origin', r.headers)
|
||||
|
||||
@async_test
|
||||
async def test_expose_headers(self):
|
||||
app = make_app(allow_origins=['*'], expose_headers=['X-Total-Count'])
|
||||
r = await request(app, 'GET', headers={'Origin': 'https://example.com'})
|
||||
self.assertEqual('X-Total-Count', r.headers.get('Access-Control-Expose-Headers'))
|
||||
|
||||
@async_test
|
||||
async def test_handler_set_cors_header_not_overwritten(self):
|
||||
app = KayaApp(mixins=[CorsMixin(allow_origins=['*'])])
|
||||
|
||||
@app.GET('/custom')
|
||||
async def custom(ctx: HttpContext) -> None:
|
||||
await ctx.send_str(200, 'custom', headers={'Access-Control-Allow-Origin': 'https://custom.example.com'})
|
||||
|
||||
r = await request(app, 'GET', '/custom', headers={'Origin': 'https://example.com'})
|
||||
self.assertEqual('https://custom.example.com', r.headers.get('Access-Control-Allow-Origin'))
|
||||
|
||||
|
||||
class CorsPreflightTest(unittest.TestCase):
|
||||
|
||||
@staticmethod
|
||||
def preflight_headers(origin: str = 'https://example.com',
|
||||
method: str = 'POST',
|
||||
headers: Optional[str] = None) -> dict[str, str]:
|
||||
result = {
|
||||
'Origin': origin,
|
||||
'Access-Control-Request-Method': method,
|
||||
}
|
||||
if headers is not None:
|
||||
result['Access-Control-Request-Headers'] = headers
|
||||
return result
|
||||
|
||||
@async_test
|
||||
async def test_preflight_allowed(self):
|
||||
app = make_app(allow_origins=['https://example.com'],
|
||||
allow_methods=('GET', 'POST'),
|
||||
allow_credentials=True)
|
||||
r = await request(app, 'OPTIONS', headers=self.preflight_headers())
|
||||
self.assertEqual(200, r.status_code)
|
||||
self.assertEqual('OK', r.text)
|
||||
self.assertEqual('https://example.com', r.headers.get('Access-Control-Allow-Origin'))
|
||||
self.assertEqual('GET, POST', r.headers.get('Access-Control-Allow-Methods'))
|
||||
self.assertEqual('600', r.headers.get('Access-Control-Max-Age'))
|
||||
self.assertEqual('true', r.headers.get('Access-Control-Allow-Credentials'))
|
||||
self.assertEqual('Origin', r.headers.get('Vary'))
|
||||
|
||||
@async_test
|
||||
async def test_preflight_wildcard_origin(self):
|
||||
app = make_app(allow_origins=['*'], allow_methods=('GET', 'POST'))
|
||||
r = await request(app, 'OPTIONS', headers=self.preflight_headers())
|
||||
self.assertEqual(200, r.status_code)
|
||||
self.assertEqual('*', r.headers.get('Access-Control-Allow-Origin'))
|
||||
|
||||
@async_test
|
||||
async def test_preflight_disallowed_origin(self):
|
||||
app = make_app(allow_origins=['https://example.com'], allow_methods=('GET', 'POST'))
|
||||
r = await request(app, 'OPTIONS', headers=self.preflight_headers(origin='https://evil.com'))
|
||||
self.assertEqual(400, r.status_code)
|
||||
self.assertEqual('Disallowed CORS origin', r.text)
|
||||
|
||||
@async_test
|
||||
async def test_preflight_disallowed_method(self):
|
||||
app = make_app(allow_origins=['https://example.com'], allow_methods=('GET',))
|
||||
r = await request(app, 'OPTIONS', headers=self.preflight_headers(method='DELETE'))
|
||||
self.assertEqual(400, r.status_code)
|
||||
self.assertEqual('Disallowed CORS method', r.text)
|
||||
|
||||
@async_test
|
||||
async def test_preflight_disallowed_headers(self):
|
||||
app = make_app(allow_origins=['https://example.com'], allow_methods=('GET', 'POST'))
|
||||
r = await request(app, 'OPTIONS', headers=self.preflight_headers(headers='X-Custom'))
|
||||
self.assertEqual(400, r.status_code)
|
||||
self.assertEqual('Disallowed CORS headers', r.text)
|
||||
|
||||
@async_test
|
||||
async def test_preflight_safelisted_headers_allowed(self):
|
||||
app = make_app(allow_origins=['https://example.com'], allow_methods=('GET', 'POST'))
|
||||
r = await request(app, 'OPTIONS', headers=self.preflight_headers(headers='Content-Type'))
|
||||
self.assertEqual(200, r.status_code)
|
||||
|
||||
@async_test
|
||||
async def test_preflight_allow_all_headers_mirrors_request(self):
|
||||
app = make_app(allow_origins=['*'], allow_methods=('GET', 'POST'), allow_headers=['*'])
|
||||
r = await request(app, 'OPTIONS', headers=self.preflight_headers(headers='X-Custom, X-Other'))
|
||||
self.assertEqual(200, r.status_code)
|
||||
self.assertEqual('X-Custom, X-Other', r.headers.get('Access-Control-Allow-Headers'))
|
||||
|
||||
@async_test
|
||||
async def test_preflight_configured_allow_headers(self):
|
||||
app = make_app(allow_origins=['https://example.com'],
|
||||
allow_methods=('GET', 'POST'),
|
||||
allow_headers=('X-Custom',))
|
||||
r = await request(app, 'OPTIONS', headers=self.preflight_headers(headers='X-Custom'))
|
||||
self.assertEqual(200, r.status_code)
|
||||
allow_headers = r.headers.get('Access-Control-Allow-Headers')
|
||||
self.assertIsNotNone(allow_headers)
|
||||
assert allow_headers is not None
|
||||
self.assertIn('X-Custom', allow_headers)
|
||||
|
||||
@async_test
|
||||
async def test_preflight_response_not_overwritten_by_handler(self):
|
||||
# even when a user-registered OPTIONS handler matches, the preflight
|
||||
# response sent by the mixin is the only one delivered to the client
|
||||
app = KayaApp(mixins=[CorsMixin(allow_origins=['https://example.com'], allow_methods=('GET', 'POST'))])
|
||||
|
||||
@app.GET('/hello')
|
||||
async def hello(ctx: HttpContext) -> None:
|
||||
await ctx.send_str(200, 'Hello World!')
|
||||
|
||||
@app.OPTIONS('/hello')
|
||||
async def options(ctx: HttpContext) -> None:
|
||||
await ctx.send_str(200, 'custom OPTIONS handler')
|
||||
|
||||
r = await request(app, 'OPTIONS', headers=self.preflight_headers())
|
||||
self.assertEqual(200, r.status_code)
|
||||
self.assertEqual('OK', r.text)
|
||||
|
||||
@async_test
|
||||
async def test_options_without_preflight_headers_routes_normally(self):
|
||||
app = KayaApp(mixins=[CorsMixin(allow_origins=['https://example.com'])])
|
||||
|
||||
@app.OPTIONS('/hello')
|
||||
async def options(ctx: HttpContext) -> None:
|
||||
await ctx.send_str(200, 'custom OPTIONS handler')
|
||||
|
||||
# an OPTIONS request without Access-Control-Request-Method is not a
|
||||
# preflight request and is routed normally
|
||||
r = await request(app, 'OPTIONS', headers={'Origin': 'https://example.com'})
|
||||
self.assertEqual(200, r.status_code)
|
||||
self.assertEqual('custom OPTIONS handler', r.text)
|
||||
self.assertEqual('https://example.com', r.headers.get('Access-Control-Allow-Origin'))
|
||||
|
||||
|
||||
class CorsRsgiTest(unittest.TestCase):
|
||||
|
||||
def test_rsgi_context_header_injection(self):
|
||||
from kaya.rsgi import RsgiContext
|
||||
|
||||
class FakeScope:
|
||||
scheme = 'http'
|
||||
method = 'GET'
|
||||
path = '/'
|
||||
query_string = ''
|
||||
headers = {'origin': 'https://example.com'}
|
||||
client = '127.0.0.1:12345'
|
||||
server = '127.0.0.1:80'
|
||||
|
||||
class FakeProtocol:
|
||||
def __init__(self) -> None:
|
||||
self.responses = []
|
||||
|
||||
def response_str(self, status: int, headers: list, body: str) -> None:
|
||||
self.responses.append((status, dict(headers), body))
|
||||
|
||||
mixin = CorsMixin(allow_origins=['https://example.com'])
|
||||
protocol = FakeProtocol()
|
||||
ctx = RsgiContext(FakeScope(), protocol) # type: ignore[arg-type]
|
||||
|
||||
async def run() -> None:
|
||||
wrapped = await mixin._before_request(ctx)
|
||||
assert wrapped is not None
|
||||
await wrapped.send_str(200, 'hi')
|
||||
|
||||
asyncio.run(run())
|
||||
self.assertEqual(1, len(protocol.responses))
|
||||
status, headers, body = protocol.responses[0]
|
||||
self.assertEqual(200, status)
|
||||
self.assertEqual('https://example.com', headers.get('Access-Control-Allow-Origin'))
|
||||
self.assertEqual('Origin', headers.get('Vary'))
|
||||
Reference in New Issue
Block a user