94 lines
2.6 KiB
Django/Jinja
94 lines
2.6 KiB
Django/Jinja
{% if database == 'postgres' -%}
|
|
from collections.abc import AsyncGenerator
|
|
|
|
from dishka import Provider, Scope, provide
|
|
from redis.asyncio import Redis
|
|
|
|
from utils.db.repositories import RedisCache
|
|
from utils.env import env
|
|
|
|
|
|
class RedisProvider(Provider):
|
|
@provide(scope=Scope.APP)
|
|
async def get_redis(self) -> AsyncGenerator[Redis]:
|
|
client = Redis.from_url(env.redis.url, decode_responses=True)
|
|
try:
|
|
yield client
|
|
finally:
|
|
await client.aclose()
|
|
|
|
@provide(scope=Scope.APP)
|
|
def get_cache(self, redis: Redis) -> RedisCache:
|
|
return RedisCache(redis, ttl=env.redis.cache_ttl)
|
|
{%- else -%}
|
|
from collections.abc import AsyncGenerator, Awaitable, Callable
|
|
from typing import TypeVar
|
|
|
|
from dishka import Provider, Scope, provide
|
|
from pydantic import BaseModel
|
|
from redis.asyncio import Redis
|
|
|
|
from utils.env import env
|
|
|
|
ModelT = TypeVar("ModelT", bound=BaseModel)
|
|
NAMESPACE = "{{ project_slug }}"
|
|
|
|
|
|
class RedisCache:
|
|
def __init__(
|
|
self, redis: Redis, *, namespace: str = NAMESPACE, ttl: int = 300
|
|
) -> None:
|
|
self._redis = redis
|
|
self._namespace = namespace
|
|
self._ttl = ttl
|
|
|
|
def _key(self, key: str) -> str:
|
|
return f"{self._namespace}:{key}"
|
|
|
|
async def get_model(self, key: str, model: type[ModelT]) -> ModelT | None:
|
|
raw = await self._redis.get(self._key(key))
|
|
if raw is None:
|
|
return None
|
|
return model.model_validate_json(raw)
|
|
|
|
async def set_model(
|
|
self, key: str, value: BaseModel, *, ttl: int | None = None
|
|
) -> None:
|
|
await self._redis.set(
|
|
self._key(key), value.model_dump_json(), ex=ttl or self._ttl
|
|
)
|
|
|
|
async def get_or_set_model(
|
|
self,
|
|
key: str,
|
|
model: type[ModelT],
|
|
loader: Callable[[], Awaitable[ModelT]],
|
|
*,
|
|
ttl: int | None = None,
|
|
) -> ModelT:
|
|
cached = await self.get_model(key, model)
|
|
if cached is not None:
|
|
return cached
|
|
value = await loader()
|
|
await self.set_model(key, value, ttl=ttl)
|
|
return value
|
|
|
|
async def invalidate(self, *keys: str) -> None:
|
|
if keys:
|
|
await self._redis.delete(*(self._key(key) for key in keys))
|
|
|
|
|
|
class RedisProvider(Provider):
|
|
@provide(scope=Scope.APP)
|
|
async def get_redis(self) -> AsyncGenerator[Redis]:
|
|
client = Redis.from_url(env.redis.url, decode_responses=True)
|
|
try:
|
|
yield client
|
|
finally:
|
|
await client.aclose()
|
|
|
|
@provide(scope=Scope.APP)
|
|
def get_cache(self, redis: Redis) -> RedisCache:
|
|
return RedisCache(redis, ttl=env.redis.cache_ttl)
|
|
{%- endif %}
|