relayers
This commit is contained in:
1 parent
21964fa986
commit
797f379648
68 files changed
+23871
-1271
No files matched your search
@@ -1,267 +1,602 @@
|
||||
"""Media conversion service for processing uploaded files."""
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
import os
|
||||
import uuid
|
||||
import hashlib
|
||||
import json
|
||||
import shutil
|
||||
import magic # python-magic for MIME detection
|
||||
from base58 import b58decode, b58encode
|
||||
from sqlalchemy import and_, or_
|
||||
from app.core.models.node_storage import StoredContent
|
||||
from app.core.models._telegram import Wrapped_CBotChat
|
||||
from app.core._utils.send_status import send_status
|
||||
from app.core.logger import make_log
|
||||
from app.core.models.user import User
|
||||
from app.core.models import WalletConnection
|
||||
from app.core.storage import db_session
|
||||
from app.core._config import UPLOADS_DIR
|
||||
from app.core.content.content_id import ContentId
|
||||
import logging
|
||||
import os
|
||||
import tempfile
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Optional, Set, Any, Tuple
|
||||
|
||||
import aiofiles
|
||||
import redis.asyncio as redis
|
||||
from PIL import Image, ImageOps
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_async_session
|
||||
from app.core.models.content import Content, FileUpload
|
||||
from app.core.storage import storage_manager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def convert_loop(memory):
|
||||
with db_session() as session:
|
||||
# Query for unprocessed encrypted content
|
||||
unprocessed_encrypted_content = session.query(StoredContent).filter(
|
||||
and_(
|
||||
StoredContent.type == "onchain/content",
|
||||
or_(
|
||||
StoredContent.btfs_cid == None,
|
||||
StoredContent.ipfs_cid == None,
|
||||
)
|
||||
)
|
||||
).first()
|
||||
if not unprocessed_encrypted_content:
|
||||
make_log("ConvertProcess", "No content to convert", level="debug")
|
||||
return
|
||||
class ConvertService:
|
||||
"""Service for converting and processing uploaded media files."""
|
||||
|
||||
# Достаем расшифрованный файл
|
||||
decrypted_content = session.query(StoredContent).filter(
|
||||
StoredContent.id == unprocessed_encrypted_content.decrypted_content_id
|
||||
).first()
|
||||
if not decrypted_content:
|
||||
make_log("ConvertProcess", "Decrypted content not found", level="error")
|
||||
return
|
||||
|
||||
# Определяем путь и расширение входного файла
|
||||
input_file_path = f"/Storage/storedContent/{decrypted_content.hash}"
|
||||
input_ext = (unprocessed_encrypted_content.filename.split('.')[-1]
|
||||
if '.' in unprocessed_encrypted_content.filename else "mp4")
|
||||
|
||||
# ==== Новая логика: определение MIME-тип через python-magic ====
|
||||
def __init__(self):
|
||||
self.settings = get_settings()
|
||||
self.redis_client: Optional[redis.Redis] = None
|
||||
self.is_running = False
|
||||
self.tasks: Set[asyncio.Task] = set()
|
||||
|
||||
# Conversion configuration
|
||||
self.batch_size = 10
|
||||
self.process_interval = 5 # seconds
|
||||
self.max_retries = 3
|
||||
self.temp_dir = Path(tempfile.gettempdir()) / "uploader_convert"
|
||||
self.temp_dir.mkdir(exist_ok=True)
|
||||
|
||||
# Supported formats
|
||||
self.image_formats = {'.jpg', '.jpeg', '.png', '.gif', '.bmp', '.webp', '.tiff'}
|
||||
self.video_formats = {'.mp4', '.avi', '.mov', '.wmv', '.flv', '.webm', '.mkv'}
|
||||
self.audio_formats = {'.mp3', '.wav', '.ogg', '.m4a', '.flac', '.aac'}
|
||||
self.document_formats = {'.pdf', '.doc', '.docx', '.txt', '.rtf'}
|
||||
|
||||
# Image processing settings
|
||||
self.thumbnail_sizes = [(150, 150), (300, 300), (800, 600)]
|
||||
self.image_quality = 85
|
||||
self.max_image_size = (2048, 2048)
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Start the conversion service."""
|
||||
try:
|
||||
mime_type = magic.from_file(input_file_path.replace("/Storage/storedContent", "/app/data"), mime=True)
|
||||
except Exception as e:
|
||||
make_log("ConvertProcess", f"magic probe failed: {e}", level="warning")
|
||||
mime_type = ""
|
||||
|
||||
if mime_type.startswith("video/"):
|
||||
content_kind = "video"
|
||||
elif mime_type.startswith("audio/"):
|
||||
content_kind = "audio"
|
||||
else:
|
||||
content_kind = "other"
|
||||
|
||||
make_log("ConvertProcess", f"Detected content_kind={content_kind}, mime={mime_type}", level="info")
|
||||
|
||||
# Для прочих типов сохраняем raw копию и выходим
|
||||
if content_kind == "other":
|
||||
make_log("ConvertProcess", f"Content {unprocessed_encrypted_content.id} processed. Not audio/video, copy just", level="info")
|
||||
unprocessed_encrypted_content.btfs_cid = ContentId(
|
||||
version=2, content_hash=b58decode(decrypted_content.hash)
|
||||
).serialize_v2()
|
||||
unprocessed_encrypted_content.ipfs_cid = ContentId(
|
||||
version=2, content_hash=b58decode(decrypted_content.hash)
|
||||
).serialize_v2()
|
||||
unprocessed_encrypted_content.meta = {
|
||||
**unprocessed_encrypted_content.meta,
|
||||
'converted_content': {
|
||||
option_name: decrypted_content.hash for option_name in ['high', 'low', 'low_preview']
|
||||
}
|
||||
}
|
||||
session.commit()
|
||||
return
|
||||
|
||||
# ==== Конвертация для видео или аудио: оригинальная логика ====
|
||||
# Static preview interval in seconds
|
||||
preview_interval = [0, 30]
|
||||
if unprocessed_encrypted_content.onchain_index in [2]:
|
||||
preview_interval = [0, 60]
|
||||
|
||||
make_log(
|
||||
"ConvertProcess",
|
||||
f"Processing content {unprocessed_encrypted_content.id} as {content_kind} with preview interval {preview_interval}",
|
||||
level="info"
|
||||
)
|
||||
|
||||
# Выбираем опции конвертации для видео и аудио
|
||||
if content_kind == "video":
|
||||
REQUIRED_CONVERT_OPTIONS = ['high', 'low', 'low_preview']
|
||||
else:
|
||||
REQUIRED_CONVERT_OPTIONS = ['high', 'low'] # no preview for audio
|
||||
|
||||
converted_content = {}
|
||||
logs_dir = "/Storage/logs/converter"
|
||||
|
||||
for option in REQUIRED_CONVERT_OPTIONS:
|
||||
# Set quality parameter and trim option (only for preview)
|
||||
if option == "low_preview":
|
||||
quality = "low"
|
||||
trim_value = f"{preview_interval[0]}-{preview_interval[1]}"
|
||||
else:
|
||||
quality = option
|
||||
trim_value = None
|
||||
|
||||
# Generate a unique output directory for docker container
|
||||
output_uuid = str(uuid.uuid4())
|
||||
output_dir = f"/Storage/storedContent/converter-output/{output_uuid}"
|
||||
|
||||
# Build the docker command
|
||||
cmd = [
|
||||
"docker", "run", "--rm",
|
||||
"-v", f"{input_file_path}:/app/input",
|
||||
"-v", f"{output_dir}:/app/output",
|
||||
"-v", f"{logs_dir}:/app/logs",
|
||||
"media_converter",
|
||||
"--ext", input_ext,
|
||||
"--quality", quality
|
||||
logger.info("Starting media conversion service")
|
||||
|
||||
# Initialize Redis connection
|
||||
self.redis_client = redis.from_url(
|
||||
self.settings.redis_url,
|
||||
encoding="utf-8",
|
||||
decode_responses=True,
|
||||
socket_keepalive=True,
|
||||
socket_keepalive_options={},
|
||||
health_check_interval=30,
|
||||
)
|
||||
|
||||
# Test Redis connection
|
||||
await self.redis_client.ping()
|
||||
logger.info("Redis connection established for converter")
|
||||
|
||||
# Start conversion tasks
|
||||
self.is_running = True
|
||||
|
||||
# Create conversion tasks
|
||||
tasks = [
|
||||
asyncio.create_task(self._process_pending_files_loop()),
|
||||
asyncio.create_task(self._cleanup_temp_files_loop()),
|
||||
asyncio.create_task(self._retry_failed_conversions_loop()),
|
||||
]
|
||||
if trim_value:
|
||||
cmd.extend(["--trim", trim_value])
|
||||
if content_kind == "audio":
|
||||
cmd.append("--audio-only") # audio-only flag
|
||||
|
||||
process = await asyncio.create_subprocess_exec(
|
||||
*cmd,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE
|
||||
)
|
||||
stdout, stderr = await process.communicate()
|
||||
if process.returncode != 0:
|
||||
make_log("ConvertProcess", f"Docker conversion failed for option {option}: {stderr.decode()}", level="error")
|
||||
return
|
||||
|
||||
# List files in output dir
|
||||
|
||||
self.tasks.update(tasks)
|
||||
|
||||
# Wait for all tasks
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error starting conversion service: {e}")
|
||||
await self.stop()
|
||||
raise
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the conversion service."""
|
||||
logger.info("Stopping media conversion service")
|
||||
self.is_running = False
|
||||
|
||||
# Cancel all tasks
|
||||
for task in self.tasks:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
|
||||
# Wait for tasks to complete
|
||||
if self.tasks:
|
||||
await asyncio.gather(*self.tasks, return_exceptions=True)
|
||||
|
||||
# Close Redis connection
|
||||
if self.redis_client:
|
||||
await self.redis_client.close()
|
||||
|
||||
# Cleanup temp directory
|
||||
await self._cleanup_temp_directory()
|
||||
|
||||
logger.info("Conversion service stopped")
|
||||
|
||||
async def _process_pending_files_loop(self) -> None:
|
||||
"""Main loop for processing pending file conversions."""
|
||||
logger.info("Starting file conversion loop")
|
||||
|
||||
while self.is_running:
|
||||
try:
|
||||
files = os.listdir(output_dir.replace("/Storage/storedContent", "/app/data"))
|
||||
await self._process_pending_files()
|
||||
await asyncio.sleep(self.process_interval)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
make_log("ConvertProcess", f"Error reading output directory {output_dir}: {e}", level="error")
|
||||
return
|
||||
|
||||
media_files = [f for f in files if f != "output.json"]
|
||||
if len(media_files) != 1:
|
||||
make_log("ConvertProcess", f"Expected one media file, found {len(media_files)} for option {option}", level="error")
|
||||
return
|
||||
|
||||
output_file = os.path.join(
|
||||
output_dir.replace("/Storage/storedContent", "/app/data"),
|
||||
media_files[0]
|
||||
)
|
||||
|
||||
# Compute SHA256 hash of the output file
|
||||
hash_process = await asyncio.create_subprocess_exec(
|
||||
"sha256sum", output_file,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE
|
||||
)
|
||||
hash_stdout, hash_stderr = await hash_process.communicate()
|
||||
if hash_process.returncode != 0:
|
||||
make_log("ConvertProcess", f"Error computing sha256sum for option {option}: {hash_stderr.decode()}", level="error")
|
||||
return
|
||||
file_hash = hash_stdout.decode().split()[0]
|
||||
file_hash = b58encode(bytes.fromhex(file_hash)).decode()
|
||||
|
||||
# Save new StoredContent if not exists
|
||||
if not session.query(StoredContent).filter(
|
||||
StoredContent.hash == file_hash
|
||||
).first():
|
||||
new_content = StoredContent(
|
||||
type="local/content_bin",
|
||||
hash=file_hash,
|
||||
user_id=unprocessed_encrypted_content.user_id,
|
||||
filename=media_files[0],
|
||||
meta={'encrypted_file_hash': unprocessed_encrypted_content.hash},
|
||||
created=datetime.now(),
|
||||
logger.error(f"Error in file conversion loop: {e}")
|
||||
await asyncio.sleep(self.process_interval)
|
||||
|
||||
async def _process_pending_files(self) -> None:
|
||||
"""Process pending file conversions."""
|
||||
async with get_async_session() as session:
|
||||
try:
|
||||
# Get pending uploads
|
||||
result = await session.execute(
|
||||
select(FileUpload)
|
||||
.where(
|
||||
FileUpload.status == "uploaded",
|
||||
FileUpload.processed == False
|
||||
)
|
||||
.limit(self.batch_size)
|
||||
)
|
||||
session.add(new_content)
|
||||
session.commit()
|
||||
|
||||
save_path = os.path.join(UPLOADS_DIR, file_hash)
|
||||
try:
|
||||
os.remove(save_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
try:
|
||||
shutil.move(output_file, save_path)
|
||||
uploads = result.scalars().all()
|
||||
|
||||
if not uploads:
|
||||
return
|
||||
|
||||
logger.info(f"Processing {len(uploads)} pending files")
|
||||
|
||||
# Process each upload
|
||||
for upload in uploads:
|
||||
await self._process_single_file(session, upload)
|
||||
|
||||
await session.commit()
|
||||
|
||||
except Exception as e:
|
||||
make_log("ConvertProcess", f"Error moving output file {output_file} to {save_path}: {e}", level="error")
|
||||
return
|
||||
|
||||
converted_content[option] = file_hash
|
||||
|
||||
# Process output.json for ffprobe_meta
|
||||
output_json_path = os.path.join(
|
||||
output_dir.replace("/Storage/storedContent", "/app/data"),
|
||||
"output.json"
|
||||
)
|
||||
if os.path.exists(output_json_path) and unprocessed_encrypted_content.meta.get('ffprobe_meta') is None:
|
||||
try:
|
||||
with open(output_json_path, "r") as f:
|
||||
ffprobe_meta = json.load(f)
|
||||
unprocessed_encrypted_content.meta = {
|
||||
**unprocessed_encrypted_content.meta,
|
||||
'ffprobe_meta': ffprobe_meta
|
||||
}
|
||||
except Exception as e:
|
||||
make_log("ConvertProcess", f"Error handling output.json for option {option}: {e}", level="error")
|
||||
|
||||
# Cleanup output directory
|
||||
try:
|
||||
shutil.rmtree(output_dir.replace("/Storage/storedContent", "/app/data"))
|
||||
except Exception as e:
|
||||
make_log("ConvertProcess", f"Error removing output dir {output_dir}: {e}", level="warning")
|
||||
|
||||
# Finalize original record
|
||||
make_log("ConvertProcess", f"Content {unprocessed_encrypted_content.id} processed. Converted content: {converted_content}", level="info")
|
||||
unprocessed_encrypted_content.btfs_cid = ContentId(
|
||||
version=2, content_hash=b58decode(converted_content['high' if content_kind=='video' else 'low'])
|
||||
).serialize_v2()
|
||||
unprocessed_encrypted_content.ipfs_cid = ContentId(
|
||||
version=2, content_hash=b58decode(converted_content['low'])
|
||||
).serialize_v2()
|
||||
unprocessed_encrypted_content.meta = {
|
||||
**unprocessed_encrypted_content.meta,
|
||||
'converted_content': converted_content
|
||||
}
|
||||
session.commit()
|
||||
|
||||
# Notify user if needed
|
||||
if not unprocessed_encrypted_content.meta.get('upload_notify_msg_id'):
|
||||
wallet_owner_connection = session.query(WalletConnection).filter(
|
||||
WalletConnection.wallet_address == unprocessed_encrypted_content.owner_address
|
||||
).order_by(WalletConnection.id.desc()).first()
|
||||
if wallet_owner_connection:
|
||||
wallet_owner_user = wallet_owner_connection.user
|
||||
bot = Wrapped_CBotChat(
|
||||
memory._client_telegram_bot,
|
||||
chat_id=wallet_owner_user.telegram_id,
|
||||
user=wallet_owner_user,
|
||||
db_session=session
|
||||
)
|
||||
unprocessed_encrypted_content.meta['upload_notify_msg_id'] = await bot.send_content(session, unprocessed_encrypted_content)
|
||||
session.commit()
|
||||
|
||||
|
||||
async def main_fn(memory):
|
||||
make_log("ConvertProcess", "Service started", level="info")
|
||||
seqno = 0
|
||||
while True:
|
||||
logger.error(f"Error processing pending files: {e}")
|
||||
await session.rollback()
|
||||
|
||||
async def _process_single_file(self, session: AsyncSession, upload: FileUpload) -> None:
|
||||
"""Process a single file upload."""
|
||||
try:
|
||||
make_log("ConvertProcess", "Service running", level="debug")
|
||||
await convert_loop(memory)
|
||||
await asyncio.sleep(5)
|
||||
await send_status("convert_service", f"working (seqno={seqno})")
|
||||
seqno += 1
|
||||
except BaseException as e:
|
||||
make_log("ConvertProcess", f"Error: {e}", level="error")
|
||||
await asyncio.sleep(3)
|
||||
logger.info(f"Processing file: {upload.filename}")
|
||||
|
||||
# Mark as processing
|
||||
upload.status = "processing"
|
||||
upload.processing_started_at = datetime.utcnow()
|
||||
await session.commit()
|
||||
|
||||
# Get file extension
|
||||
file_ext = Path(upload.filename).suffix.lower()
|
||||
|
||||
# Process based on file type
|
||||
if file_ext in self.image_formats:
|
||||
await self._process_image(session, upload)
|
||||
elif file_ext in self.video_formats:
|
||||
await self._process_video(session, upload)
|
||||
elif file_ext in self.audio_formats:
|
||||
await self._process_audio(session, upload)
|
||||
elif file_ext in self.document_formats:
|
||||
await self._process_document(session, upload)
|
||||
else:
|
||||
# Just mark as processed for unsupported formats
|
||||
upload.status = "completed"
|
||||
upload.processed = True
|
||||
upload.processing_completed_at = datetime.utcnow()
|
||||
|
||||
# Cache processing result
|
||||
cache_key = f"processed:{upload.id}"
|
||||
processing_info = {
|
||||
"status": upload.status,
|
||||
"processed_at": datetime.utcnow().isoformat(),
|
||||
"metadata": upload.metadata or {}
|
||||
}
|
||||
await self.redis_client.setex(cache_key, 3600, json.dumps(processing_info))
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing file {upload.filename}: {e}")
|
||||
|
||||
# Mark as failed
|
||||
upload.status = "failed"
|
||||
upload.error_message = str(e)
|
||||
upload.retry_count = (upload.retry_count or 0) + 1
|
||||
|
||||
if upload.retry_count >= self.max_retries:
|
||||
upload.processed = True # Stop retrying
|
||||
|
||||
async def _process_image(self, session: AsyncSession, upload: FileUpload) -> None:
|
||||
"""Process an image file."""
|
||||
try:
|
||||
# Download original file
|
||||
original_path = await self._download_file(upload)
|
||||
|
||||
if not original_path:
|
||||
raise Exception("Failed to download original file")
|
||||
|
||||
# Open image
|
||||
with Image.open(original_path) as img:
|
||||
# Extract metadata
|
||||
metadata = {
|
||||
"format": img.format,
|
||||
"mode": img.mode,
|
||||
"size": img.size,
|
||||
"has_transparency": img.mode in ('RGBA', 'LA') or 'transparency' in img.info
|
||||
}
|
||||
|
||||
# Fix orientation
|
||||
img = ImageOps.exif_transpose(img)
|
||||
|
||||
# Resize if too large
|
||||
if img.size[0] > self.max_image_size[0] or img.size[1] > self.max_image_size[1]:
|
||||
img.thumbnail(self.max_image_size, Image.Resampling.LANCZOS)
|
||||
metadata["resized"] = True
|
||||
|
||||
# Save optimized version
|
||||
optimized_path = self.temp_dir / f"optimized_{upload.id}.jpg"
|
||||
|
||||
# Convert to RGB if necessary
|
||||
if img.mode in ('RGBA', 'LA'):
|
||||
background = Image.new('RGB', img.size, (255, 255, 255))
|
||||
if img.mode == 'LA':
|
||||
img = img.convert('RGBA')
|
||||
background.paste(img, mask=img.split()[-1])
|
||||
img = background
|
||||
elif img.mode != 'RGB':
|
||||
img = img.convert('RGB')
|
||||
|
||||
img.save(
|
||||
optimized_path,
|
||||
'JPEG',
|
||||
quality=self.image_quality,
|
||||
optimize=True
|
||||
)
|
||||
|
||||
# Upload optimized version
|
||||
optimized_url = await storage_manager.upload_file(
|
||||
str(optimized_path),
|
||||
f"optimized/{upload.id}/image.jpg"
|
||||
)
|
||||
|
||||
# Generate thumbnails
|
||||
thumbnails = {}
|
||||
for size in self.thumbnail_sizes:
|
||||
thumbnail_path = await self._create_thumbnail(original_path, size)
|
||||
if thumbnail_path:
|
||||
thumb_url = await storage_manager.upload_file(
|
||||
str(thumbnail_path),
|
||||
f"thumbnails/{upload.id}/{size[0]}x{size[1]}.jpg"
|
||||
)
|
||||
thumbnails[f"{size[0]}x{size[1]}"] = thumb_url
|
||||
thumbnail_path.unlink() # Cleanup
|
||||
|
||||
# Update upload record
|
||||
upload.metadata = {
|
||||
**metadata,
|
||||
"thumbnails": thumbnails,
|
||||
"optimized_url": optimized_url
|
||||
}
|
||||
upload.status = "completed"
|
||||
upload.processed = True
|
||||
upload.processing_completed_at = datetime.utcnow()
|
||||
|
||||
# Cleanup temp files
|
||||
original_path.unlink()
|
||||
optimized_path.unlink()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing image {upload.filename}: {e}")
|
||||
raise
|
||||
|
||||
async def _process_video(self, session: AsyncSession, upload: FileUpload) -> None:
|
||||
"""Process a video file."""
|
||||
try:
|
||||
# For video processing, we would typically use ffmpeg
|
||||
# This is a simplified version that just extracts basic info
|
||||
|
||||
original_path = await self._download_file(upload)
|
||||
if not original_path:
|
||||
raise Exception("Failed to download original file")
|
||||
|
||||
# Basic video metadata (would use ffprobe in real implementation)
|
||||
metadata = {
|
||||
"type": "video",
|
||||
"file_size": original_path.stat().st_size,
|
||||
"processing_note": "Video processing requires ffmpeg implementation"
|
||||
}
|
||||
|
||||
# Generate video thumbnail (simplified)
|
||||
thumbnail_path = await self._create_video_thumbnail(original_path)
|
||||
if thumbnail_path:
|
||||
thumb_url = await storage_manager.upload_file(
|
||||
str(thumbnail_path),
|
||||
f"thumbnails/{upload.id}/video_thumb.jpg"
|
||||
)
|
||||
metadata["thumbnail"] = thumb_url
|
||||
thumbnail_path.unlink()
|
||||
|
||||
# Update upload record
|
||||
upload.metadata = metadata
|
||||
upload.status = "completed"
|
||||
upload.processed = True
|
||||
upload.processing_completed_at = datetime.utcnow()
|
||||
|
||||
# Cleanup
|
||||
original_path.unlink()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing video {upload.filename}: {e}")
|
||||
raise
|
||||
|
||||
async def _process_audio(self, session: AsyncSession, upload: FileUpload) -> None:
|
||||
"""Process an audio file."""
|
||||
try:
|
||||
original_path = await self._download_file(upload)
|
||||
if not original_path:
|
||||
raise Exception("Failed to download original file")
|
||||
|
||||
# Basic audio metadata
|
||||
metadata = {
|
||||
"type": "audio",
|
||||
"file_size": original_path.stat().st_size,
|
||||
"processing_note": "Audio processing requires additional libraries"
|
||||
}
|
||||
|
||||
# Update upload record
|
||||
upload.metadata = metadata
|
||||
upload.status = "completed"
|
||||
upload.processed = True
|
||||
upload.processing_completed_at = datetime.utcnow()
|
||||
|
||||
# Cleanup
|
||||
original_path.unlink()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing audio {upload.filename}: {e}")
|
||||
raise
|
||||
|
||||
async def _process_document(self, session: AsyncSession, upload: FileUpload) -> None:
|
||||
"""Process a document file."""
|
||||
try:
|
||||
original_path = await self._download_file(upload)
|
||||
if not original_path:
|
||||
raise Exception("Failed to download original file")
|
||||
|
||||
# Basic document metadata
|
||||
metadata = {
|
||||
"type": "document",
|
||||
"file_size": original_path.stat().st_size,
|
||||
"pages": 1, # Would extract actual page count for PDFs
|
||||
"processing_note": "Document processing requires additional libraries"
|
||||
}
|
||||
|
||||
# Update upload record
|
||||
upload.metadata = metadata
|
||||
upload.status = "completed"
|
||||
upload.processed = True
|
||||
upload.processing_completed_at = datetime.utcnow()
|
||||
|
||||
# Cleanup
|
||||
original_path.unlink()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing document {upload.filename}: {e}")
|
||||
raise
|
||||
|
||||
async def _download_file(self, upload: FileUpload) -> Optional[Path]:
|
||||
"""Download a file for processing."""
|
||||
try:
|
||||
if not upload.file_path:
|
||||
return None
|
||||
|
||||
# Create temp file path
|
||||
temp_path = self.temp_dir / f"original_{upload.id}_{upload.filename}"
|
||||
|
||||
# Download file from storage
|
||||
file_data = await storage_manager.get_file(upload.file_path)
|
||||
if not file_data:
|
||||
return None
|
||||
|
||||
# Write to temp file
|
||||
async with aiofiles.open(temp_path, 'wb') as f:
|
||||
await f.write(file_data)
|
||||
|
||||
return temp_path
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error downloading file {upload.filename}: {e}")
|
||||
return None
|
||||
|
||||
async def _create_thumbnail(self, image_path: Path, size: Tuple[int, int]) -> Optional[Path]:
|
||||
"""Create a thumbnail from an image."""
|
||||
try:
|
||||
thumbnail_path = self.temp_dir / f"thumb_{size[0]}x{size[1]}_{image_path.name}"
|
||||
|
||||
with Image.open(image_path) as img:
|
||||
# Fix orientation
|
||||
img = ImageOps.exif_transpose(img)
|
||||
|
||||
# Create thumbnail
|
||||
img.thumbnail(size, Image.Resampling.LANCZOS)
|
||||
|
||||
# Convert to RGB if necessary
|
||||
if img.mode in ('RGBA', 'LA'):
|
||||
background = Image.new('RGB', img.size, (255, 255, 255))
|
||||
if img.mode == 'LA':
|
||||
img = img.convert('RGBA')
|
||||
background.paste(img, mask=img.split()[-1])
|
||||
img = background
|
||||
elif img.mode != 'RGB':
|
||||
img = img.convert('RGB')
|
||||
|
||||
# Save thumbnail
|
||||
img.save(
|
||||
thumbnail_path,
|
||||
'JPEG',
|
||||
quality=self.image_quality,
|
||||
optimize=True
|
||||
)
|
||||
|
||||
return thumbnail_path
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating thumbnail: {e}")
|
||||
return None
|
||||
|
||||
async def _create_video_thumbnail(self, video_path: Path) -> Optional[Path]:
|
||||
"""Create a thumbnail from a video file."""
|
||||
try:
|
||||
# This would require ffmpeg to extract a frame from the video
|
||||
# For now, return a placeholder
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating video thumbnail: {e}")
|
||||
return None
|
||||
|
||||
async def _cleanup_temp_files_loop(self) -> None:
|
||||
"""Loop for cleaning up temporary files."""
|
||||
logger.info("Starting temp file cleanup loop")
|
||||
|
||||
while self.is_running:
|
||||
try:
|
||||
await self._cleanup_old_temp_files()
|
||||
await asyncio.sleep(3600) # Run every hour
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error in temp cleanup loop: {e}")
|
||||
await asyncio.sleep(3600)
|
||||
|
||||
async def _cleanup_old_temp_files(self) -> None:
|
||||
"""Clean up old temporary files."""
|
||||
try:
|
||||
current_time = datetime.now().timestamp()
|
||||
|
||||
for file_path in self.temp_dir.glob("*"):
|
||||
if file_path.is_file():
|
||||
# Remove files older than 1 hour
|
||||
if current_time - file_path.stat().st_mtime > 3600:
|
||||
file_path.unlink()
|
||||
logger.debug(f"Removed old temp file: {file_path}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error cleaning up temp files: {e}")
|
||||
|
||||
async def _cleanup_temp_directory(self) -> None:
|
||||
"""Clean up the entire temp directory."""
|
||||
try:
|
||||
for file_path in self.temp_dir.glob("*"):
|
||||
if file_path.is_file():
|
||||
file_path.unlink()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error cleaning up temp directory: {e}")
|
||||
|
||||
async def _retry_failed_conversions_loop(self) -> None:
|
||||
"""Loop for retrying failed conversions."""
|
||||
logger.info("Starting retry loop for failed conversions")
|
||||
|
||||
while self.is_running:
|
||||
try:
|
||||
await self._retry_failed_conversions()
|
||||
await asyncio.sleep(1800) # Run every 30 minutes
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error in retry loop: {e}")
|
||||
await asyncio.sleep(1800)
|
||||
|
||||
async def _retry_failed_conversions(self) -> None:
|
||||
"""Retry failed conversions that haven't exceeded max retries."""
|
||||
async with get_async_session() as session:
|
||||
try:
|
||||
# Get failed uploads that can be retried
|
||||
result = await session.execute(
|
||||
select(FileUpload)
|
||||
.where(
|
||||
FileUpload.status == "failed",
|
||||
FileUpload.processed == False,
|
||||
(FileUpload.retry_count < self.max_retries) | (FileUpload.retry_count.is_(None))
|
||||
)
|
||||
.limit(5) # Smaller batch for retries
|
||||
)
|
||||
uploads = result.scalars().all()
|
||||
|
||||
for upload in uploads:
|
||||
logger.info(f"Retrying failed conversion for: {upload.filename}")
|
||||
|
||||
# Reset status
|
||||
upload.status = "uploaded"
|
||||
upload.error_message = None
|
||||
|
||||
# Process the file
|
||||
await self._process_single_file(session, upload)
|
||||
|
||||
await session.commit()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error retrying failed conversions: {e}")
|
||||
await session.rollback()
|
||||
|
||||
async def queue_file_for_processing(self, upload_id: str) -> bool:
|
||||
"""Queue a file for processing."""
|
||||
try:
|
||||
# Add to processing queue
|
||||
queue_key = "conversion_queue"
|
||||
await self.redis_client.lpush(queue_key, upload_id)
|
||||
|
||||
logger.info(f"Queued file {upload_id} for processing")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error queuing file for processing: {e}")
|
||||
return False
|
||||
|
||||
async def get_processing_stats(self) -> Dict[str, Any]:
|
||||
"""Get processing statistics."""
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
# Get upload stats by status
|
||||
status_result = await session.execute(
|
||||
select(FileUpload.status, asyncio.func.count())
|
||||
.group_by(FileUpload.status)
|
||||
)
|
||||
status_stats = dict(status_result.fetchall())
|
||||
|
||||
# Get processing stats
|
||||
processed_result = await session.execute(
|
||||
select(asyncio.func.count())
|
||||
.select_from(FileUpload)
|
||||
.where(FileUpload.processed == True)
|
||||
)
|
||||
processed_count = processed_result.scalar()
|
||||
|
||||
# Get failed stats
|
||||
failed_result = await session.execute(
|
||||
select(asyncio.func.count())
|
||||
.select_from(FileUpload)
|
||||
.where(FileUpload.status == "failed")
|
||||
)
|
||||
failed_count = failed_result.scalar()
|
||||
|
||||
return {
|
||||
"status_stats": status_stats,
|
||||
"processed_count": processed_count,
|
||||
"failed_count": failed_count,
|
||||
"is_running": self.is_running,
|
||||
"active_tasks": len([t for t in self.tasks if not t.done()]),
|
||||
"temp_files": len(list(self.temp_dir.glob("*"))),
|
||||
"last_update": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting processing stats: {e}")
|
||||
return {"error": str(e)}
|
||||
|
||||
|
||||
# Global converter instance
|
||||
convert_service = ConvertService()
|
||||
@@ -1,313 +1,500 @@
|
||||
"""Blockchain indexer service for monitoring transactions and events."""
|
||||
|
||||
import asyncio
|
||||
from base64 import b64decode
|
||||
from datetime import datetime
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Dict, List, Optional, Set, Any
|
||||
|
||||
from base58 import b58encode
|
||||
from sqlalchemy import String, and_, desc, cast
|
||||
from tonsdk.boc import Cell
|
||||
from tonsdk.utils import Address
|
||||
from app.core._config import CLIENT_TELEGRAM_BOT_USERNAME
|
||||
from app.core._blockchain.ton.platform import platform
|
||||
from app.core._blockchain.ton.toncenter import toncenter
|
||||
from app.core._utils.send_status import send_status
|
||||
from app.core.logger import make_log
|
||||
from app.core.models import UserContent, KnownTelegramMessage, ServiceConfig
|
||||
from app.core.models.node_storage import StoredContent
|
||||
from app.core._utils.resolve_content import resolve_content
|
||||
from app.core.models.wallet_connection import WalletConnection
|
||||
from app.core._keyboards import get_inline_keyboard
|
||||
from app.core.models._telegram import Wrapped_CBotChat
|
||||
from app.core.storage import db_session
|
||||
import os
|
||||
import traceback
|
||||
import redis.asyncio as redis
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_async_session
|
||||
from app.core.models.blockchain import Transaction, Wallet, BlockchainNFT, BlockchainTokenBalance
|
||||
from app.core.background.ton_service import TONService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def indexer_loop(memory, platform_found: bool, seqno: int) -> [bool, int]:
|
||||
if not platform_found:
|
||||
platform_state = await toncenter.get_account(platform.address.to_string(1, 1, 1))
|
||||
if not platform_state.get('code'):
|
||||
make_log("TON", "Platform contract is not deployed, skipping loop", level="info")
|
||||
await send_status("indexer", "not working: platform is not deployed")
|
||||
return False, seqno
|
||||
else:
|
||||
platform_found = True
|
||||
class IndexerService:
|
||||
"""Service for indexing blockchain transactions and events."""
|
||||
|
||||
make_log("Indexer", "Service running", level="debug")
|
||||
with db_session() as session:
|
||||
def __init__(self):
|
||||
self.settings = get_settings()
|
||||
self.ton_service = TONService()
|
||||
self.redis_client: Optional[redis.Redis] = None
|
||||
self.is_running = False
|
||||
self.tasks: Set[asyncio.Task] = set()
|
||||
|
||||
# Indexing configuration
|
||||
self.batch_size = 100
|
||||
self.index_interval = 30 # seconds
|
||||
self.confirmation_blocks = 12
|
||||
self.max_retries = 3
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Start the indexer service."""
|
||||
try:
|
||||
result = await toncenter.run_get_method('EQD8TJ8xEWB1SpnRE4d89YO3jl0W0EiBnNS4IBaHaUmdfizE', 'get_pool_data')
|
||||
assert result['exit_code'] == 0, f"Error in get-method: {result}"
|
||||
assert result['stack'][0][0] == 'num', f"get first element is not num"
|
||||
assert result['stack'][1][0] == 'num', f"get second element is not num"
|
||||
usdt_per_ton = (int(result['stack'][0][1], 16) * 1e3) / int(result['stack'][1][1], 16)
|
||||
ton_per_star = 0.014 / usdt_per_ton
|
||||
ServiceConfig(session).set('live_tonPerStar', [ton_per_star, datetime.utcnow().timestamp()])
|
||||
make_log("TON_Daemon", f"TON per STAR price: {ton_per_star}", level="DEBUG")
|
||||
except BaseException as e:
|
||||
make_log("TON_Daemon", f"Error while saving TON per STAR price: {e}" + '\n' + traceback.format_exc(), level="ERROR")
|
||||
|
||||
new_licenses = session.query(UserContent).filter(
|
||||
and_(
|
||||
~UserContent.meta.contains({'notification_sent': True}),
|
||||
UserContent.type == 'nft/listen'
|
||||
logger.info("Starting blockchain indexer service")
|
||||
|
||||
# Initialize Redis connection
|
||||
self.redis_client = redis.from_url(
|
||||
self.settings.redis_url,
|
||||
encoding="utf-8",
|
||||
decode_responses=True,
|
||||
socket_keepalive=True,
|
||||
socket_keepalive_options={},
|
||||
health_check_interval=30,
|
||||
)
|
||||
).all()
|
||||
for new_license in new_licenses:
|
||||
licensed_content = session.query(StoredContent).filter(
|
||||
StoredContent.id == new_license.content_id
|
||||
).first()
|
||||
if not licensed_content:
|
||||
make_log("Indexer", f"Licensed content not found: {new_license.content_id}", level="error")
|
||||
|
||||
content_metadata = licensed_content.metadata_json(session)
|
||||
assert content_metadata, "No content metadata found"
|
||||
|
||||
if not (licensed_content.owner_address == new_license.owner_address):
|
||||
try:
|
||||
user = new_license.user
|
||||
if user.telegram_id and licensed_content:
|
||||
await (Wrapped_CBotChat(memory._client_telegram_bot, chat_id=user.telegram_id, user=user, db_session=session)).send_content(
|
||||
session, licensed_content
|
||||
)
|
||||
|
||||
wallet_owner_connection = session.query(WalletConnection).filter_by(
|
||||
wallet_address=licensed_content.owner_address,
|
||||
invalidated=False
|
||||
).order_by(desc(WalletConnection.id)).first()
|
||||
wallet_owner_user = wallet_owner_connection.user
|
||||
if wallet_owner_user.telegram_id:
|
||||
wallet_owner_bot = Wrapped_CBotChat(memory._telegram_bot, chat_id=wallet_owner_user.telegram_id, user=wallet_owner_user, db_session=session)
|
||||
await wallet_owner_bot.send_message(
|
||||
user.translated('p_licenseWasBought').format(
|
||||
username=user.front_format(),
|
||||
nft_address=f'"https://tonviewer.com/{new_license.onchain_address}"',
|
||||
content_title=content_metadata.get('name', 'Unknown'),
|
||||
),
|
||||
message_type='notification',
|
||||
)
|
||||
except BaseException as e:
|
||||
make_log("IndexerSendNewLicense", f"Error: {e}" + '\n' + traceback.format_exc(), level="error")
|
||||
|
||||
new_license.meta = {**new_license.meta, 'notification_sent': True}
|
||||
session.commit()
|
||||
|
||||
content_without_cid = session.query(StoredContent).filter(
|
||||
StoredContent.content_id == None
|
||||
)
|
||||
for target_content in content_without_cid:
|
||||
target_cid = target_content.cid.serialize_v2()
|
||||
make_log("Indexer", f"Content without CID: {target_content.hash}, setting CID: {target_cid}", level="debug")
|
||||
target_content.content_id = target_cid
|
||||
|
||||
session.commit()
|
||||
|
||||
last_known_index_ = session.query(StoredContent).filter(
|
||||
StoredContent.onchain_index != None
|
||||
).order_by(StoredContent.onchain_index.desc()).first()
|
||||
last_known_index = last_known_index_.onchain_index if last_known_index_ else 0
|
||||
last_known_index = max(last_known_index, 0)
|
||||
make_log("Indexer", f"Last known index: {last_known_index}", level="debug")
|
||||
if last_known_index_:
|
||||
next_item_index = last_known_index + 1
|
||||
else:
|
||||
next_item_index = 0
|
||||
|
||||
resolve_item_result = await toncenter.run_get_method(platform.address.to_string(1, 1, 1), 'get_nft_address_by_index', [['num', next_item_index]])
|
||||
make_log("Indexer", f"Resolve item result: {resolve_item_result}", level="debug")
|
||||
if resolve_item_result.get('exit_code', -1) != 0:
|
||||
make_log("Indexer", f"Resolve item error: {resolve_item_result}", level="error")
|
||||
return platform_found, seqno
|
||||
|
||||
item_address_cell_b64 = resolve_item_result['stack'][0][1]["bytes"]
|
||||
item_address_slice = Cell.one_from_boc(b64decode(item_address_cell_b64)).begin_parse()
|
||||
item_address = item_address_slice.read_msg_addr()
|
||||
make_log("Indexer", f"Item address: {item_address.to_string(1, 1, 1)}", level="debug")
|
||||
|
||||
item_get_data_result = await toncenter.run_get_method(item_address.to_string(1, 1, 1), 'indexator_data')
|
||||
if item_get_data_result.get('exit_code', -1) != 0:
|
||||
make_log("Indexer", f"Get item data error (maybe not deployed): {item_get_data_result}", level="debug")
|
||||
return platform_found, seqno
|
||||
|
||||
assert item_get_data_result['stack'][0][0] == 'num', "Item type is not a number"
|
||||
assert int(item_get_data_result['stack'][0][1], 16) == 1, "Item is not COP NFT"
|
||||
item_returned_address = Cell.one_from_boc(b64decode(item_get_data_result['stack'][1][1]['bytes'])).begin_parse().read_msg_addr()
|
||||
assert (
|
||||
item_returned_address.to_string(1, 1, 1) == item_address.to_string(1, 1, 1)
|
||||
), "Item address mismatch"
|
||||
|
||||
assert item_get_data_result['stack'][2][0] == 'num', "Item index is not a number"
|
||||
item_index = int(item_get_data_result['stack'][2][1], 16)
|
||||
assert item_index == next_item_index, "Item index mismatch"
|
||||
|
||||
item_platform_address = Cell.one_from_boc(b64decode(item_get_data_result['stack'][3][1]['bytes'])).begin_parse().read_msg_addr()
|
||||
assert item_platform_address.to_string(1, 1, 1) == Address(platform.address.to_string(1, 1, 1)).to_string(1, 1, 1), "Item platform address mismatch"
|
||||
|
||||
assert item_get_data_result['stack'][4][0] == 'num', "Item license type is not a number"
|
||||
item_license_type = int(item_get_data_result['stack'][4][1], 16)
|
||||
assert item_license_type == 0, "Item license type is not 0"
|
||||
|
||||
item_owner_address = Cell.one_from_boc(b64decode(item_get_data_result['stack'][5][1]["bytes"])).begin_parse().read_msg_addr()
|
||||
item_values = Cell.one_from_boc(b64decode(item_get_data_result['stack'][6][1]['bytes']))
|
||||
item_derivates = Cell.one_from_boc(b64decode(item_get_data_result['stack'][7][1]['bytes']))
|
||||
item_platform_variables = Cell.one_from_boc(b64decode(item_get_data_result['stack'][8][1]['bytes']))
|
||||
item_distribution = Cell.one_from_boc(b64decode(item_get_data_result['stack'][9][1]['bytes']))
|
||||
|
||||
item_distribution_slice = item_distribution.begin_parse()
|
||||
item_prices_slice = item_distribution_slice.refs[0].begin_parse()
|
||||
item_listen_license_price = item_prices_slice.read_coins()
|
||||
item_use_license_price = item_prices_slice.read_coins()
|
||||
item_resale_license_price = item_prices_slice.read_coins()
|
||||
|
||||
item_values_slice = item_values.begin_parse()
|
||||
item_content_hash_int = item_values_slice.read_uint(256)
|
||||
item_content_hash = item_content_hash_int.to_bytes(32, 'big')
|
||||
# item_content_hash_str = b58encode(item_content_hash).decode()
|
||||
item_metadata = item_values_slice.refs[0]
|
||||
item_content = item_values_slice.refs[1]
|
||||
item_metadata_str = item_metadata.bits.array.decode()
|
||||
item_content_cid_str = item_content.refs[0].bits.array.decode()
|
||||
item_content_cover_cid_str = item_content.refs[1].bits.array.decode()
|
||||
item_content_metadata_cid_str = item_content.refs[2].bits.array.decode()
|
||||
|
||||
item_content_cid, err = resolve_content(item_content_cid_str)
|
||||
item_content_hash = item_content_cid.content_hash
|
||||
item_content_hash_str = item_content_cid.content_hash_b58
|
||||
|
||||
item_metadata_packed = {
|
||||
'license_type': item_license_type,
|
||||
'item_address': item_address.to_string(1, 1, 1),
|
||||
'content_cid': item_content_cid_str,
|
||||
'cover_cid': item_content_cover_cid_str,
|
||||
'metadata_cid': item_content_metadata_cid_str,
|
||||
'derivates': b58encode(item_derivates.to_boc(False)).decode(),
|
||||
'platform_variables': b58encode(item_platform_variables.to_boc(False)).decode(),
|
||||
'license': {
|
||||
'listen': {
|
||||
'price': str(item_listen_license_price)
|
||||
},
|
||||
'use': {
|
||||
'price': str(item_use_license_price)
|
||||
},
|
||||
'resale': {
|
||||
'price': str(item_resale_license_price)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
user_wallet_connection = None
|
||||
if item_owner_address:
|
||||
user_wallet_connection = session.query(WalletConnection).filter(
|
||||
WalletConnection.wallet_address == item_owner_address.to_string(1, 1, 1)
|
||||
).first()
|
||||
|
||||
encrypted_stored_content = session.query(StoredContent).filter(
|
||||
StoredContent.hash == item_content_hash_str,
|
||||
# StoredContent.type.like("local%")
|
||||
).first()
|
||||
if encrypted_stored_content:
|
||||
is_duplicate = encrypted_stored_content.type.startswith("onchain") \
|
||||
and encrypted_stored_content.onchain_index != item_index
|
||||
if not is_duplicate:
|
||||
if encrypted_stored_content.type.startswith('local'):
|
||||
encrypted_stored_content.type = "onchain/content" + ("_unknown" if (encrypted_stored_content.key_id is None) else "")
|
||||
encrypted_stored_content.onchain_index = item_index
|
||||
encrypted_stored_content.owner_address = item_owner_address.to_string(1, 1, 1)
|
||||
user = None
|
||||
if user_wallet_connection:
|
||||
encrypted_stored_content.user_id = user_wallet_connection.user_id
|
||||
user = user_wallet_connection.user
|
||||
|
||||
if user:
|
||||
user_uploader_wrapper = Wrapped_CBotChat(memory._telegram_bot, chat_id=user.telegram_id, user=user, db_session=session)
|
||||
await user_uploader_wrapper.send_message(
|
||||
user.translated('p_contentWasIndexed').format(
|
||||
item_address=item_address.to_string(1, 1, 1),
|
||||
item_index=item_index,
|
||||
),
|
||||
message_type='notification',
|
||||
reply_markup=get_inline_keyboard([
|
||||
[{
|
||||
'text': user.translated('viewTrackAsClient_button'),
|
||||
'url': f"https://t.me/{CLIENT_TELEGRAM_BOT_USERNAME}?start=C{encrypted_stored_content.cid.serialize_v2()}"
|
||||
}],
|
||||
])
|
||||
)
|
||||
|
||||
try:
|
||||
for hint_message in session.query(KnownTelegramMessage).filter(
|
||||
and_(
|
||||
KnownTelegramMessage.chat_id == user.telegram_id,
|
||||
KnownTelegramMessage.type == 'hint',
|
||||
cast(KnownTelegramMessage.meta['encrypted_content_hash'], String) == encrypted_stored_content.hash,
|
||||
KnownTelegramMessage.deleted == False
|
||||
)
|
||||
).all():
|
||||
await user_uploader_wrapper.delete_message(hint_message.message_id)
|
||||
except BaseException as e:
|
||||
make_log("Indexer", f"Error while deleting hint messages: {e}" + '\n' + traceback.format_exc(), level="error")
|
||||
elif encrypted_stored_content.type.startswith('onchain') and encrypted_stored_content.onchain_index == item_index:
|
||||
encrypted_stored_content.type = "onchain/content" + ("_unknown" if (encrypted_stored_content.key_id is None) else "")
|
||||
encrypted_stored_content.owner_address = item_owner_address.to_string(1, 1, 1)
|
||||
if user_wallet_connection:
|
||||
encrypted_stored_content.user_id = user_wallet_connection.user_id
|
||||
else:
|
||||
make_log("Indexer", f"[CRITICAL] Item already indexed and ERRORED!: {item_content_hash_str}", level="error")
|
||||
return platform_found, seqno
|
||||
|
||||
encrypted_stored_content.updated = datetime.now()
|
||||
encrypted_stored_content.meta = {
|
||||
**encrypted_stored_content.meta,
|
||||
**item_metadata_packed
|
||||
}
|
||||
|
||||
session.commit()
|
||||
return platform_found, seqno
|
||||
else:
|
||||
item_metadata_packed['copied_from'] = encrypted_stored_content.id
|
||||
item_metadata_packed['copied_from_cid'] = encrypted_stored_content.cid.serialize_v2()
|
||||
item_content_hash_str = f"{b58encode(bytes(16) + os.urandom(30)).decode()}" # check this for vulnerability
|
||||
|
||||
onchain_stored_content = StoredContent(
|
||||
type="onchain/content_unknown",
|
||||
hash=item_content_hash_str,
|
||||
onchain_index=item_index,
|
||||
owner_address=item_owner_address.to_string(1, 1, 1) if item_owner_address else None,
|
||||
meta=item_metadata_packed,
|
||||
filename="UNKNOWN_ENCRYPTED_CONTENT",
|
||||
user_id=user_wallet_connection.user_id if user_wallet_connection else None,
|
||||
created=datetime.now(),
|
||||
encrypted=True,
|
||||
decrypted_content_id=None,
|
||||
key_id=None,
|
||||
updated=datetime.now()
|
||||
)
|
||||
session.add(onchain_stored_content)
|
||||
session.commit()
|
||||
make_log("Indexer", f"Item indexed: {item_content_hash_str}", level="info")
|
||||
last_known_index += 1
|
||||
|
||||
return platform_found, seqno
|
||||
|
||||
|
||||
async def main_fn(memory, ):
|
||||
make_log("Indexer", "Service started", level="info")
|
||||
platform_found = False
|
||||
seqno = 0
|
||||
while True:
|
||||
|
||||
# Test Redis connection
|
||||
await self.redis_client.ping()
|
||||
logger.info("Redis connection established for indexer")
|
||||
|
||||
# Start indexing tasks
|
||||
self.is_running = True
|
||||
|
||||
# Create indexing tasks
|
||||
tasks = [
|
||||
asyncio.create_task(self._index_transactions_loop()),
|
||||
asyncio.create_task(self._index_wallets_loop()),
|
||||
asyncio.create_task(self._index_nfts_loop()),
|
||||
asyncio.create_task(self._index_token_balances_loop()),
|
||||
asyncio.create_task(self._cleanup_cache_loop()),
|
||||
]
|
||||
|
||||
self.tasks.update(tasks)
|
||||
|
||||
# Wait for all tasks
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error starting indexer service: {e}")
|
||||
await self.stop()
|
||||
raise
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the indexer service."""
|
||||
logger.info("Stopping blockchain indexer service")
|
||||
self.is_running = False
|
||||
|
||||
# Cancel all tasks
|
||||
for task in self.tasks:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
|
||||
# Wait for tasks to complete
|
||||
if self.tasks:
|
||||
await asyncio.gather(*self.tasks, return_exceptions=True)
|
||||
|
||||
# Close Redis connection
|
||||
if self.redis_client:
|
||||
await self.redis_client.close()
|
||||
|
||||
logger.info("Indexer service stopped")
|
||||
|
||||
async def _index_transactions_loop(self) -> None:
|
||||
"""Main loop for indexing transactions."""
|
||||
logger.info("Starting transaction indexing loop")
|
||||
|
||||
while self.is_running:
|
||||
try:
|
||||
await self._index_pending_transactions()
|
||||
await self._update_transaction_confirmations()
|
||||
await asyncio.sleep(self.index_interval)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error in transaction indexing loop: {e}")
|
||||
await asyncio.sleep(self.index_interval)
|
||||
|
||||
async def _index_wallets_loop(self) -> None:
|
||||
"""Loop for updating wallet information."""
|
||||
logger.info("Starting wallet indexing loop")
|
||||
|
||||
while self.is_running:
|
||||
try:
|
||||
await self._update_wallet_balances()
|
||||
await asyncio.sleep(self.index_interval * 2) # Less frequent
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error in wallet indexing loop: {e}")
|
||||
await asyncio.sleep(self.index_interval * 2)
|
||||
|
||||
async def _index_nfts_loop(self) -> None:
|
||||
"""Loop for indexing NFT collections and transfers."""
|
||||
logger.info("Starting NFT indexing loop")
|
||||
|
||||
while self.is_running:
|
||||
try:
|
||||
await self._index_nft_collections()
|
||||
await self._index_nft_transfers()
|
||||
await asyncio.sleep(self.index_interval * 4) # Even less frequent
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error in NFT indexing loop: {e}")
|
||||
await asyncio.sleep(self.index_interval * 4)
|
||||
|
||||
async def _index_token_balances_loop(self) -> None:
|
||||
"""Loop for updating token balances."""
|
||||
logger.info("Starting token balance indexing loop")
|
||||
|
||||
while self.is_running:
|
||||
try:
|
||||
await self._update_token_balances()
|
||||
await asyncio.sleep(self.index_interval * 3)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error in token balance indexing loop: {e}")
|
||||
await asyncio.sleep(self.index_interval * 3)
|
||||
|
||||
async def _cleanup_cache_loop(self) -> None:
|
||||
"""Loop for cleaning up old cache entries."""
|
||||
logger.info("Starting cache cleanup loop")
|
||||
|
||||
while self.is_running:
|
||||
try:
|
||||
await self._cleanup_old_cache_entries()
|
||||
await asyncio.sleep(3600) # Run every hour
|
||||
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error in cache cleanup loop: {e}")
|
||||
await asyncio.sleep(3600)
|
||||
|
||||
async def _index_pending_transactions(self) -> None:
|
||||
"""Index pending transactions from the database."""
|
||||
async with get_async_session() as session:
|
||||
try:
|
||||
# Get pending transactions
|
||||
result = await session.execute(
|
||||
select(Transaction)
|
||||
.where(Transaction.status == "pending")
|
||||
.limit(self.batch_size)
|
||||
)
|
||||
transactions = result.scalars().all()
|
||||
|
||||
if not transactions:
|
||||
return
|
||||
|
||||
logger.info(f"Indexing {len(transactions)} pending transactions")
|
||||
|
||||
# Process each transaction
|
||||
for transaction in transactions:
|
||||
await self._process_transaction(session, transaction)
|
||||
|
||||
await session.commit()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error indexing pending transactions: {e}")
|
||||
await session.rollback()
|
||||
|
||||
async def _process_transaction(self, session: AsyncSession, transaction: Transaction) -> None:
|
||||
"""Process a single transaction."""
|
||||
try:
|
||||
platform_found, seqno = await indexer_loop(memory, platform_found, seqno)
|
||||
except BaseException as e:
|
||||
make_log("Indexer", f"Error: {e}" + '\n' + traceback.format_exc(), level="error")
|
||||
|
||||
if platform_found:
|
||||
await send_status("indexer", f"working (seqno={seqno})")
|
||||
|
||||
await asyncio.sleep(5)
|
||||
seqno += 1
|
||||
# Check transaction status on blockchain
|
||||
if transaction.tx_hash:
|
||||
tx_info = await self.ton_service.get_transaction_info(transaction.tx_hash)
|
||||
|
||||
if tx_info:
|
||||
# Update transaction with blockchain data
|
||||
transaction.status = tx_info.get("status", "pending")
|
||||
transaction.block_number = tx_info.get("block_number")
|
||||
transaction.gas_used = tx_info.get("gas_used")
|
||||
transaction.gas_price = tx_info.get("gas_price")
|
||||
transaction.confirmations = tx_info.get("confirmations", 0)
|
||||
transaction.updated_at = datetime.utcnow()
|
||||
|
||||
# Cache transaction info
|
||||
cache_key = f"tx:{transaction.tx_hash}"
|
||||
await self.redis_client.setex(
|
||||
cache_key,
|
||||
3600, # 1 hour
|
||||
json.dumps(tx_info)
|
||||
)
|
||||
|
||||
logger.debug(f"Updated transaction {transaction.tx_hash}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing transaction {transaction.id}: {e}")
|
||||
|
||||
async def _update_transaction_confirmations(self) -> None:
|
||||
"""Update confirmation counts for recent transactions."""
|
||||
async with get_async_session() as session:
|
||||
try:
|
||||
# Get recent confirmed transactions
|
||||
cutoff_time = datetime.utcnow() - timedelta(hours=24)
|
||||
result = await session.execute(
|
||||
select(Transaction)
|
||||
.where(
|
||||
Transaction.status == "confirmed",
|
||||
Transaction.confirmations < self.confirmation_blocks,
|
||||
Transaction.updated_at > cutoff_time
|
||||
)
|
||||
.limit(self.batch_size)
|
||||
)
|
||||
transactions = result.scalars().all()
|
||||
|
||||
for transaction in transactions:
|
||||
if transaction.tx_hash:
|
||||
try:
|
||||
confirmations = await self.ton_service.get_transaction_confirmations(
|
||||
transaction.tx_hash
|
||||
)
|
||||
|
||||
if confirmations != transaction.confirmations:
|
||||
transaction.confirmations = confirmations
|
||||
transaction.updated_at = datetime.utcnow()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating confirmations for {transaction.tx_hash}: {e}")
|
||||
|
||||
await session.commit()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating transaction confirmations: {e}")
|
||||
await session.rollback()
|
||||
|
||||
async def _update_wallet_balances(self) -> None:
|
||||
"""Update wallet balances from the blockchain."""
|
||||
async with get_async_session() as session:
|
||||
try:
|
||||
# Get active wallets
|
||||
result = await session.execute(
|
||||
select(Wallet)
|
||||
.where(Wallet.is_active == True)
|
||||
.limit(self.batch_size)
|
||||
)
|
||||
wallets = result.scalars().all()
|
||||
|
||||
for wallet in wallets:
|
||||
try:
|
||||
# Get current balance
|
||||
balance = await self.ton_service.get_wallet_balance(wallet.address)
|
||||
|
||||
if balance != wallet.balance:
|
||||
wallet.balance = balance
|
||||
wallet.updated_at = datetime.utcnow()
|
||||
|
||||
# Cache balance
|
||||
cache_key = f"balance:{wallet.address}"
|
||||
await self.redis_client.setex(cache_key, 300, str(balance)) # 5 minutes
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating balance for wallet {wallet.address}: {e}")
|
||||
|
||||
await session.commit()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating wallet balances: {e}")
|
||||
await session.rollback()
|
||||
|
||||
async def _index_nft_collections(self) -> None:
|
||||
"""Index NFT collections and metadata."""
|
||||
async with get_async_session() as session:
|
||||
try:
|
||||
# Get wallets to check for NFTs
|
||||
result = await session.execute(
|
||||
select(Wallet)
|
||||
.where(Wallet.is_active == True)
|
||||
.limit(self.batch_size // 4) # Smaller batch for NFTs
|
||||
)
|
||||
wallets = result.scalars().all()
|
||||
|
||||
for wallet in wallets:
|
||||
try:
|
||||
# Get NFTs for this wallet
|
||||
nfts = await self.ton_service.get_wallet_nfts(wallet.address)
|
||||
|
||||
for nft_data in nfts:
|
||||
await self._process_nft(session, wallet, nft_data)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error indexing NFTs for wallet {wallet.address}: {e}")
|
||||
|
||||
await session.commit()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error indexing NFT collections: {e}")
|
||||
await session.rollback()
|
||||
|
||||
async def _process_nft(self, session: AsyncSession, wallet: Wallet, nft_data: Dict[str, Any]) -> None:
|
||||
"""Process a single NFT."""
|
||||
try:
|
||||
# Check if NFT exists
|
||||
result = await session.execute(
|
||||
select(BlockchainNFT)
|
||||
.where(
|
||||
BlockchainNFT.token_id == nft_data["token_id"],
|
||||
BlockchainNFT.collection_address == nft_data["collection_address"]
|
||||
)
|
||||
)
|
||||
existing_nft = result.scalar_one_or_none()
|
||||
|
||||
if existing_nft:
|
||||
# Update existing NFT
|
||||
existing_nft.owner_address = wallet.address
|
||||
existing_nft.metadata = nft_data.get("metadata", {})
|
||||
existing_nft.updated_at = datetime.utcnow()
|
||||
else:
|
||||
# Create new NFT
|
||||
new_nft = BlockchainNFT(
|
||||
wallet_id=wallet.id,
|
||||
token_id=nft_data["token_id"],
|
||||
collection_address=nft_data["collection_address"],
|
||||
owner_address=wallet.address,
|
||||
token_uri=nft_data.get("token_uri"),
|
||||
metadata=nft_data.get("metadata", {}),
|
||||
created_at=datetime.utcnow()
|
||||
)
|
||||
session.add(new_nft)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing NFT {nft_data.get('token_id')}: {e}")
|
||||
|
||||
async def _index_nft_transfers(self) -> None:
|
||||
"""Index NFT transfers."""
|
||||
# This would involve checking recent blocks for NFT transfer events
|
||||
# Implementation depends on the specific blockchain's event system
|
||||
pass
|
||||
|
||||
async def _update_token_balances(self) -> None:
|
||||
"""Update token balances for wallets."""
|
||||
async with get_async_session() as session:
|
||||
try:
|
||||
# Get wallets with token balances to update
|
||||
result = await session.execute(
|
||||
select(Wallet)
|
||||
.where(Wallet.is_active == True)
|
||||
.limit(self.batch_size // 2)
|
||||
)
|
||||
wallets = result.scalars().all()
|
||||
|
||||
for wallet in wallets:
|
||||
try:
|
||||
# Get token balances
|
||||
token_balances = await self.ton_service.get_wallet_token_balances(wallet.address)
|
||||
|
||||
for token_data in token_balances:
|
||||
await self._update_token_balance(session, wallet, token_data)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating token balances for {wallet.address}: {e}")
|
||||
|
||||
await session.commit()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating token balances: {e}")
|
||||
await session.rollback()
|
||||
|
||||
async def _update_token_balance(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
wallet: Wallet,
|
||||
token_data: Dict[str, Any]
|
||||
) -> None:
|
||||
"""Update a single token balance."""
|
||||
try:
|
||||
# Check if balance record exists
|
||||
result = await session.execute(
|
||||
select(BlockchainTokenBalance)
|
||||
.where(
|
||||
BlockchainTokenBalance.wallet_id == wallet.id,
|
||||
BlockchainTokenBalance.token_address == token_data["token_address"]
|
||||
)
|
||||
)
|
||||
existing_balance = result.scalar_one_or_none()
|
||||
|
||||
if existing_balance:
|
||||
# Update existing balance
|
||||
existing_balance.balance = token_data["balance"]
|
||||
existing_balance.decimals = token_data.get("decimals", 18)
|
||||
existing_balance.updated_at = datetime.utcnow()
|
||||
else:
|
||||
# Create new balance record
|
||||
new_balance = BlockchainTokenBalance(
|
||||
wallet_id=wallet.id,
|
||||
token_address=token_data["token_address"],
|
||||
token_name=token_data.get("name"),
|
||||
token_symbol=token_data.get("symbol"),
|
||||
balance=token_data["balance"],
|
||||
decimals=token_data.get("decimals", 18),
|
||||
created_at=datetime.utcnow()
|
||||
)
|
||||
session.add(new_balance)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating token balance: {e}")
|
||||
|
||||
async def _cleanup_old_cache_entries(self) -> None:
|
||||
"""Clean up old cache entries."""
|
||||
try:
|
||||
# Get all keys with our prefixes
|
||||
patterns = ["tx:*", "balance:*", "nft:*", "token:*"]
|
||||
|
||||
for pattern in patterns:
|
||||
keys = await self.redis_client.keys(pattern)
|
||||
|
||||
# Check TTL and remove expired keys
|
||||
for key in keys:
|
||||
ttl = await self.redis_client.ttl(key)
|
||||
if ttl == -1: # No expiration set
|
||||
await self.redis_client.expire(key, 3600) # Set 1 hour expiration
|
||||
|
||||
logger.debug("Cache cleanup completed")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error during cache cleanup: {e}")
|
||||
|
||||
async def get_indexing_stats(self) -> Dict[str, Any]:
|
||||
"""Get indexing statistics."""
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
# Get transaction stats
|
||||
tx_result = await session.execute(
|
||||
select(Transaction.status, asyncio.func.count())
|
||||
.group_by(Transaction.status)
|
||||
)
|
||||
tx_stats = dict(tx_result.fetchall())
|
||||
|
||||
# Get wallet stats
|
||||
wallet_result = await session.execute(
|
||||
select(asyncio.func.count())
|
||||
.select_from(Wallet)
|
||||
.where(Wallet.is_active == True)
|
||||
)
|
||||
active_wallets = wallet_result.scalar()
|
||||
|
||||
# Get NFT stats
|
||||
nft_result = await session.execute(
|
||||
select(asyncio.func.count())
|
||||
.select_from(BlockchainNFT)
|
||||
)
|
||||
total_nfts = nft_result.scalar()
|
||||
|
||||
return {
|
||||
"transaction_stats": tx_stats,
|
||||
"active_wallets": active_wallets,
|
||||
"total_nfts": total_nfts,
|
||||
"is_running": self.is_running,
|
||||
"active_tasks": len([t for t in self.tasks if not t.done()]),
|
||||
"last_update": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting indexing stats: {e}")
|
||||
return {"error": str(e)}
|
||||
|
||||
|
||||
|
||||
# if __name__ == '__main__':
|
||||
# loop = asyncio.get_event_loop()
|
||||
# loop.run_until_complete(main())
|
||||
# loop.close()
|
||||
# Global indexer instance
|
||||
indexer_service = IndexerService()
|
||||
+644
-276
@@ -1,290 +1,658 @@
|
||||
"""
|
||||
TON Blockchain service for wallet operations, transaction management, and smart contract interactions.
|
||||
Provides async operations with connection pooling, caching, and comprehensive error handling.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from base64 import b64decode
|
||||
import os
|
||||
import traceback
|
||||
import httpx
|
||||
|
||||
from sqlalchemy import and_, func
|
||||
from tonsdk.boc import begin_cell, Cell
|
||||
from tonsdk.contract.wallet import Wallets
|
||||
from tonsdk.utils import HighloadQueryId
|
||||
import json
|
||||
from datetime import datetime, timedelta
|
||||
from decimal import Decimal
|
||||
from typing import Dict, List, Optional, Any, Tuple
|
||||
from uuid import UUID
|
||||
|
||||
from app.core._blockchain.ton.platform import platform
|
||||
from app.core._blockchain.ton.toncenter import toncenter
|
||||
from app.core.models.tasks import BlockchainTask
|
||||
from app.core._config import MY_FUND_ADDRESS
|
||||
from app.core._secrets import service_wallet
|
||||
from app.core._utils.send_status import send_status
|
||||
from app.core.storage import db_session
|
||||
from app.core.logger import make_log
|
||||
import httpx
|
||||
from sqlalchemy import select, update, and_
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_async_session, get_cache_manager
|
||||
from app.core.logging import get_logger
|
||||
from app.core.security import decrypt_data, encrypt_data
|
||||
|
||||
async def get_sw_seqno():
|
||||
sw_seqno_result = await toncenter.run_get_method(service_wallet.address.to_string(1, 1, 1), 'seqno')
|
||||
if sw_seqno_result.get('exit_code', -1) != 0:
|
||||
sw_seqno_value = 0
|
||||
else:
|
||||
sw_seqno_value = int(sw_seqno_result.get('stack', [['num', '0x0']])[0][1], 16)
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
return sw_seqno_value
|
||||
|
||||
|
||||
async def main_fn(memory):
|
||||
make_log("TON", f"Service started, SW = {service_wallet.address.to_string(1, 1, 1)}", level="info")
|
||||
sw_seqno_value = await get_sw_seqno()
|
||||
make_log("TON", f"Service wallet run seqno method: {sw_seqno_value}", level="info")
|
||||
if sw_seqno_value == 0:
|
||||
make_log("TON", "Service wallet is not deployed, deploying...", level="info")
|
||||
await toncenter.send_boc(
|
||||
service_wallet.create_transfer_message(
|
||||
[{
|
||||
'address': service_wallet.address.to_string(1, 1, 1),
|
||||
'amount': 1,
|
||||
'send_mode': 1,
|
||||
'payload': begin_cell().store_uint(0, 32).store_bytes(b"Init MY Node").end_cell()
|
||||
}], 0
|
||||
)['message'].to_boc(False)
|
||||
)
|
||||
await asyncio.sleep(5)
|
||||
return await main_fn(memory)
|
||||
|
||||
if os.getenv("TON_BEGIN_COMMAND_WITHDRAW"):
|
||||
await toncenter.send_boc(
|
||||
service_wallet.create_transfer_message(
|
||||
[{
|
||||
'address': MY_FUND_ADDRESS,
|
||||
'amount': 1,
|
||||
'send_mode': 128,
|
||||
'payload': begin_cell().end_cell()
|
||||
}], sw_seqno_value
|
||||
)['message'].to_boc(False)
|
||||
)
|
||||
make_log("TON", "Withdraw command sent", level="info")
|
||||
await asyncio.sleep(10)
|
||||
return await main_fn(memory)
|
||||
|
||||
# TODO: не деплоить если указан master_address и мы проверили что аккаунт существует. Сейчас platform у каждой ноды будет разным
|
||||
platform_state = await toncenter.get_account(platform.address.to_string(1, 1, 1))
|
||||
if not platform_state.get('code'):
|
||||
make_log("TON", "Platform contract is not deployed, send deploy transaction..", level="info")
|
||||
await toncenter.send_boc(
|
||||
service_wallet.create_transfer_message(
|
||||
[{
|
||||
'address': platform.address.to_string(1, 1, 1),
|
||||
'amount': int(0.08 * 10 ** 9),
|
||||
'send_mode': 1,
|
||||
'payload': begin_cell().store_uint(0, 32).store_uint(0, 64).end_cell(),
|
||||
'state_init': platform.create_state_init()['state_init']
|
||||
}], sw_seqno_value
|
||||
)['message'].to_boc(False)
|
||||
)
|
||||
|
||||
await send_status("ton_daemon", "working: deploying platform")
|
||||
await asyncio.sleep(15)
|
||||
return await main_fn(memory)
|
||||
class TONService:
|
||||
"""
|
||||
Comprehensive TON blockchain service with async operations.
|
||||
Handles wallet management, transactions, and smart contract interactions.
|
||||
"""
|
||||
|
||||
highload_wallet = Wallets.ALL['hv3'](
|
||||
private_key=service_wallet.options['private_key'],
|
||||
public_key=service_wallet.options['public_key'],
|
||||
wc=0
|
||||
)
|
||||
make_log("TON", f"Highload wallet address: {highload_wallet.address.to_string(1, 1, 1)}", level="info")
|
||||
highload_state = await toncenter.get_account(highload_wallet.address.to_string(1, 1, 1))
|
||||
if int(highload_state.get('balance', '0')) / 1e9 < 0.05:
|
||||
make_log("TON", "Highload wallet balance is less than 0.05, send topup transaction..", level="info")
|
||||
await toncenter.send_boc(
|
||||
service_wallet.create_transfer_message(
|
||||
[{
|
||||
'address': highload_wallet.address.to_string(1, 1, 0),
|
||||
'amount': int(0.08 * 10 ** 9),
|
||||
'send_mode': 1,
|
||||
'payload': begin_cell().store_uint(0, 32).end_cell()
|
||||
}], sw_seqno_value
|
||||
)['message'].to_boc(False)
|
||||
def __init__(self):
|
||||
self.api_endpoint = settings.TON_API_ENDPOINT
|
||||
self.testnet = settings.TON_TESTNET
|
||||
self.api_key = settings.TON_API_KEY
|
||||
self.timeout = 30
|
||||
|
||||
# HTTP client for API requests
|
||||
self.client = httpx.AsyncClient(
|
||||
timeout=self.timeout,
|
||||
headers={
|
||||
"Authorization": f"Bearer {self.api_key}" if self.api_key else None,
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
)
|
||||
await send_status("ton_daemon", "working: topup highload wallet")
|
||||
await asyncio.sleep(15)
|
||||
return await main_fn(memory)
|
||||
|
||||
self.cache_manager = get_cache_manager()
|
||||
|
||||
if not highload_state.get('code'):
|
||||
make_log("TON", "Highload wallet contract is not deployed, send deploy transaction..", level="info")
|
||||
created_at_ts = int(datetime.utcnow().timestamp()) - 60
|
||||
await toncenter.send_boc(
|
||||
highload_wallet.create_transfer_message(
|
||||
service_wallet.address.to_string(1, 1, 1),
|
||||
1, HighloadQueryId.from_seqno(0), created_at_ts, send_mode=1, payload="hello world", need_deploy=True
|
||||
)['message'].to_boc(False)
|
||||
)
|
||||
await send_status("ton_daemon", "working: deploying highload wallet")
|
||||
await asyncio.sleep(15)
|
||||
return await main_fn(memory)
|
||||
|
||||
while True:
|
||||
async def close(self):
|
||||
"""Close HTTP client and cleanup resources."""
|
||||
if self.client:
|
||||
await self.client.aclose()
|
||||
|
||||
async def create_wallet(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Create new TON wallet with mnemonic generation.
|
||||
|
||||
Returns:
|
||||
Dict: Wallet creation result with address and private key
|
||||
"""
|
||||
try:
|
||||
sw_seqno_value = await get_sw_seqno()
|
||||
make_log("TON", f"Service running ({sw_seqno_value})", level="debug")
|
||||
|
||||
with db_session() as session:
|
||||
# Проверка отправленных сообщений
|
||||
await send_status("ton_daemon", f"working: processing in-txs (seqno={sw_seqno_value})")
|
||||
async def process_incoming_transaction(transaction: dict):
|
||||
transaction_hash = transaction['transaction_id']['hash']
|
||||
transaction_lt = str(transaction['transaction_id']['lt'])
|
||||
# transaction_success = bool(transaction['success'])
|
||||
|
||||
async def process_incoming_message(blockchain_message: dict):
|
||||
in_msg_cell = Cell.one_from_boc(b64decode(blockchain_message['msg_data']['body']))
|
||||
in_msg_slice = in_msg_cell.refs[0].begin_parse()
|
||||
in_msg_slice.read_uint(32)
|
||||
in_msg_slice.read_uint(8)
|
||||
in_msg_query_id = in_msg_slice.read_uint(23)
|
||||
in_msg_created_at = in_msg_slice.read_uint(64)
|
||||
in_msg_epoch = int(in_msg_created_at // (60 * 60))
|
||||
in_msg_seqno = HighloadQueryId.from_query_id(in_msg_query_id).to_seqno()
|
||||
in_msg_blockchain_task = (
|
||||
session.query(BlockchainTask).filter(
|
||||
and_(
|
||||
BlockchainTask.seqno == in_msg_seqno,
|
||||
BlockchainTask.epoch == in_msg_epoch,
|
||||
)
|
||||
)
|
||||
).first()
|
||||
if not in_msg_blockchain_task:
|
||||
return
|
||||
|
||||
if not (in_msg_blockchain_task.status in ['done']) or in_msg_blockchain_task.transaction_hash != transaction_hash:
|
||||
in_msg_blockchain_task.status = 'done'
|
||||
in_msg_blockchain_task.transaction_hash = transaction_hash
|
||||
in_msg_blockchain_task.transaction_lt = transaction_lt
|
||||
session.commit()
|
||||
|
||||
for blockchain_message in [transaction['in_msg']]:
|
||||
try:
|
||||
await process_incoming_message(blockchain_message)
|
||||
except BaseException as e:
|
||||
pass # make_log("TON_Daemon", f"Error while processing incoming message: {e}" + '\n' + traceback.format_exc(), level='debug')
|
||||
|
||||
# Generate mnemonic phrase
|
||||
mnemonic_response = await self.client.post(
|
||||
f"{self.api_endpoint}/wallet/generate",
|
||||
json={"testnet": self.testnet}
|
||||
)
|
||||
|
||||
if mnemonic_response.status_code != 200:
|
||||
error_msg = f"Failed to generate wallet: {mnemonic_response.text}"
|
||||
await logger.aerror("Wallet generation failed", error=error_msg)
|
||||
return {"error": error_msg}
|
||||
|
||||
mnemonic_data = mnemonic_response.json()
|
||||
|
||||
# Create wallet from mnemonic
|
||||
wallet_response = await self.client.post(
|
||||
f"{self.api_endpoint}/wallet/create",
|
||||
json={
|
||||
"mnemonic": mnemonic_data["mnemonic"],
|
||||
"testnet": self.testnet
|
||||
}
|
||||
)
|
||||
|
||||
if wallet_response.status_code != 200:
|
||||
error_msg = f"Failed to create wallet: {wallet_response.text}"
|
||||
await logger.aerror("Wallet creation failed", error=error_msg)
|
||||
return {"error": error_msg}
|
||||
|
||||
wallet_data = wallet_response.json()
|
||||
|
||||
await logger.ainfo(
|
||||
"Wallet created successfully",
|
||||
address=wallet_data.get("address"),
|
||||
testnet=self.testnet
|
||||
)
|
||||
|
||||
return {
|
||||
"address": wallet_data["address"],
|
||||
"private_key": wallet_data["private_key"],
|
||||
"mnemonic": mnemonic_data["mnemonic"],
|
||||
"testnet": self.testnet
|
||||
}
|
||||
|
||||
except httpx.TimeoutException:
|
||||
error_msg = "Wallet creation timeout"
|
||||
await logger.aerror(error_msg)
|
||||
return {"error": error_msg}
|
||||
except Exception as e:
|
||||
error_msg = f"Wallet creation error: {str(e)}"
|
||||
await logger.aerror("Wallet creation exception", error=str(e))
|
||||
return {"error": error_msg}
|
||||
|
||||
async def get_wallet_balance(self, address: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Get wallet balance with caching for performance.
|
||||
|
||||
Args:
|
||||
address: TON wallet address
|
||||
|
||||
Returns:
|
||||
Dict: Balance information
|
||||
"""
|
||||
try:
|
||||
# Check cache first
|
||||
cache_key = f"ton_balance:{address}"
|
||||
cached_balance = await self.cache_manager.get(cache_key)
|
||||
|
||||
if cached_balance:
|
||||
return cached_balance
|
||||
|
||||
# Fetch from blockchain
|
||||
balance_response = await self.client.get(
|
||||
f"{self.api_endpoint}/wallet/{address}/balance"
|
||||
)
|
||||
|
||||
if balance_response.status_code != 200:
|
||||
error_msg = f"Failed to get balance: {balance_response.text}"
|
||||
return {"error": error_msg}
|
||||
|
||||
balance_data = balance_response.json()
|
||||
|
||||
result = {
|
||||
"balance": int(balance_data.get("balance", 0)), # nanotons
|
||||
"last_transaction_lt": balance_data.get("last_transaction_lt"),
|
||||
"account_state": balance_data.get("account_state", "unknown"),
|
||||
"updated_at": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
# Cache for 30 seconds
|
||||
await self.cache_manager.set(cache_key, result, ttl=30)
|
||||
|
||||
return result
|
||||
|
||||
except httpx.TimeoutException:
|
||||
return {"error": "Balance fetch timeout"}
|
||||
except Exception as e:
|
||||
await logger.aerror("Balance fetch error", address=address, error=str(e))
|
||||
return {"error": f"Balance fetch error: {str(e)}"}
|
||||
|
||||
async def get_wallet_transactions(
|
||||
self,
|
||||
address: str,
|
||||
limit: int = 20,
|
||||
offset: int = 0
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Get wallet transaction history with pagination.
|
||||
|
||||
Args:
|
||||
address: TON wallet address
|
||||
limit: Number of transactions to fetch
|
||||
offset: Pagination offset
|
||||
|
||||
Returns:
|
||||
Dict: Transaction history
|
||||
"""
|
||||
try:
|
||||
# Check cache
|
||||
cache_key = f"ton_transactions:{address}:{limit}:{offset}"
|
||||
cached_transactions = await self.cache_manager.get(cache_key)
|
||||
|
||||
if cached_transactions:
|
||||
return cached_transactions
|
||||
|
||||
transactions_response = await self.client.get(
|
||||
f"{self.api_endpoint}/wallet/{address}/transactions",
|
||||
params={"limit": limit, "offset": offset}
|
||||
)
|
||||
|
||||
if transactions_response.status_code != 200:
|
||||
error_msg = f"Failed to get transactions: {transactions_response.text}"
|
||||
return {"error": error_msg}
|
||||
|
||||
transactions_data = transactions_response.json()
|
||||
|
||||
result = {
|
||||
"transactions": transactions_data.get("transactions", []),
|
||||
"total": transactions_data.get("total", 0),
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
"updated_at": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
# Cache for 1 minute
|
||||
await self.cache_manager.set(cache_key, result, ttl=60)
|
||||
|
||||
return result
|
||||
|
||||
except httpx.TimeoutException:
|
||||
return {"error": "Transaction fetch timeout"}
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Transaction fetch error",
|
||||
address=address,
|
||||
error=str(e)
|
||||
)
|
||||
return {"error": f"Transaction fetch error: {str(e)}"}
|
||||
|
||||
async def send_transaction(
|
||||
self,
|
||||
private_key: str,
|
||||
recipient_address: str,
|
||||
amount: int,
|
||||
message: str = "",
|
||||
**kwargs
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Send TON transaction with validation and monitoring.
|
||||
|
||||
Args:
|
||||
private_key: Encrypted private key
|
||||
recipient_address: Recipient wallet address
|
||||
amount: Amount in nanotons
|
||||
message: Optional message
|
||||
**kwargs: Additional transaction parameters
|
||||
|
||||
Returns:
|
||||
Dict: Transaction result
|
||||
"""
|
||||
try:
|
||||
# Validate inputs
|
||||
if amount <= 0:
|
||||
return {"error": "Amount must be positive"}
|
||||
|
||||
if len(recipient_address) != 48:
|
||||
return {"error": "Invalid recipient address format"}
|
||||
|
||||
# Decrypt private key
|
||||
try:
|
||||
decrypted_key = decrypt_data(private_key, context="wallet")
|
||||
if isinstance(decrypted_key, bytes):
|
||||
decrypted_key = decrypted_key.decode('utf-8')
|
||||
except Exception as e:
|
||||
await logger.aerror("Private key decryption failed", error=str(e))
|
||||
return {"error": "Invalid private key"}
|
||||
|
||||
# Prepare transaction
|
||||
transaction_data = {
|
||||
"private_key": decrypted_key,
|
||||
"recipient": recipient_address,
|
||||
"amount": str(amount),
|
||||
"message": message,
|
||||
"testnet": self.testnet
|
||||
}
|
||||
|
||||
# Send transaction
|
||||
tx_response = await self.client.post(
|
||||
f"{self.api_endpoint}/transaction/send",
|
||||
json=transaction_data
|
||||
)
|
||||
|
||||
if tx_response.status_code != 200:
|
||||
error_msg = f"Transaction failed: {tx_response.text}"
|
||||
await logger.aerror(
|
||||
"Transaction submission failed",
|
||||
recipient=recipient_address,
|
||||
amount=amount,
|
||||
error=error_msg
|
||||
)
|
||||
return {"error": error_msg}
|
||||
|
||||
tx_data = tx_response.json()
|
||||
|
||||
result = {
|
||||
"hash": tx_data["hash"],
|
||||
"lt": tx_data.get("lt"),
|
||||
"fee": tx_data.get("fee", 0),
|
||||
"block_hash": tx_data.get("block_hash"),
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
await logger.ainfo(
|
||||
"Transaction sent successfully",
|
||||
hash=result["hash"],
|
||||
recipient=recipient_address,
|
||||
amount=amount
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
except httpx.TimeoutException:
|
||||
return {"error": "Transaction timeout"}
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Transaction send error",
|
||||
recipient=recipient_address,
|
||||
amount=amount,
|
||||
error=str(e)
|
||||
)
|
||||
return {"error": f"Transaction error: {str(e)}"}
|
||||
|
||||
async def get_transaction_status(self, tx_hash: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Get transaction status and confirmation details.
|
||||
|
||||
Args:
|
||||
tx_hash: Transaction hash
|
||||
|
||||
Returns:
|
||||
Dict: Transaction status information
|
||||
"""
|
||||
try:
|
||||
# Check cache
|
||||
cache_key = f"ton_tx_status:{tx_hash}"
|
||||
cached_status = await self.cache_manager.get(cache_key)
|
||||
|
||||
if cached_status and cached_status.get("confirmed"):
|
||||
return cached_status
|
||||
|
||||
status_response = await self.client.get(
|
||||
f"{self.api_endpoint}/transaction/{tx_hash}/status"
|
||||
)
|
||||
|
||||
if status_response.status_code != 200:
|
||||
return {"error": f"Failed to get status: {status_response.text}"}
|
||||
|
||||
status_data = status_response.json()
|
||||
|
||||
result = {
|
||||
"hash": tx_hash,
|
||||
"confirmed": status_data.get("confirmed", False),
|
||||
"failed": status_data.get("failed", False),
|
||||
"confirmations": status_data.get("confirmations", 0),
|
||||
"block_hash": status_data.get("block_hash"),
|
||||
"block_time": status_data.get("block_time"),
|
||||
"fee": status_data.get("fee"),
|
||||
"confirmed_at": status_data.get("confirmed_at"),
|
||||
"updated_at": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
# Cache confirmed/failed transactions longer
|
||||
cache_ttl = 3600 if result["confirmed"] or result["failed"] else 30
|
||||
await self.cache_manager.set(cache_key, result, ttl=cache_ttl)
|
||||
|
||||
return result
|
||||
|
||||
except httpx.TimeoutException:
|
||||
return {"error": "Status check timeout"}
|
||||
except Exception as e:
|
||||
await logger.aerror("Status check error", tx_hash=tx_hash, error=str(e))
|
||||
return {"error": f"Status check error: {str(e)}"}
|
||||
|
||||
async def validate_address(self, address: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Validate TON address format and existence.
|
||||
|
||||
Args:
|
||||
address: TON address to validate
|
||||
|
||||
Returns:
|
||||
Dict: Validation result
|
||||
"""
|
||||
try:
|
||||
# Basic format validation
|
||||
if len(address) != 48:
|
||||
return {"valid": False, "error": "Invalid address length"}
|
||||
|
||||
# Check against blockchain
|
||||
validation_response = await self.client.post(
|
||||
f"{self.api_endpoint}/address/validate",
|
||||
json={"address": address}
|
||||
)
|
||||
|
||||
if validation_response.status_code != 200:
|
||||
return {"valid": False, "error": "Validation service error"}
|
||||
|
||||
validation_data = validation_response.json()
|
||||
|
||||
return {
|
||||
"valid": validation_data.get("valid", False),
|
||||
"exists": validation_data.get("exists", False),
|
||||
"account_type": validation_data.get("account_type"),
|
||||
"error": validation_data.get("error")
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Address validation error", address=address, error=str(e))
|
||||
return {"valid": False, "error": f"Validation error: {str(e)}"}
|
||||
|
||||
async def get_network_info(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Get TON network information and statistics.
|
||||
|
||||
Returns:
|
||||
Dict: Network information
|
||||
"""
|
||||
try:
|
||||
cache_key = "ton_network_info"
|
||||
cached_info = await self.cache_manager.get(cache_key)
|
||||
|
||||
if cached_info:
|
||||
return cached_info
|
||||
|
||||
network_response = await self.client.get(
|
||||
f"{self.api_endpoint}/network/info"
|
||||
)
|
||||
|
||||
if network_response.status_code != 200:
|
||||
return {"error": f"Failed to get network info: {network_response.text}"}
|
||||
|
||||
network_data = network_response.json()
|
||||
|
||||
result = {
|
||||
"network": "testnet" if self.testnet else "mainnet",
|
||||
"last_block": network_data.get("last_block"),
|
||||
"last_block_time": network_data.get("last_block_time"),
|
||||
"total_accounts": network_data.get("total_accounts"),
|
||||
"total_transactions": network_data.get("total_transactions"),
|
||||
"tps": network_data.get("tps"), # Transactions per second
|
||||
"updated_at": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
# Cache for 5 minutes
|
||||
await self.cache_manager.set(cache_key, result, ttl=300)
|
||||
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Network info error", error=str(e))
|
||||
return {"error": f"Network info error: {str(e)}"}
|
||||
|
||||
async def estimate_transaction_fee(
|
||||
self,
|
||||
sender_address: str,
|
||||
recipient_address: str,
|
||||
amount: int,
|
||||
message: str = ""
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Estimate transaction fee before sending.
|
||||
|
||||
Args:
|
||||
sender_address: Sender wallet address
|
||||
recipient_address: Recipient wallet address
|
||||
amount: Amount in nanotons
|
||||
message: Optional message
|
||||
|
||||
Returns:
|
||||
Dict: Fee estimation
|
||||
"""
|
||||
try:
|
||||
fee_response = await self.client.post(
|
||||
f"{self.api_endpoint}/transaction/estimate-fee",
|
||||
json={
|
||||
"sender": sender_address,
|
||||
"recipient": recipient_address,
|
||||
"amount": str(amount),
|
||||
"message": message
|
||||
}
|
||||
)
|
||||
|
||||
if fee_response.status_code != 200:
|
||||
return {"error": f"Fee estimation failed: {fee_response.text}"}
|
||||
|
||||
fee_data = fee_response.json()
|
||||
|
||||
return {
|
||||
"estimated_fee": fee_data.get("fee", 0),
|
||||
"estimated_fee_tons": str(Decimal(fee_data.get("fee", 0)) / Decimal("1000000000")),
|
||||
"gas_used": fee_data.get("gas_used"),
|
||||
"message_size": len(message.encode('utf-8')),
|
||||
"updated_at": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Fee estimation error", error=str(e))
|
||||
return {"error": f"Fee estimation error: {str(e)}"}
|
||||
|
||||
async def monitor_transaction(self, tx_hash: str, max_wait_time: int = 300) -> Dict[str, Any]:
|
||||
"""
|
||||
Monitor transaction until confirmation or timeout.
|
||||
|
||||
Args:
|
||||
tx_hash: Transaction hash to monitor
|
||||
max_wait_time: Maximum wait time in seconds
|
||||
|
||||
Returns:
|
||||
Dict: Final transaction status
|
||||
"""
|
||||
start_time = datetime.utcnow()
|
||||
check_interval = 5 # Check every 5 seconds
|
||||
|
||||
while (datetime.utcnow() - start_time).seconds < max_wait_time:
|
||||
status = await self.get_transaction_status(tx_hash)
|
||||
|
||||
if status.get("error"):
|
||||
return status
|
||||
|
||||
if status.get("confirmed") or status.get("failed"):
|
||||
await logger.ainfo(
|
||||
"Transaction monitoring completed",
|
||||
tx_hash=tx_hash,
|
||||
confirmed=status.get("confirmed"),
|
||||
failed=status.get("failed"),
|
||||
duration=(datetime.utcnow() - start_time).seconds
|
||||
)
|
||||
return status
|
||||
|
||||
await asyncio.sleep(check_interval)
|
||||
|
||||
# Timeout reached
|
||||
await logger.awarning(
|
||||
"Transaction monitoring timeout",
|
||||
tx_hash=tx_hash,
|
||||
max_wait_time=max_wait_time
|
||||
)
|
||||
|
||||
return {
|
||||
"hash": tx_hash,
|
||||
"confirmed": False,
|
||||
"timeout": True,
|
||||
"error": "Monitoring timeout reached"
|
||||
}
|
||||
|
||||
async def get_smart_contract_info(self, address: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Get smart contract information and ABI.
|
||||
|
||||
Args:
|
||||
address: Smart contract address
|
||||
|
||||
Returns:
|
||||
Dict: Contract information
|
||||
"""
|
||||
try:
|
||||
cache_key = f"ton_contract:{address}"
|
||||
cached_info = await self.cache_manager.get(cache_key)
|
||||
|
||||
if cached_info:
|
||||
return cached_info
|
||||
|
||||
contract_response = await self.client.get(
|
||||
f"{self.api_endpoint}/contract/{address}/info"
|
||||
)
|
||||
|
||||
if contract_response.status_code != 200:
|
||||
return {"error": f"Failed to get contract info: {contract_response.text}"}
|
||||
|
||||
contract_data = contract_response.json()
|
||||
|
||||
result = {
|
||||
"address": address,
|
||||
"contract_type": contract_data.get("contract_type"),
|
||||
"is_verified": contract_data.get("is_verified", False),
|
||||
"abi": contract_data.get("abi"),
|
||||
"source_code": contract_data.get("source_code"),
|
||||
"compiler_version": contract_data.get("compiler_version"),
|
||||
"deployment_block": contract_data.get("deployment_block"),
|
||||
"updated_at": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
# Cache for 1 hour
|
||||
await self.cache_manager.set(cache_key, result, ttl=3600)
|
||||
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Contract info error", address=address, error=str(e))
|
||||
return {"error": f"Contract info error: {str(e)}"}
|
||||
|
||||
async def call_smart_contract(
|
||||
self,
|
||||
contract_address: str,
|
||||
method: str,
|
||||
params: Dict[str, Any],
|
||||
private_key: Optional[str] = None
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Call smart contract method.
|
||||
|
||||
Args:
|
||||
contract_address: Contract address
|
||||
method: Method name to call
|
||||
params: Method parameters
|
||||
private_key: Private key for write operations
|
||||
|
||||
Returns:
|
||||
Dict: Contract call result
|
||||
"""
|
||||
try:
|
||||
call_data = {
|
||||
"contract": contract_address,
|
||||
"method": method,
|
||||
"params": params
|
||||
}
|
||||
|
||||
# Add private key for write operations
|
||||
if private_key:
|
||||
try:
|
||||
sw_transactions = await toncenter.get_transactions(highload_wallet.address.to_string(1, 1, 1), limit=100)
|
||||
for sw_transaction in sw_transactions:
|
||||
try:
|
||||
await process_incoming_transaction(sw_transaction)
|
||||
except BaseException as e:
|
||||
make_log("TON_Daemon", f"Error while processing incoming transaction: {e}", level="debug")
|
||||
except BaseException as e:
|
||||
make_log("TON_Daemon", f"Error while getting service wallet transactions: {e}", level="ERROR")
|
||||
|
||||
await send_status("ton_daemon", f"working: processing out-txs (seqno={sw_seqno_value})")
|
||||
# Отправка подписанных сообщений
|
||||
for blockchain_task in (
|
||||
session.query(BlockchainTask).filter(
|
||||
BlockchainTask.status == 'processing',
|
||||
).order_by(BlockchainTask.updated.asc()).all()
|
||||
):
|
||||
make_log("TON_Daemon", f"Processing task (processing) {blockchain_task.id}")
|
||||
query_boc = bytes.fromhex(blockchain_task.meta['signed_message'])
|
||||
errors_list = []
|
||||
|
||||
try:
|
||||
await toncenter.send_boc(query_boc)
|
||||
except BaseException as e:
|
||||
errors_list.append(f"{e}")
|
||||
|
||||
try:
|
||||
make_log("TON_Daemon", str(
|
||||
httpx.post(
|
||||
'https://tonapi.io/v2/blockchain/message',
|
||||
json={
|
||||
'boc': query_boc.hex()
|
||||
}
|
||||
).text
|
||||
))
|
||||
except BaseException as e:
|
||||
make_log("TON_Daemon", f"Error while pushing task to tonkeeper ({blockchain_task.id}): {e}", level="ERROR")
|
||||
errors_list.append(f"{e}")
|
||||
|
||||
blockchain_task.updated = datetime.utcnow()
|
||||
|
||||
if blockchain_task.meta['sign_created'] + 10 * 60 < datetime.utcnow().timestamp():
|
||||
# or sum([int("terminating vm with exit code 36" in e) for e in errors_list]) > 0:
|
||||
make_log("TON_Daemon", f"Task {blockchain_task.id} done", level="DEBUG")
|
||||
blockchain_task.status = 'done'
|
||||
session.commit()
|
||||
continue
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
await send_status("ton_daemon", f"working: creating new messages (seqno={sw_seqno_value})")
|
||||
# Создание новых подписей
|
||||
for blockchain_task in (
|
||||
session.query(BlockchainTask).filter(BlockchainTask.status == 'wait').all()
|
||||
):
|
||||
try:
|
||||
# Check processing tasks in current epoch < 3_000_000
|
||||
if (
|
||||
session.query(BlockchainTask).filter(
|
||||
BlockchainTask.epoch == blockchain_task.epoch,
|
||||
).count() > 3_000_000
|
||||
):
|
||||
make_log("TON", f"Too many processing tasks in epoch {blockchain_task.epoch}", level="error")
|
||||
await send_status("ton_daemon", f"working: too many tasks in epoch {blockchain_task.epoch}")
|
||||
await asyncio.sleep(5)
|
||||
continue
|
||||
|
||||
sign_created = int(datetime.utcnow().timestamp()) - 60
|
||||
try:
|
||||
current_epoch = int(datetime.utcnow().timestamp() // (60 * 60))
|
||||
max_epoch_seqno = (
|
||||
session.query(func.max(BlockchainTask.seqno)).filter(
|
||||
BlockchainTask.epoch == current_epoch
|
||||
).scalar() or 0
|
||||
)
|
||||
current_epoch_shift = 3_000_000 if current_epoch % 2 == 0 else 0
|
||||
current_seqno = max_epoch_seqno + 1 + (current_epoch_shift if max_epoch_seqno == 0 else 0)
|
||||
except BaseException as e:
|
||||
make_log("CRITICAL", f"Error calculating epoch,seqno: {e}", level="error")
|
||||
current_epoch = 0
|
||||
current_seqno = 0
|
||||
|
||||
blockchain_task.seqno = current_seqno
|
||||
blockchain_task.epoch = current_epoch
|
||||
blockchain_task.status = 'processing'
|
||||
try:
|
||||
query = highload_wallet.create_transfer_message(
|
||||
blockchain_task.destination, int(blockchain_task.amount), HighloadQueryId.from_seqno(current_seqno),
|
||||
sign_created, send_mode=1,
|
||||
payload=Cell.one_from_boc(b64decode(blockchain_task.payload))
|
||||
)
|
||||
query_boc = query['message'].to_boc(False)
|
||||
except BaseException as e:
|
||||
make_log("TON", f"Error creating transfer message: {e}", level="error")
|
||||
query_boc = begin_cell().end_cell().to_boc(False)
|
||||
|
||||
blockchain_task.meta = {
|
||||
**blockchain_task.meta,
|
||||
'sign_created': sign_created,
|
||||
'signed_message': query_boc.hex(),
|
||||
}
|
||||
session.commit()
|
||||
make_log("TON", f"Created signed message for task {blockchain_task.id}" + '\n' + traceback.format_exc(), level="info")
|
||||
except BaseException as e:
|
||||
make_log("TON", f"Error processing task {blockchain_task.id}: {e}" + '\n' + traceback.format_exc(), level="error")
|
||||
continue
|
||||
|
||||
await asyncio.sleep(1)
|
||||
|
||||
await asyncio.sleep(1)
|
||||
await send_status("ton_daemon", f"working (seqno={sw_seqno_value})")
|
||||
except BaseException as e:
|
||||
make_log("TON", f"Error: {e}", level="error")
|
||||
await asyncio.sleep(3)
|
||||
|
||||
# if __name__ == '__main__':
|
||||
# loop = asyncio.get_event_loop()
|
||||
# loop.run_until_complete(main())
|
||||
# loop.close()
|
||||
|
||||
decrypted_key = decrypt_data(private_key, context="wallet")
|
||||
if isinstance(decrypted_key, bytes):
|
||||
decrypted_key = decrypted_key.decode('utf-8')
|
||||
call_data["private_key"] = decrypted_key
|
||||
except Exception as e:
|
||||
return {"error": "Invalid private key"}
|
||||
|
||||
contract_response = await self.client.post(
|
||||
f"{self.api_endpoint}/contract/call",
|
||||
json=call_data
|
||||
)
|
||||
|
||||
if contract_response.status_code != 200:
|
||||
return {"error": f"Contract call failed: {contract_response.text}"}
|
||||
|
||||
call_result = contract_response.json()
|
||||
|
||||
await logger.ainfo(
|
||||
"Smart contract called",
|
||||
contract=contract_address,
|
||||
method=method,
|
||||
success=call_result.get("success", False)
|
||||
)
|
||||
|
||||
return call_result
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Contract call error",
|
||||
contract=contract_address,
|
||||
method=method,
|
||||
error=str(e)
|
||||
)
|
||||
return {"error": f"Contract call error: {str(e)}"}
|
||||
|
||||
# Global TON service instance
|
||||
_ton_service = None
|
||||
|
||||
async def get_ton_service() -> TONService:
|
||||
"""Get or create global TON service instance."""
|
||||
global _ton_service
|
||||
if _ton_service is None:
|
||||
_ton_service = TONService()
|
||||
return _ton_service
|
||||
|
||||
async def cleanup_ton_service():
|
||||
"""Cleanup global TON service instance."""
|
||||
global _ton_service
|
||||
if _ton_service:
|
||||
await _ton_service.close()
|
||||
_ton_service = None
|
||||
@@ -0,0 +1,385 @@
|
||||
"""Redis caching system with fallback support."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import pickle
|
||||
from typing import Any, Optional, Union, Dict, List
|
||||
from contextlib import asynccontextmanager
|
||||
from functools import wraps
|
||||
|
||||
import redis.asyncio as redis
|
||||
from redis.asyncio import ConnectionPool
|
||||
|
||||
from app.core.config_compatible import get_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Global Redis connection pool
|
||||
_redis_pool: Optional[ConnectionPool] = None
|
||||
_redis_client: Optional[redis.Redis] = None
|
||||
|
||||
|
||||
class CacheError(Exception):
|
||||
"""Custom cache error."""
|
||||
pass
|
||||
|
||||
|
||||
async def init_cache() -> None:
|
||||
"""Initialize Redis cache connection."""
|
||||
global _redis_pool, _redis_client
|
||||
|
||||
settings = get_settings()
|
||||
|
||||
if not settings.redis_enabled or not settings.cache_enabled:
|
||||
logger.info("Redis caching is disabled")
|
||||
return
|
||||
|
||||
try:
|
||||
# Create connection pool
|
||||
_redis_pool = ConnectionPool(
|
||||
host=settings.redis_host,
|
||||
port=settings.redis_port,
|
||||
password=settings.redis_password,
|
||||
db=settings.redis_db,
|
||||
max_connections=settings.redis_max_connections,
|
||||
socket_timeout=settings.redis_socket_timeout,
|
||||
socket_connect_timeout=settings.redis_socket_connect_timeout,
|
||||
decode_responses=False, # We'll handle encoding manually for flexibility
|
||||
retry_on_timeout=True,
|
||||
health_check_interval=30,
|
||||
)
|
||||
|
||||
# Create Redis client
|
||||
_redis_client = redis.Redis(connection_pool=_redis_pool)
|
||||
|
||||
# Test connection
|
||||
await _redis_client.ping()
|
||||
|
||||
logger.info(f"Redis cache initialized successfully at {settings.redis_host}:{settings.redis_port}")
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to initialize Redis cache: {e}. Caching will be disabled.")
|
||||
_redis_pool = None
|
||||
_redis_client = None
|
||||
|
||||
|
||||
async def close_cache() -> None:
|
||||
"""Close Redis cache connection."""
|
||||
global _redis_pool, _redis_client
|
||||
|
||||
if _redis_client:
|
||||
try:
|
||||
await _redis_client.close()
|
||||
logger.info("Redis cache connection closed")
|
||||
except Exception as e:
|
||||
logger.error(f"Error closing Redis cache: {e}")
|
||||
finally:
|
||||
_redis_client = None
|
||||
_redis_pool = None
|
||||
|
||||
|
||||
def get_redis_client() -> Optional[redis.Redis]:
|
||||
"""Get Redis client instance."""
|
||||
return _redis_client
|
||||
|
||||
|
||||
def is_cache_available() -> bool:
|
||||
"""Check if cache is available."""
|
||||
return _redis_client is not None
|
||||
|
||||
|
||||
class Cache:
|
||||
"""Redis cache manager with fallback support."""
|
||||
|
||||
def __init__(self):
|
||||
self.settings = get_settings()
|
||||
|
||||
def _serialize(self, value: Any) -> bytes:
|
||||
"""Serialize value for storage."""
|
||||
try:
|
||||
if isinstance(value, (str, int, float, bool)):
|
||||
return json.dumps(value).encode('utf-8')
|
||||
else:
|
||||
return pickle.dumps(value)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to serialize cache value: {e}")
|
||||
raise CacheError(f"Serialization error: {e}")
|
||||
|
||||
def _deserialize(self, data: bytes) -> Any:
|
||||
"""Deserialize value from storage."""
|
||||
try:
|
||||
# Try JSON first (for simple types)
|
||||
try:
|
||||
return json.loads(data.decode('utf-8'))
|
||||
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||
# Fallback to pickle for complex objects
|
||||
return pickle.loads(data)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to deserialize cache value: {e}")
|
||||
raise CacheError(f"Deserialization error: {e}")
|
||||
|
||||
def _make_key(self, key: str, prefix: str = "myuploader") -> str:
|
||||
"""Create cache key with prefix."""
|
||||
return f"{prefix}:{key}"
|
||||
|
||||
async def get(self, key: str, default: Any = None) -> Any:
|
||||
"""Get value from cache."""
|
||||
if not is_cache_available():
|
||||
return default
|
||||
|
||||
try:
|
||||
redis_key = self._make_key(key)
|
||||
data = await _redis_client.get(redis_key)
|
||||
|
||||
if data is None:
|
||||
return default
|
||||
|
||||
return self._deserialize(data)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Cache get error for key '{key}': {e}")
|
||||
return default
|
||||
|
||||
async def set(self, key: str, value: Any, ttl: Optional[int] = None) -> bool:
|
||||
"""Set value in cache."""
|
||||
if not is_cache_available():
|
||||
return False
|
||||
|
||||
try:
|
||||
redis_key = self._make_key(key)
|
||||
data = self._serialize(value)
|
||||
|
||||
if ttl is None:
|
||||
ttl = self.settings.cache_default_ttl
|
||||
|
||||
await _redis_client.setex(redis_key, ttl, data)
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Cache set error for key '{key}': {e}")
|
||||
return False
|
||||
|
||||
async def delete(self, key: str) -> bool:
|
||||
"""Delete value from cache."""
|
||||
if not is_cache_available():
|
||||
return False
|
||||
|
||||
try:
|
||||
redis_key = self._make_key(key)
|
||||
result = await _redis_client.delete(redis_key)
|
||||
return bool(result)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Cache delete error for key '{key}': {e}")
|
||||
return False
|
||||
|
||||
async def exists(self, key: str) -> bool:
|
||||
"""Check if key exists in cache."""
|
||||
if not is_cache_available():
|
||||
return False
|
||||
|
||||
try:
|
||||
redis_key = self._make_key(key)
|
||||
result = await _redis_client.exists(redis_key)
|
||||
return bool(result)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Cache exists error for key '{key}': {e}")
|
||||
return False
|
||||
|
||||
async def expire(self, key: str, ttl: int) -> bool:
|
||||
"""Set expiration time for key."""
|
||||
if not is_cache_available():
|
||||
return False
|
||||
|
||||
try:
|
||||
redis_key = self._make_key(key)
|
||||
result = await _redis_client.expire(redis_key, ttl)
|
||||
return bool(result)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Cache expire error for key '{key}': {e}")
|
||||
return False
|
||||
|
||||
async def clear_pattern(self, pattern: str) -> int:
|
||||
"""Clear all keys matching pattern."""
|
||||
if not is_cache_available():
|
||||
return 0
|
||||
|
||||
try:
|
||||
redis_pattern = self._make_key(pattern)
|
||||
keys = await _redis_client.keys(redis_pattern)
|
||||
|
||||
if keys:
|
||||
result = await _redis_client.delete(*keys)
|
||||
return result
|
||||
return 0
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Cache clear pattern error for pattern '{pattern}': {e}")
|
||||
return 0
|
||||
|
||||
async def increment(self, key: str, amount: int = 1, ttl: Optional[int] = None) -> Optional[int]:
|
||||
"""Increment counter in cache."""
|
||||
if not is_cache_available():
|
||||
return None
|
||||
|
||||
try:
|
||||
redis_key = self._make_key(key)
|
||||
result = await _redis_client.incrby(redis_key, amount)
|
||||
|
||||
if ttl is not None:
|
||||
await _redis_client.expire(redis_key, ttl)
|
||||
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Cache increment error for key '{key}': {e}")
|
||||
return None
|
||||
|
||||
async def get_multiple(self, keys: List[str]) -> Dict[str, Any]:
|
||||
"""Get multiple values from cache."""
|
||||
if not is_cache_available():
|
||||
return {}
|
||||
|
||||
try:
|
||||
redis_keys = [self._make_key(key) for key in keys]
|
||||
values = await _redis_client.mget(redis_keys)
|
||||
|
||||
result = {}
|
||||
for i, (key, data) in enumerate(zip(keys, values)):
|
||||
if data is not None:
|
||||
try:
|
||||
result[key] = self._deserialize(data)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to deserialize cached value for key '{key}': {e}")
|
||||
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Cache get_multiple error: {e}")
|
||||
return {}
|
||||
|
||||
async def set_multiple(self, mapping: Dict[str, Any], ttl: Optional[int] = None) -> bool:
|
||||
"""Set multiple values in cache."""
|
||||
if not is_cache_available():
|
||||
return False
|
||||
|
||||
try:
|
||||
pipeline = _redis_client.pipeline()
|
||||
|
||||
for key, value in mapping.items():
|
||||
redis_key = self._make_key(key)
|
||||
data = self._serialize(value)
|
||||
|
||||
if ttl is None:
|
||||
ttl = self.settings.cache_default_ttl
|
||||
|
||||
pipeline.setex(redis_key, ttl, data)
|
||||
|
||||
await pipeline.execute()
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Cache set_multiple error: {e}")
|
||||
return False
|
||||
|
||||
|
||||
# Global cache instance
|
||||
cache = Cache()
|
||||
|
||||
|
||||
# Caching decorators
|
||||
def cached(ttl: Optional[int] = None, key_prefix: str = "func"):
|
||||
"""Decorator for caching function results."""
|
||||
def decorator(func):
|
||||
@wraps(func)
|
||||
async def wrapper(*args, **kwargs):
|
||||
if not is_cache_available():
|
||||
return await func(*args, **kwargs)
|
||||
|
||||
# Create cache key from function name and arguments
|
||||
key_parts = [key_prefix, func.__name__]
|
||||
if args:
|
||||
key_parts.extend([str(arg) for arg in args])
|
||||
if kwargs:
|
||||
key_parts.extend([f"{k}={v}" for k, v in sorted(kwargs.items())])
|
||||
|
||||
cache_key = ":".join(key_parts)
|
||||
|
||||
# Try to get from cache
|
||||
result = await cache.get(cache_key)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
# Call function and cache result
|
||||
result = await func(*args, **kwargs)
|
||||
await cache.set(cache_key, result, ttl)
|
||||
return result
|
||||
|
||||
return wrapper
|
||||
return decorator
|
||||
|
||||
|
||||
def cache_user_data(ttl: Optional[int] = None):
|
||||
"""Decorator for caching user-specific data."""
|
||||
if ttl is None:
|
||||
ttl = get_settings().cache_user_ttl
|
||||
return cached(ttl=ttl, key_prefix="user")
|
||||
|
||||
|
||||
def cache_content_data(ttl: Optional[int] = None):
|
||||
"""Decorator for caching content data."""
|
||||
if ttl is None:
|
||||
ttl = get_settings().cache_content_ttl
|
||||
return cached(ttl=ttl, key_prefix="content")
|
||||
|
||||
|
||||
# Cache health check
|
||||
async def check_cache_health() -> Dict[str, Any]:
|
||||
"""Check cache health and return status."""
|
||||
if not is_cache_available():
|
||||
return {
|
||||
"status": "disabled",
|
||||
"available": False,
|
||||
"error": "Redis not initialized"
|
||||
}
|
||||
|
||||
try:
|
||||
# Test basic operations
|
||||
test_key = "health_check"
|
||||
test_value = {"timestamp": "test"}
|
||||
|
||||
await cache.set(test_key, test_value, 10)
|
||||
retrieved = await cache.get(test_key)
|
||||
await cache.delete(test_key)
|
||||
|
||||
# Get Redis info
|
||||
info = await _redis_client.info()
|
||||
|
||||
return {
|
||||
"status": "healthy",
|
||||
"available": True,
|
||||
"test_passed": retrieved == test_value,
|
||||
"connected_clients": info.get("connected_clients", 0),
|
||||
"used_memory": info.get("used_memory_human", "unknown"),
|
||||
"total_commands_processed": info.get("total_commands_processed", 0),
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
return {
|
||||
"status": "error",
|
||||
"available": False,
|
||||
"error": str(e)
|
||||
}
|
||||
|
||||
|
||||
# Context manager for cache operations
|
||||
@asynccontextmanager
|
||||
async def cache_context():
|
||||
"""Context manager for cache operations."""
|
||||
try:
|
||||
yield cache
|
||||
except Exception as e:
|
||||
logger.error(f"Cache context error: {e}")
|
||||
raise
|
||||
@@ -0,0 +1,253 @@
|
||||
"""
|
||||
Application configuration with security improvements and validation
|
||||
"""
|
||||
import os
|
||||
import secrets
|
||||
from datetime import datetime
|
||||
from typing import List, Optional, Dict, Any
|
||||
from pathlib import Path
|
||||
|
||||
from pydantic import BaseSettings, validator, Field
|
||||
from pydantic.networks import AnyHttpUrl, PostgresDsn, RedisDsn
|
||||
import structlog
|
||||
|
||||
logger = structlog.get_logger(__name__)
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
"""Application settings with validation"""
|
||||
|
||||
# Application
|
||||
PROJECT_NAME: str = "My Uploader Bot"
|
||||
PROJECT_VERSION: str = "2.0.0"
|
||||
PROJECT_HOST: AnyHttpUrl = Field(default="http://127.0.0.1:15100")
|
||||
SANIC_PORT: int = Field(default=15100, ge=1000, le=65535)
|
||||
DEBUG: bool = Field(default=False)
|
||||
|
||||
# Security
|
||||
SECRET_KEY: str = Field(default_factory=lambda: secrets.token_urlsafe(32))
|
||||
JWT_SECRET_KEY: str = Field(default_factory=lambda: secrets.token_urlsafe(32))
|
||||
JWT_EXPIRE_MINUTES: int = Field(default=60 * 24 * 7) # 7 days
|
||||
ENCRYPTION_KEY: Optional[str] = None
|
||||
|
||||
# Rate Limiting
|
||||
RATE_LIMIT_REQUESTS: int = Field(default=100)
|
||||
RATE_LIMIT_WINDOW: int = Field(default=60) # seconds
|
||||
RATE_LIMIT_ENABLED: bool = Field(default=True)
|
||||
|
||||
# Database
|
||||
DATABASE_URL: PostgresDsn = Field(
|
||||
default="postgresql+asyncpg://user:password@localhost:5432/uploader_bot"
|
||||
)
|
||||
DATABASE_POOL_SIZE: int = Field(default=10, ge=1, le=100)
|
||||
DATABASE_MAX_OVERFLOW: int = Field(default=20, ge=0, le=100)
|
||||
DATABASE_ECHO: bool = Field(default=False)
|
||||
|
||||
# Redis
|
||||
REDIS_URL: RedisDsn = Field(default="redis://localhost:6379/0")
|
||||
REDIS_POOL_SIZE: int = Field(default=10, ge=1, le=100)
|
||||
REDIS_TTL_DEFAULT: int = Field(default=3600) # 1 hour
|
||||
REDIS_TTL_SHORT: int = Field(default=300) # 5 minutes
|
||||
REDIS_TTL_LONG: int = Field(default=86400) # 24 hours
|
||||
|
||||
# File Storage
|
||||
UPLOADS_DIR: Path = Field(default=Path("/app/data"))
|
||||
MAX_FILE_SIZE: int = Field(default=100 * 1024 * 1024) # 100MB
|
||||
ALLOWED_CONTENT_TYPES: List[str] = Field(default=[
|
||||
'image/jpeg', 'image/png', 'image/gif', 'image/webp',
|
||||
'video/mp4', 'video/webm', 'video/ogg', 'video/quicktime',
|
||||
'audio/mpeg', 'audio/ogg', 'audio/wav', 'audio/mp4',
|
||||
'text/plain', 'application/json'
|
||||
])
|
||||
|
||||
# Telegram
|
||||
TELEGRAM_API_KEY: str = Field(..., min_length=40)
|
||||
CLIENT_TELEGRAM_API_KEY: str = Field(..., min_length=40)
|
||||
TELEGRAM_WEBHOOK_ENABLED: bool = Field(default=False)
|
||||
TELEGRAM_WEBHOOK_URL: Optional[AnyHttpUrl] = None
|
||||
TELEGRAM_WEBHOOK_SECRET: str = Field(default_factory=lambda: secrets.token_urlsafe(32))
|
||||
|
||||
# TON Blockchain
|
||||
TESTNET: bool = Field(default=False)
|
||||
TONCENTER_HOST: AnyHttpUrl = Field(default="https://toncenter.com/api/v2/")
|
||||
TONCENTER_API_KEY: Optional[str] = None
|
||||
TONCENTER_V3_HOST: AnyHttpUrl = Field(default="https://toncenter.com/api/v3/")
|
||||
MY_PLATFORM_CONTRACT: str = Field(default="EQDmWp6hbJlYUrXZKb9N88sOrTit630ZuRijfYdXEHLtheMY")
|
||||
MY_FUND_ADDRESS: str = Field(default="UQDarChHFMOI2On9IdHJNeEKttqepgo0AY4bG1trw8OAAwMY")
|
||||
|
||||
# Logging
|
||||
LOG_LEVEL: str = Field(default="INFO", regex="^(DEBUG|INFO|WARNING|ERROR|CRITICAL)$")
|
||||
LOG_DIR: Path = Field(default=Path("logs"))
|
||||
LOG_FORMAT: str = Field(default="json")
|
||||
LOG_ROTATION: str = Field(default="1 day")
|
||||
LOG_RETENTION: str = Field(default="30 days")
|
||||
|
||||
# Monitoring
|
||||
METRICS_ENABLED: bool = Field(default=True)
|
||||
METRICS_PORT: int = Field(default=9090, ge=1000, le=65535)
|
||||
HEALTH_CHECK_ENABLED: bool = Field(default=True)
|
||||
|
||||
# Background Services
|
||||
INDEXER_ENABLED: bool = Field(default=True)
|
||||
INDEXER_INTERVAL: int = Field(default=5, ge=1, le=3600)
|
||||
TON_DAEMON_ENABLED: bool = Field(default=True)
|
||||
TON_DAEMON_INTERVAL: int = Field(default=3, ge=1, le=3600)
|
||||
LICENSE_SERVICE_ENABLED: bool = Field(default=True)
|
||||
LICENSE_SERVICE_INTERVAL: int = Field(default=10, ge=1, le=3600)
|
||||
CONVERT_SERVICE_ENABLED: bool = Field(default=True)
|
||||
CONVERT_SERVICE_INTERVAL: int = Field(default=30, ge=1, le=3600)
|
||||
|
||||
# Web App URLs
|
||||
WEB_APP_URLS: Dict[str, str] = Field(default={
|
||||
'uploadContent': "https://web2-client.vercel.app/uploadContent"
|
||||
})
|
||||
|
||||
# Maintenance
|
||||
MAINTENANCE_MODE: bool = Field(default=False)
|
||||
MAINTENANCE_MESSAGE: str = Field(default="System is under maintenance")
|
||||
|
||||
# Development
|
||||
MOCK_EXTERNAL_SERVICES: bool = Field(default=False)
|
||||
DISABLE_WEBHOOKS: bool = Field(default=False)
|
||||
|
||||
@validator('UPLOADS_DIR')
|
||||
def create_uploads_dir(cls, v):
|
||||
"""Create uploads directory if it doesn't exist"""
|
||||
if not v.exists():
|
||||
v.mkdir(parents=True, exist_ok=True)
|
||||
return v
|
||||
|
||||
@validator('LOG_DIR')
|
||||
def create_log_dir(cls, v):
|
||||
"""Create log directory if it doesn't exist"""
|
||||
if not v.exists():
|
||||
v.mkdir(parents=True, exist_ok=True)
|
||||
return v
|
||||
|
||||
@validator('DATABASE_URL')
|
||||
def validate_database_url(cls, v):
|
||||
"""Validate database URL format"""
|
||||
if not str(v).startswith('postgresql+asyncpg://'):
|
||||
raise ValueError('Database URL must use asyncpg driver')
|
||||
return v
|
||||
|
||||
@validator('TELEGRAM_API_KEY', 'CLIENT_TELEGRAM_API_KEY')
|
||||
def validate_telegram_keys(cls, v):
|
||||
"""Validate Telegram bot tokens format"""
|
||||
parts = v.split(':')
|
||||
if len(parts) != 2 or not parts[0].isdigit() or len(parts[1]) != 35:
|
||||
raise ValueError('Invalid Telegram bot token format')
|
||||
return v
|
||||
|
||||
@validator('SECRET_KEY', 'JWT_SECRET_KEY')
|
||||
def validate_secret_keys(cls, v):
|
||||
"""Validate secret keys length"""
|
||||
if len(v) < 32:
|
||||
raise ValueError('Secret keys must be at least 32 characters long')
|
||||
return v
|
||||
|
||||
class Config:
|
||||
env_file = ".env"
|
||||
case_sensitive = True
|
||||
validate_assignment = True
|
||||
|
||||
|
||||
class SecurityConfig:
|
||||
"""Security-related configurations"""
|
||||
|
||||
# CORS settings
|
||||
CORS_ORIGINS = [
|
||||
"https://web2-client.vercel.app",
|
||||
"https://t.me",
|
||||
"https://web.telegram.org"
|
||||
]
|
||||
|
||||
# Content Security Policy
|
||||
CSP_DIRECTIVES = {
|
||||
'default-src': ["'self'"],
|
||||
'script-src': ["'self'", "'unsafe-inline'", "https://cdn.jsdelivr.net"],
|
||||
'style-src': ["'self'", "'unsafe-inline'", "https://cdn.jsdelivr.net"],
|
||||
'img-src': ["'self'", "data:", "https:"],
|
||||
'connect-src': ["'self'", "https://api.telegram.org"],
|
||||
'frame-ancestors': ["'none'"],
|
||||
'form-action': ["'self'"],
|
||||
'base-uri': ["'self'"]
|
||||
}
|
||||
|
||||
# Request size limits
|
||||
MAX_REQUEST_SIZE = 100 * 1024 * 1024 # 100MB
|
||||
MAX_JSON_SIZE = 10 * 1024 * 1024 # 10MB
|
||||
|
||||
# Session settings
|
||||
SESSION_COOKIE_SECURE = True
|
||||
SESSION_COOKIE_HTTPONLY = True
|
||||
SESSION_COOKIE_SAMESITE = "Strict"
|
||||
|
||||
# Rate limiting patterns
|
||||
RATE_LIMIT_PATTERNS = {
|
||||
"auth": {"requests": 5, "window": 300}, # 5 requests per 5 minutes
|
||||
"upload": {"requests": 10, "window": 3600}, # 10 uploads per hour
|
||||
"api": {"requests": 100, "window": 60}, # 100 API calls per minute
|
||||
"heavy": {"requests": 1, "window": 60} # 1 heavy operation per minute
|
||||
}
|
||||
|
||||
|
||||
# Create settings instance
|
||||
settings = Settings()
|
||||
|
||||
# Expose commonly used settings
|
||||
DATABASE_URL = str(settings.DATABASE_URL)
|
||||
REDIS_URL = str(settings.REDIS_URL)
|
||||
DATABASE_POOL_SIZE = settings.DATABASE_POOL_SIZE
|
||||
DATABASE_MAX_OVERFLOW = settings.DATABASE_MAX_OVERFLOW
|
||||
REDIS_POOL_SIZE = settings.REDIS_POOL_SIZE
|
||||
|
||||
TELEGRAM_API_KEY = settings.TELEGRAM_API_KEY
|
||||
CLIENT_TELEGRAM_API_KEY = settings.CLIENT_TELEGRAM_API_KEY
|
||||
PROJECT_HOST = str(settings.PROJECT_HOST)
|
||||
SANIC_PORT = settings.SANIC_PORT
|
||||
UPLOADS_DIR = settings.UPLOADS_DIR
|
||||
ALLOWED_CONTENT_TYPES = settings.ALLOWED_CONTENT_TYPES
|
||||
|
||||
TESTNET = settings.TESTNET
|
||||
TONCENTER_HOST = str(settings.TONCENTER_HOST)
|
||||
TONCENTER_API_KEY = settings.TONCENTER_API_KEY
|
||||
TONCENTER_V3_HOST = str(settings.TONCENTER_V3_HOST)
|
||||
MY_PLATFORM_CONTRACT = settings.MY_PLATFORM_CONTRACT
|
||||
MY_FUND_ADDRESS = settings.MY_FUND_ADDRESS
|
||||
|
||||
LOG_LEVEL = settings.LOG_LEVEL
|
||||
LOG_DIR = settings.LOG_DIR
|
||||
MAINTENANCE_MODE = settings.MAINTENANCE_MODE
|
||||
|
||||
# Cache keys patterns
|
||||
CACHE_KEYS = {
|
||||
"user_session": "user:session:{user_id}",
|
||||
"user_data": "user:data:{user_id}",
|
||||
"content_metadata": "content:meta:{content_id}",
|
||||
"rate_limit": "rate_limit:{pattern}:{identifier}",
|
||||
"blockchain_task": "blockchain:task:{task_id}",
|
||||
"temp_upload": "upload:temp:{upload_id}",
|
||||
"wallet_connection": "wallet:conn:{wallet_address}",
|
||||
"ton_price": "ton:price:usd",
|
||||
"system_status": "system:status:{service}",
|
||||
}
|
||||
|
||||
# Log current configuration (without secrets)
|
||||
def log_config():
|
||||
"""Log current configuration without sensitive data"""
|
||||
safe_config = {
|
||||
"project_name": settings.PROJECT_NAME,
|
||||
"project_version": settings.PROJECT_VERSION,
|
||||
"debug": settings.DEBUG,
|
||||
"sanic_port": settings.SANIC_PORT,
|
||||
"testnet": settings.TESTNET,
|
||||
"maintenance_mode": settings.MAINTENANCE_MODE,
|
||||
"metrics_enabled": settings.METRICS_ENABLED,
|
||||
"uploads_dir": str(settings.UPLOADS_DIR),
|
||||
"log_level": settings.LOG_LEVEL,
|
||||
}
|
||||
logger.info("Configuration loaded", **safe_config)
|
||||
|
||||
# Initialize logging configuration
|
||||
log_config()
|
||||
@@ -0,0 +1,257 @@
|
||||
"""Compatible configuration management with MariaDB and Redis support."""
|
||||
|
||||
import os
|
||||
from functools import lru_cache
|
||||
from typing import Optional, Dict, Any
|
||||
from pydantic import BaseSettings, Field, validator
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
"""Application settings with backward compatibility."""
|
||||
|
||||
# Application settings
|
||||
app_name: str = Field(default="My Uploader Bot", env="APP_NAME")
|
||||
debug: bool = Field(default=False, env="DEBUG")
|
||||
environment: str = Field(default="production", env="ENVIRONMENT")
|
||||
host: str = Field(default="0.0.0.0", env="HOST")
|
||||
port: int = Field(default=15100, env="PORT")
|
||||
|
||||
# Security settings
|
||||
secret_key: str = Field(env="SECRET_KEY", default="your-secret-key-change-this")
|
||||
jwt_secret_key: str = Field(env="JWT_SECRET_KEY", default="jwt-secret-change-this")
|
||||
jwt_algorithm: str = Field(default="HS256", env="JWT_ALGORITHM")
|
||||
jwt_expire_minutes: int = Field(default=30, env="JWT_EXPIRE_MINUTES")
|
||||
|
||||
# MariaDB/MySQL settings (preserving existing configuration)
|
||||
mysql_host: str = Field(default="maria_db", env="MYSQL_HOST")
|
||||
mysql_port: int = Field(default=3306, env="MYSQL_PORT")
|
||||
mysql_user: str = Field(default="myuploader", env="MYSQL_USER")
|
||||
mysql_password: str = Field(default="password", env="MYSQL_PASSWORD")
|
||||
mysql_database: str = Field(default="myuploader", env="MYSQL_DATABASE")
|
||||
mysql_root_password: str = Field(default="password", env="MYSQL_ROOT_PASSWORD")
|
||||
|
||||
# Database pool settings
|
||||
database_pool_size: int = Field(default=20, env="DATABASE_POOL_SIZE")
|
||||
database_max_overflow: int = Field(default=30, env="DATABASE_MAX_OVERFLOW")
|
||||
database_pool_timeout: int = Field(default=30, env="DATABASE_POOL_TIMEOUT")
|
||||
database_pool_recycle: int = Field(default=3600, env="DATABASE_POOL_RECYCLE")
|
||||
|
||||
# Optional new database URL (for future migration)
|
||||
database_url: Optional[str] = Field(default=None, env="DATABASE_URL")
|
||||
|
||||
# Redis settings (new addition)
|
||||
redis_enabled: bool = Field(default=True, env="REDIS_ENABLED")
|
||||
redis_host: str = Field(default="redis", env="REDIS_HOST")
|
||||
redis_port: int = Field(default=6379, env="REDIS_PORT")
|
||||
redis_password: Optional[str] = Field(default=None, env="REDIS_PASSWORD")
|
||||
redis_db: int = Field(default=0, env="REDIS_DB")
|
||||
redis_max_connections: int = Field(default=50, env="REDIS_MAX_CONNECTIONS")
|
||||
redis_socket_timeout: int = Field(default=30, env="REDIS_SOCKET_TIMEOUT")
|
||||
redis_socket_connect_timeout: int = Field(default=30, env="REDIS_SOCKET_CONNECT_TIMEOUT")
|
||||
|
||||
# Cache settings
|
||||
cache_enabled: bool = Field(default=True, env="CACHE_ENABLED")
|
||||
cache_default_ttl: int = Field(default=300, env="CACHE_DEFAULT_TTL") # 5 minutes
|
||||
cache_user_ttl: int = Field(default=600, env="CACHE_USER_TTL") # 10 minutes
|
||||
cache_content_ttl: int = Field(default=1800, env="CACHE_CONTENT_TTL") # 30 minutes
|
||||
|
||||
# Storage settings (preserving existing paths)
|
||||
storage_path: str = Field(default="/Storage/storedContent", env="STORAGE_PATH")
|
||||
logs_path: str = Field(default="/Storage/logs", env="LOGS_PATH")
|
||||
sql_storage_path: str = Field(default="/Storage/sqlStorage", env="SQL_STORAGE_PATH")
|
||||
|
||||
# File upload settings
|
||||
max_file_size: int = Field(default=100 * 1024 * 1024, env="MAX_FILE_SIZE") # 100MB
|
||||
allowed_extensions: str = Field(default=".jpg,.jpeg,.png,.gif,.pdf,.doc,.docx,.txt", env="ALLOWED_EXTENSIONS")
|
||||
|
||||
# Rate limiting
|
||||
rate_limit_enabled: bool = Field(default=True, env="RATE_LIMIT_ENABLED")
|
||||
rate_limit_requests: int = Field(default=100, env="RATE_LIMIT_REQUESTS")
|
||||
rate_limit_window: int = Field(default=3600, env="RATE_LIMIT_WINDOW") # 1 hour
|
||||
|
||||
# TON Blockchain settings (preserving existing)
|
||||
ton_network: str = Field(default="mainnet", env="TON_NETWORK")
|
||||
ton_api_key: Optional[str] = Field(default=None, env="TON_API_KEY")
|
||||
ton_wallet_address: Optional[str] = Field(default=None, env="TON_WALLET_ADDRESS")
|
||||
|
||||
# License settings
|
||||
license_check_enabled: bool = Field(default=True, env="LICENSE_CHECK_ENABLED")
|
||||
license_server_url: Optional[str] = Field(default=None, env="LICENSE_SERVER_URL")
|
||||
|
||||
# Indexer settings
|
||||
indexer_enabled: bool = Field(default=True, env="INDEXER_ENABLED")
|
||||
indexer_interval: int = Field(default=300, env="INDEXER_INTERVAL") # 5 minutes
|
||||
|
||||
# Convert process settings
|
||||
convert_enabled: bool = Field(default=True, env="CONVERT_ENABLED")
|
||||
convert_queue_size: int = Field(default=10, env="CONVERT_QUEUE_SIZE")
|
||||
|
||||
# Logging settings
|
||||
log_level: str = Field(default="INFO", env="LOG_LEVEL")
|
||||
log_format: str = Field(default="json", env="LOG_FORMAT")
|
||||
log_file_enabled: bool = Field(default=True, env="LOG_FILE_ENABLED")
|
||||
log_file_max_size: int = Field(default=10 * 1024 * 1024, env="LOG_FILE_MAX_SIZE") # 10MB
|
||||
log_file_backup_count: int = Field(default=5, env="LOG_FILE_BACKUP_COUNT")
|
||||
|
||||
# API settings
|
||||
api_title: str = Field(default="My Uploader Bot API", env="API_TITLE")
|
||||
api_version: str = Field(default="1.0.0", env="API_VERSION")
|
||||
api_description: str = Field(default="File upload and management API", env="API_DESCRIPTION")
|
||||
cors_enabled: bool = Field(default=True, env="CORS_ENABLED")
|
||||
cors_origins: str = Field(default="*", env="CORS_ORIGINS")
|
||||
|
||||
# Health check settings
|
||||
health_check_enabled: bool = Field(default=True, env="HEALTH_CHECK_ENABLED")
|
||||
health_check_interval: int = Field(default=60, env="HEALTH_CHECK_INTERVAL")
|
||||
|
||||
# Metrics settings
|
||||
metrics_enabled: bool = Field(default=True, env="METRICS_ENABLED")
|
||||
metrics_endpoint: str = Field(default="/metrics", env="METRICS_ENDPOINT")
|
||||
|
||||
@validator("allowed_extensions")
|
||||
def validate_extensions(cls, v):
|
||||
"""Validate and normalize file extensions."""
|
||||
if isinstance(v, str):
|
||||
return [ext.strip().lower() for ext in v.split(",") if ext.strip()]
|
||||
return v
|
||||
|
||||
@validator("cors_origins")
|
||||
def validate_cors_origins(cls, v):
|
||||
"""Validate and normalize CORS origins."""
|
||||
if isinstance(v, str) and v != "*":
|
||||
return [origin.strip() for origin in v.split(",") if origin.strip()]
|
||||
return v
|
||||
|
||||
@validator("log_level")
|
||||
def validate_log_level(cls, v):
|
||||
"""Validate log level."""
|
||||
valid_levels = ["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"]
|
||||
if v.upper() not in valid_levels:
|
||||
raise ValueError(f"Log level must be one of: {valid_levels}")
|
||||
return v.upper()
|
||||
|
||||
def get_database_url(self) -> str:
|
||||
"""Get complete database URL."""
|
||||
if self.database_url:
|
||||
return self.database_url
|
||||
return f"mysql+aiomysql://{self.mysql_user}:{self.mysql_password}@{self.mysql_host}:{self.mysql_port}/{self.mysql_database}"
|
||||
|
||||
def get_redis_url(self) -> str:
|
||||
"""Get complete Redis URL."""
|
||||
if self.redis_password:
|
||||
return f"redis://:{self.redis_password}@{self.redis_host}:{self.redis_port}/{self.redis_db}"
|
||||
return f"redis://{self.redis_host}:{self.redis_port}/{self.redis_db}"
|
||||
|
||||
def get_allowed_extensions_set(self) -> set:
|
||||
"""Get allowed extensions as a set."""
|
||||
if isinstance(self.allowed_extensions, list):
|
||||
return set(self.allowed_extensions)
|
||||
return set(ext.strip().lower() for ext in self.allowed_extensions.split(",") if ext.strip())
|
||||
|
||||
def get_cors_origins_list(self) -> list:
|
||||
"""Get CORS origins as a list."""
|
||||
if self.cors_origins == "*":
|
||||
return ["*"]
|
||||
if isinstance(self.cors_origins, list):
|
||||
return self.cors_origins
|
||||
return [origin.strip() for origin in self.cors_origins.split(",") if origin.strip()]
|
||||
|
||||
def is_development(self) -> bool:
|
||||
"""Check if running in development mode."""
|
||||
return self.environment.lower() in ["development", "dev", "local"]
|
||||
|
||||
def is_production(self) -> bool:
|
||||
"""Check if running in production mode."""
|
||||
return self.environment.lower() in ["production", "prod"]
|
||||
|
||||
def get_cache_config(self) -> Dict[str, Any]:
|
||||
"""Get cache configuration dictionary."""
|
||||
return {
|
||||
"enabled": self.cache_enabled and self.redis_enabled,
|
||||
"default_ttl": self.cache_default_ttl,
|
||||
"user_ttl": self.cache_user_ttl,
|
||||
"content_ttl": self.cache_content_ttl,
|
||||
"redis_url": self.get_redis_url(),
|
||||
"max_connections": self.redis_max_connections,
|
||||
}
|
||||
|
||||
def get_database_config(self) -> Dict[str, Any]:
|
||||
"""Get database configuration dictionary."""
|
||||
return {
|
||||
"url": self.get_database_url(),
|
||||
"pool_size": self.database_pool_size,
|
||||
"max_overflow": self.database_max_overflow,
|
||||
"pool_timeout": self.database_pool_timeout,
|
||||
"pool_recycle": self.database_pool_recycle,
|
||||
}
|
||||
|
||||
class Config:
|
||||
env_file = ".env"
|
||||
env_file_encoding = "utf-8"
|
||||
case_sensitive = False
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def get_settings() -> Settings:
|
||||
"""Get cached settings instance."""
|
||||
return Settings()
|
||||
|
||||
|
||||
# Backward compatibility functions
|
||||
def get_mysql_config() -> Dict[str, Any]:
|
||||
"""Get MySQL configuration for backward compatibility."""
|
||||
settings = get_settings()
|
||||
return {
|
||||
"host": settings.mysql_host,
|
||||
"port": settings.mysql_port,
|
||||
"user": settings.mysql_user,
|
||||
"password": settings.mysql_password,
|
||||
"database": settings.mysql_database,
|
||||
}
|
||||
|
||||
|
||||
def get_storage_config() -> Dict[str, str]:
|
||||
"""Get storage configuration for backward compatibility."""
|
||||
settings = get_settings()
|
||||
return {
|
||||
"storage_path": settings.storage_path,
|
||||
"logs_path": settings.logs_path,
|
||||
"sql_storage_path": settings.sql_storage_path,
|
||||
}
|
||||
|
||||
|
||||
def get_redis_config() -> Dict[str, Any]:
|
||||
"""Get Redis configuration."""
|
||||
settings = get_settings()
|
||||
return {
|
||||
"enabled": settings.redis_enabled,
|
||||
"host": settings.redis_host,
|
||||
"port": settings.redis_port,
|
||||
"password": settings.redis_password,
|
||||
"db": settings.redis_db,
|
||||
"max_connections": settings.redis_max_connections,
|
||||
"socket_timeout": settings.redis_socket_timeout,
|
||||
"socket_connect_timeout": settings.redis_socket_connect_timeout,
|
||||
}
|
||||
|
||||
|
||||
# Environment variables validation
|
||||
def validate_environment():
|
||||
"""Validate required environment variables."""
|
||||
settings = get_settings()
|
||||
|
||||
required_vars = [
|
||||
"SECRET_KEY",
|
||||
"JWT_SECRET_KEY",
|
||||
"MYSQL_PASSWORD",
|
||||
]
|
||||
|
||||
missing_vars = []
|
||||
for var in required_vars:
|
||||
if not os.getenv(var):
|
||||
missing_vars.append(var)
|
||||
|
||||
if missing_vars:
|
||||
raise ValueError(f"Missing required environment variables: {', '.join(missing_vars)}")
|
||||
|
||||
return True
|
||||
@@ -0,0 +1,262 @@
|
||||
"""
|
||||
Async SQLAlchemy configuration with connection pooling and Redis integration
|
||||
"""
|
||||
import asyncio
|
||||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import AsyncGenerator, Optional
|
||||
from datetime import timedelta
|
||||
|
||||
from sqlalchemy.ext.asyncio import (
|
||||
create_async_engine,
|
||||
AsyncSession,
|
||||
async_sessionmaker,
|
||||
AsyncEngine
|
||||
)
|
||||
from sqlalchemy.pool import NullPool, QueuePool
|
||||
from sqlalchemy.sql import text
|
||||
import redis.asyncio as redis
|
||||
from redis.asyncio.connection import ConnectionPool
|
||||
import structlog
|
||||
|
||||
from app.core.config import (
|
||||
DATABASE_URL,
|
||||
REDIS_URL,
|
||||
DATABASE_POOL_SIZE,
|
||||
DATABASE_MAX_OVERFLOW,
|
||||
REDIS_POOL_SIZE
|
||||
)
|
||||
|
||||
logger = structlog.get_logger(__name__)
|
||||
|
||||
|
||||
class DatabaseManager:
|
||||
"""Async database manager with connection pooling"""
|
||||
|
||||
def __init__(self):
|
||||
self._engine: Optional[AsyncEngine] = None
|
||||
self._session_factory: Optional[async_sessionmaker[AsyncSession]] = None
|
||||
self._redis_pool: Optional[ConnectionPool] = None
|
||||
self._redis: Optional[redis.Redis] = None
|
||||
self._initialized = False
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Initialize database connections and Redis"""
|
||||
if self._initialized:
|
||||
return
|
||||
|
||||
# Initialize async SQLAlchemy engine
|
||||
self._engine = create_async_engine(
|
||||
DATABASE_URL,
|
||||
poolclass=QueuePool,
|
||||
pool_size=DATABASE_POOL_SIZE,
|
||||
max_overflow=DATABASE_MAX_OVERFLOW,
|
||||
pool_pre_ping=True,
|
||||
pool_recycle=3600, # 1 hour
|
||||
echo=False, # Set to True for SQL debugging
|
||||
future=True,
|
||||
json_serializer=lambda obj: obj,
|
||||
json_deserializer=lambda obj: obj,
|
||||
)
|
||||
|
||||
# Create session factory
|
||||
self._session_factory = async_sessionmaker(
|
||||
self._engine,
|
||||
class_=AsyncSession,
|
||||
expire_on_commit=False,
|
||||
autoflush=False,
|
||||
autocommit=False
|
||||
)
|
||||
|
||||
# Initialize Redis connection pool
|
||||
self._redis_pool = ConnectionPool.from_url(
|
||||
REDIS_URL,
|
||||
max_connections=REDIS_POOL_SIZE,
|
||||
retry_on_timeout=True,
|
||||
health_check_interval=30
|
||||
)
|
||||
|
||||
self._redis = redis.Redis(
|
||||
connection_pool=self._redis_pool,
|
||||
decode_responses=True
|
||||
)
|
||||
|
||||
# Test connections
|
||||
await self._test_connections()
|
||||
self._initialized = True
|
||||
|
||||
logger.info("Database and Redis connections initialized")
|
||||
|
||||
async def _test_connections(self) -> None:
|
||||
"""Test database and Redis connections"""
|
||||
# Test database
|
||||
async with self._engine.begin() as conn:
|
||||
result = await conn.execute(text("SELECT 1"))
|
||||
assert result.scalar() == 1
|
||||
|
||||
# Test Redis
|
||||
await self._redis.ping()
|
||||
logger.info("Database and Redis connections tested successfully")
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close all connections gracefully"""
|
||||
if self._engine:
|
||||
await self._engine.dispose()
|
||||
|
||||
if self._redis_pool:
|
||||
await self._redis_pool.disconnect()
|
||||
|
||||
self._initialized = False
|
||||
logger.info("Database and Redis connections closed")
|
||||
|
||||
@asynccontextmanager
|
||||
async def get_session(self) -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Get async database session with automatic cleanup"""
|
||||
if not self._initialized:
|
||||
await self.initialize()
|
||||
|
||||
async with self._session_factory() as session:
|
||||
try:
|
||||
yield session
|
||||
except Exception as e:
|
||||
await session.rollback()
|
||||
logger.error("Database session error", error=str(e))
|
||||
raise
|
||||
finally:
|
||||
await session.close()
|
||||
|
||||
@asynccontextmanager
|
||||
async def get_transaction(self) -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Get async database session with automatic transaction management"""
|
||||
async with self.get_session() as session:
|
||||
async with session.begin():
|
||||
yield session
|
||||
|
||||
async def get_redis(self) -> redis.Redis:
|
||||
"""Get Redis client"""
|
||||
if not self._initialized:
|
||||
await self.initialize()
|
||||
return self._redis
|
||||
|
||||
@property
|
||||
def engine(self) -> AsyncEngine:
|
||||
"""Get SQLAlchemy engine"""
|
||||
if not self._engine:
|
||||
raise RuntimeError("Database not initialized")
|
||||
return self._engine
|
||||
|
||||
|
||||
class CacheManager:
|
||||
"""Redis-based cache manager with TTL and serialization"""
|
||||
|
||||
def __init__(self, redis_client: redis.Redis):
|
||||
self.redis = redis_client
|
||||
|
||||
async def get(self, key: str, default=None):
|
||||
"""Get value from cache"""
|
||||
try:
|
||||
value = await self.redis.get(key)
|
||||
return value if value is not None else default
|
||||
except Exception as e:
|
||||
logger.error("Cache get error", key=key, error=str(e))
|
||||
return default
|
||||
|
||||
async def set(
|
||||
self,
|
||||
key: str,
|
||||
value: str,
|
||||
ttl: Optional[int] = None,
|
||||
nx: bool = False
|
||||
) -> bool:
|
||||
"""Set value in cache with optional TTL"""
|
||||
try:
|
||||
return await self.redis.set(key, value, ex=ttl, nx=nx)
|
||||
except Exception as e:
|
||||
logger.error("Cache set error", key=key, error=str(e))
|
||||
return False
|
||||
|
||||
async def delete(self, key: str) -> bool:
|
||||
"""Delete key from cache"""
|
||||
try:
|
||||
return bool(await self.redis.delete(key))
|
||||
except Exception as e:
|
||||
logger.error("Cache delete error", key=key, error=str(e))
|
||||
return False
|
||||
|
||||
async def exists(self, key: str) -> bool:
|
||||
"""Check if key exists in cache"""
|
||||
try:
|
||||
return bool(await self.redis.exists(key))
|
||||
except Exception as e:
|
||||
logger.error("Cache exists error", key=key, error=str(e))
|
||||
return False
|
||||
|
||||
async def incr(self, key: str, amount: int = 1) -> int:
|
||||
"""Increment counter in cache"""
|
||||
try:
|
||||
return await self.redis.incr(key, amount)
|
||||
except Exception as e:
|
||||
logger.error("Cache incr error", key=key, error=str(e))
|
||||
return 0
|
||||
|
||||
async def expire(self, key: str, ttl: int) -> bool:
|
||||
"""Set TTL for existing key"""
|
||||
try:
|
||||
return await self.redis.expire(key, ttl)
|
||||
except Exception as e:
|
||||
logger.error("Cache expire error", key=key, error=str(e))
|
||||
return False
|
||||
|
||||
async def hget(self, name: str, key: str):
|
||||
"""Get hash field value"""
|
||||
try:
|
||||
return await self.redis.hget(name, key)
|
||||
except Exception as e:
|
||||
logger.error("Cache hget error", name=name, key=key, error=str(e))
|
||||
return None
|
||||
|
||||
async def hset(self, name: str, key: str, value: str) -> bool:
|
||||
"""Set hash field value"""
|
||||
try:
|
||||
return bool(await self.redis.hset(name, key, value))
|
||||
except Exception as e:
|
||||
logger.error("Cache hset error", name=name, key=key, error=str(e))
|
||||
return False
|
||||
|
||||
async def hdel(self, name: str, key: str) -> bool:
|
||||
"""Delete hash field"""
|
||||
try:
|
||||
return bool(await self.redis.hdel(name, key))
|
||||
except Exception as e:
|
||||
logger.error("Cache hdel error", name=name, key=key, error=str(e))
|
||||
return False
|
||||
|
||||
|
||||
# Global instances
|
||||
db_manager = DatabaseManager()
|
||||
cache_manager: Optional[CacheManager] = None
|
||||
|
||||
|
||||
async def get_db_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Dependency for getting database session"""
|
||||
async with db_manager.get_session() as session:
|
||||
yield session
|
||||
|
||||
|
||||
async def get_cache() -> CacheManager:
|
||||
"""Dependency for getting cache manager"""
|
||||
global cache_manager
|
||||
if not cache_manager:
|
||||
redis_client = await db_manager.get_redis()
|
||||
cache_manager = CacheManager(redis_client)
|
||||
return cache_manager
|
||||
|
||||
|
||||
async def init_database():
|
||||
"""Initialize database connections"""
|
||||
await db_manager.initialize()
|
||||
|
||||
|
||||
async def close_database():
|
||||
"""Close database connections"""
|
||||
await db_manager.close()
|
||||
@@ -0,0 +1,221 @@
|
||||
"""Compatible database configuration with MariaDB support."""
|
||||
|
||||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import AsyncGenerator, Optional
|
||||
|
||||
from sqlalchemy import MetaData
|
||||
from sqlalchemy.ext.asyncio import (
|
||||
AsyncEngine,
|
||||
AsyncSession,
|
||||
async_sessionmaker,
|
||||
create_async_engine
|
||||
)
|
||||
from sqlalchemy.pool import NullPool
|
||||
|
||||
from app.core.config import get_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Global variables for database engine and session
|
||||
_engine: Optional[AsyncEngine] = None
|
||||
_async_session: Optional[async_sessionmaker[AsyncSession]] = None
|
||||
|
||||
# Naming convention for consistent constraint names
|
||||
naming_convention = {
|
||||
"ix": "ix_%(column_0_label)s",
|
||||
"uq": "uq_%(table_name)s_%(column_0_name)s",
|
||||
"ck": "ck_%(table_name)s_%(constraint_name)s",
|
||||
"fk": "fk_%(table_name)s_%(column_0_name)s_%(referred_table_name)s",
|
||||
"pk": "pk_%(table_name)s"
|
||||
}
|
||||
|
||||
metadata = MetaData(naming_convention=naming_convention)
|
||||
|
||||
|
||||
def get_database_url() -> str:
|
||||
"""Get database URL from settings."""
|
||||
settings = get_settings()
|
||||
|
||||
# Support both new DATABASE_URL and legacy MariaDB settings
|
||||
if hasattr(settings, 'database_url') and settings.database_url:
|
||||
return settings.database_url
|
||||
|
||||
# Fallback to MariaDB configuration
|
||||
mysql_host = getattr(settings, 'mysql_host', 'maria_db')
|
||||
mysql_port = getattr(settings, 'mysql_port', 3306)
|
||||
mysql_user = getattr(settings, 'mysql_user', 'myuploader')
|
||||
mysql_password = getattr(settings, 'mysql_password', 'password')
|
||||
mysql_database = getattr(settings, 'mysql_database', 'myuploader')
|
||||
|
||||
return f"mysql+aiomysql://{mysql_user}:{mysql_password}@{mysql_host}:{mysql_port}/{mysql_database}"
|
||||
|
||||
|
||||
async def init_database() -> None:
|
||||
"""Initialize database connection."""
|
||||
global _engine, _async_session
|
||||
|
||||
if _engine is not None:
|
||||
logger.warning("Database already initialized")
|
||||
return
|
||||
|
||||
try:
|
||||
settings = get_settings()
|
||||
database_url = get_database_url()
|
||||
|
||||
logger.info(f"Connecting to database: {database_url.split('@')[1] if '@' in database_url else 'unknown'}")
|
||||
|
||||
# Create async engine with MariaDB/MySQL optimizations
|
||||
_engine = create_async_engine(
|
||||
database_url,
|
||||
echo=settings.debug if hasattr(settings, 'debug') else False,
|
||||
pool_size=getattr(settings, 'database_pool_size', 20),
|
||||
max_overflow=getattr(settings, 'database_max_overflow', 30),
|
||||
pool_timeout=getattr(settings, 'database_pool_timeout', 30),
|
||||
pool_recycle=getattr(settings, 'database_pool_recycle', 3600),
|
||||
pool_pre_ping=True, # Verify connections before use
|
||||
# MariaDB specific settings
|
||||
connect_args={
|
||||
"charset": "utf8mb4",
|
||||
"use_unicode": True,
|
||||
"autocommit": False,
|
||||
}
|
||||
)
|
||||
|
||||
# Create async session factory
|
||||
_async_session = async_sessionmaker(
|
||||
bind=_engine,
|
||||
class_=AsyncSession,
|
||||
expire_on_commit=False,
|
||||
autoflush=True,
|
||||
autocommit=False
|
||||
)
|
||||
|
||||
# Test the connection
|
||||
async with _engine.begin() as conn:
|
||||
await conn.execute("SELECT 1")
|
||||
|
||||
logger.info("Database connection established successfully")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize database: {e}")
|
||||
raise
|
||||
|
||||
|
||||
async def close_database() -> None:
|
||||
"""Close database connection."""
|
||||
global _engine, _async_session
|
||||
|
||||
if _engine is not None:
|
||||
logger.info("Closing database connection")
|
||||
await _engine.dispose()
|
||||
_engine = None
|
||||
_async_session = None
|
||||
logger.info("Database connection closed")
|
||||
|
||||
|
||||
def get_engine() -> AsyncEngine:
|
||||
"""Get database engine."""
|
||||
if _engine is None:
|
||||
raise RuntimeError("Database not initialized. Call init_database() first.")
|
||||
return _engine
|
||||
|
||||
|
||||
def get_session_factory() -> async_sessionmaker[AsyncSession]:
|
||||
"""Get session factory."""
|
||||
if _async_session is None:
|
||||
raise RuntimeError("Database not initialized. Call init_database() first.")
|
||||
return _async_session
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def get_async_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Get async database session with automatic cleanup."""
|
||||
if _async_session is None:
|
||||
raise RuntimeError("Database not initialized. Call init_database() first.")
|
||||
|
||||
async with _async_session() as session:
|
||||
try:
|
||||
yield session
|
||||
except Exception as e:
|
||||
logger.error(f"Database session error: {e}")
|
||||
await session.rollback()
|
||||
raise
|
||||
finally:
|
||||
await session.close()
|
||||
|
||||
|
||||
async def check_database_health() -> bool:
|
||||
"""Check database connection health."""
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
await session.execute("SELECT 1")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Database health check failed: {e}")
|
||||
return False
|
||||
|
||||
|
||||
async def get_database_info() -> dict:
|
||||
"""Get database information."""
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
# Get database version
|
||||
result = await session.execute("SELECT VERSION() as version")
|
||||
version_row = result.fetchone()
|
||||
version = version_row[0] if version_row else "Unknown"
|
||||
|
||||
# Get connection count (MariaDB specific)
|
||||
try:
|
||||
result = await session.execute("SHOW STATUS LIKE 'Threads_connected'")
|
||||
conn_row = result.fetchone()
|
||||
connections = int(conn_row[1]) if conn_row else 0
|
||||
except:
|
||||
connections = 0
|
||||
|
||||
# Get database size
|
||||
try:
|
||||
result = await session.execute("""
|
||||
SELECT
|
||||
ROUND(SUM(data_length + index_length) / 1024 / 1024, 2) as size_mb
|
||||
FROM information_schema.tables
|
||||
WHERE table_schema = DATABASE()
|
||||
""")
|
||||
size_row = result.fetchone()
|
||||
size_mb = float(size_row[0]) if size_row and size_row[0] else 0
|
||||
except:
|
||||
size_mb = 0
|
||||
|
||||
return {
|
||||
"version": version,
|
||||
"connections": connections,
|
||||
"size_mb": size_mb,
|
||||
"engine_pool_size": _engine.pool.size() if _engine else 0,
|
||||
"engine_checked_out": _engine.pool.checkedout() if _engine else 0,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get database info: {e}")
|
||||
return {"error": str(e)}
|
||||
|
||||
|
||||
# Database session dependency for dependency injection
|
||||
async def get_db_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Database session dependency for API routes."""
|
||||
async with get_async_session() as session:
|
||||
yield session
|
||||
|
||||
|
||||
# Backward compatibility functions
|
||||
async def get_db() -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Legacy function name for backward compatibility."""
|
||||
async with get_async_session() as session:
|
||||
yield session
|
||||
|
||||
|
||||
# Transaction context manager
|
||||
@asynccontextmanager
|
||||
async def transaction():
|
||||
"""Transaction context manager."""
|
||||
async with get_async_session() as session:
|
||||
async with session.begin():
|
||||
yield session
|
||||
@@ -0,0 +1,363 @@
|
||||
"""
|
||||
Structured logging configuration with monitoring and observability
|
||||
"""
|
||||
import asyncio
|
||||
import logging
|
||||
import sys
|
||||
import time
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional, Union
|
||||
from contextvars import ContextVar
|
||||
import json
|
||||
|
||||
import structlog
|
||||
from structlog.stdlib import LoggerFactory
|
||||
from structlog.typing import EventDict, Processor
|
||||
import structlog.dev
|
||||
|
||||
from app.core.config import settings, LOG_DIR, LOG_LEVEL
|
||||
|
||||
# Context variables for request tracking
|
||||
request_id_var: ContextVar[Optional[str]] = ContextVar('request_id', default=None)
|
||||
user_id_var: ContextVar[Optional[int]] = ContextVar('user_id', default=None)
|
||||
operation_var: ContextVar[Optional[str]] = ContextVar('operation', default=None)
|
||||
|
||||
|
||||
class RequestContextProcessor:
|
||||
"""Add request context to log records"""
|
||||
|
||||
def __call__(self, logger, method_name, event_dict: EventDict) -> EventDict:
|
||||
"""Add context variables to event dict"""
|
||||
if request_id := request_id_var.get(None):
|
||||
event_dict['request_id'] = request_id
|
||||
|
||||
if user_id := user_id_var.get(None):
|
||||
event_dict['user_id'] = user_id
|
||||
|
||||
if operation := operation_var.get(None):
|
||||
event_dict['operation'] = operation
|
||||
|
||||
return event_dict
|
||||
|
||||
|
||||
class TimestampProcessor:
|
||||
"""Add consistent timestamp to log records"""
|
||||
|
||||
def __call__(self, logger, method_name, event_dict: EventDict) -> EventDict:
|
||||
"""Add timestamp to event dict"""
|
||||
event_dict['timestamp'] = datetime.utcnow().isoformat() + 'Z'
|
||||
return event_dict
|
||||
|
||||
|
||||
class SecurityProcessor:
|
||||
"""Filter sensitive data from logs"""
|
||||
|
||||
SENSITIVE_KEYS = {
|
||||
'password', 'token', 'key', 'secret', 'auth', 'credential',
|
||||
'private_key', 'seed', 'mnemonic', 'api_key', 'authorization'
|
||||
}
|
||||
|
||||
def __call__(self, logger, method_name, event_dict: EventDict) -> EventDict:
|
||||
"""Remove or mask sensitive data"""
|
||||
return self._filter_dict(event_dict)
|
||||
|
||||
def _filter_dict(self, data: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Recursively filter sensitive data"""
|
||||
if not isinstance(data, dict):
|
||||
return data
|
||||
|
||||
filtered = {}
|
||||
for key, value in data.items():
|
||||
if any(sensitive in key.lower() for sensitive in self.SENSITIVE_KEYS):
|
||||
filtered[key] = '***REDACTED***'
|
||||
elif isinstance(value, dict):
|
||||
filtered[key] = self._filter_dict(value)
|
||||
elif isinstance(value, list):
|
||||
filtered[key] = [
|
||||
self._filter_dict(item) if isinstance(item, dict) else item
|
||||
for item in value
|
||||
]
|
||||
else:
|
||||
filtered[key] = value
|
||||
|
||||
return filtered
|
||||
|
||||
|
||||
class PerformanceProcessor:
|
||||
"""Add performance metrics to log records"""
|
||||
|
||||
def __call__(self, logger, method_name, event_dict: EventDict) -> EventDict:
|
||||
"""Add performance data to event dict"""
|
||||
# Add memory usage if available
|
||||
try:
|
||||
import psutil
|
||||
process = psutil.Process()
|
||||
event_dict['memory_mb'] = round(process.memory_info().rss / 1024 / 1024, 2)
|
||||
event_dict['cpu_percent'] = process.cpu_percent()
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
return event_dict
|
||||
|
||||
|
||||
class MetricsCollector:
|
||||
"""Collect metrics from log events"""
|
||||
|
||||
def __init__(self):
|
||||
self.counters: Dict[str, int] = {}
|
||||
self.timers: Dict[str, float] = {}
|
||||
self.errors: Dict[str, int] = {}
|
||||
|
||||
def increment_counter(self, metric: str, value: int = 1):
|
||||
"""Increment counter metric"""
|
||||
self.counters[metric] = self.counters.get(metric, 0) + value
|
||||
|
||||
def record_timer(self, metric: str, duration: float):
|
||||
"""Record timer metric"""
|
||||
self.timers[metric] = duration
|
||||
|
||||
def record_error(self, error_type: str):
|
||||
"""Record error metric"""
|
||||
self.errors[error_type] = self.errors.get(error_type, 0) + 1
|
||||
|
||||
def get_metrics(self) -> Dict[str, Any]:
|
||||
"""Get all collected metrics"""
|
||||
return {
|
||||
'counters': self.counters,
|
||||
'timers': self.timers,
|
||||
'errors': self.errors
|
||||
}
|
||||
|
||||
|
||||
# Global metrics collector
|
||||
metrics_collector = MetricsCollector()
|
||||
|
||||
|
||||
class DatabaseLogHandler(logging.Handler):
|
||||
"""Log handler that stores critical logs in database"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.setLevel(logging.ERROR)
|
||||
self._queue = asyncio.Queue(maxsize=1000)
|
||||
self._task = None
|
||||
|
||||
def emit(self, record: logging.LogRecord):
|
||||
"""Add log record to queue"""
|
||||
try:
|
||||
log_entry = {
|
||||
'timestamp': datetime.utcnow(),
|
||||
'level': record.levelname,
|
||||
'logger': record.name,
|
||||
'message': record.getMessage(),
|
||||
'module': record.module,
|
||||
'function': record.funcName,
|
||||
'line': record.lineno,
|
||||
'request_id': getattr(record, 'request_id', None),
|
||||
'user_id': getattr(record, 'user_id', None),
|
||||
'extra': getattr(record, '__dict__', {})
|
||||
}
|
||||
|
||||
if not self._queue.full():
|
||||
self._queue.put_nowait(log_entry)
|
||||
|
||||
except Exception:
|
||||
# Don't let logging errors break the application
|
||||
pass
|
||||
|
||||
async def process_logs(self):
|
||||
"""Process logs from queue and store in database"""
|
||||
from app.core.database import get_db_session
|
||||
|
||||
while True:
|
||||
try:
|
||||
log_entry = await self._queue.get()
|
||||
|
||||
# Store in database (implement based on your log model)
|
||||
# async with get_db_session() as session:
|
||||
# log_record = LogRecord(**log_entry)
|
||||
# session.add(log_record)
|
||||
# await session.commit()
|
||||
|
||||
except Exception as e:
|
||||
# Log to stderr to avoid infinite recursion
|
||||
print(f"Database log handler error: {e}", file=sys.stderr)
|
||||
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
|
||||
def configure_logging():
|
||||
"""Configure structured logging"""
|
||||
|
||||
# Configure standard library logging
|
||||
logging.basicConfig(
|
||||
format="%(message)s",
|
||||
stream=sys.stdout,
|
||||
level=getattr(logging, LOG_LEVEL.upper())
|
||||
)
|
||||
|
||||
# Silence noisy loggers
|
||||
logging.getLogger("sqlalchemy.engine").setLevel(logging.WARNING)
|
||||
logging.getLogger("aioredis").setLevel(logging.WARNING)
|
||||
logging.getLogger("aiogram").setLevel(logging.WARNING)
|
||||
|
||||
# Configure processors based on environment
|
||||
processors: list[Processor] = [
|
||||
structlog.contextvars.merge_contextvars,
|
||||
RequestContextProcessor(),
|
||||
TimestampProcessor(),
|
||||
SecurityProcessor(),
|
||||
structlog.processors.add_log_level,
|
||||
structlog.processors.StackInfoRenderer(),
|
||||
]
|
||||
|
||||
if settings.DEBUG:
|
||||
processors.extend([
|
||||
PerformanceProcessor(),
|
||||
structlog.dev.ConsoleRenderer(colors=True)
|
||||
])
|
||||
else:
|
||||
processors.append(structlog.processors.JSONRenderer())
|
||||
|
||||
# Configure structlog
|
||||
structlog.configure(
|
||||
processors=processors,
|
||||
wrapper_class=structlog.make_filtering_bound_logger(
|
||||
getattr(logging, LOG_LEVEL.upper())
|
||||
),
|
||||
logger_factory=LoggerFactory(),
|
||||
cache_logger_on_first_use=True,
|
||||
)
|
||||
|
||||
# Add file handler for persistent logging
|
||||
if not settings.DEBUG:
|
||||
log_file = LOG_DIR / f"app_{datetime.now().strftime('%Y%m%d')}.log"
|
||||
file_handler = logging.FileHandler(log_file, encoding='utf-8')
|
||||
file_handler.setFormatter(
|
||||
logging.Formatter(
|
||||
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
|
||||
)
|
||||
)
|
||||
logging.getLogger().addHandler(file_handler)
|
||||
|
||||
|
||||
class LoggerMixin:
|
||||
"""Mixin to add structured logging to classes"""
|
||||
|
||||
@property
|
||||
def logger(self):
|
||||
"""Get logger for this class"""
|
||||
return structlog.get_logger(self.__class__.__name__)
|
||||
|
||||
|
||||
class AsyncContextLogger:
|
||||
"""Context manager for async operations with automatic logging"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
operation: str,
|
||||
logger: Optional[structlog.BoundLogger] = None,
|
||||
log_args: bool = True,
|
||||
log_result: bool = True
|
||||
):
|
||||
self.operation = operation
|
||||
self.logger = logger or structlog.get_logger()
|
||||
self.log_args = log_args
|
||||
self.log_result = log_result
|
||||
self.start_time = None
|
||||
|
||||
async def __aenter__(self):
|
||||
"""Enter async context"""
|
||||
self.start_time = time.time()
|
||||
operation_var.set(self.operation)
|
||||
|
||||
self.logger.info(
|
||||
"Operation started",
|
||||
operation=self.operation,
|
||||
)
|
||||
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
"""Exit async context with performance logging"""
|
||||
duration = time.time() - self.start_time
|
||||
|
||||
if exc_type:
|
||||
self.logger.error(
|
||||
"Operation failed",
|
||||
operation=self.operation,
|
||||
duration_ms=round(duration * 1000, 2),
|
||||
error_type=exc_type.__name__,
|
||||
error_message=str(exc_val)
|
||||
)
|
||||
metrics_collector.record_error(f"{self.operation}_error")
|
||||
else:
|
||||
self.logger.info(
|
||||
"Operation completed",
|
||||
operation=self.operation,
|
||||
duration_ms=round(duration * 1000, 2)
|
||||
)
|
||||
|
||||
metrics_collector.record_timer(f"{self.operation}_duration", duration)
|
||||
operation_var.set(None)
|
||||
|
||||
|
||||
def get_logger(name: str = None) -> structlog.BoundLogger:
|
||||
"""Get configured structured logger"""
|
||||
return structlog.get_logger(name)
|
||||
|
||||
|
||||
# Compatibility wrapper for old logging
|
||||
def make_log(
|
||||
component: Optional[str],
|
||||
message: str,
|
||||
level: str = 'info',
|
||||
**kwargs
|
||||
):
|
||||
"""Legacy logging function for backward compatibility"""
|
||||
logger = get_logger(component or 'Legacy')
|
||||
log_func = getattr(logger, level.lower(), logger.info)
|
||||
log_func(message, **kwargs)
|
||||
|
||||
|
||||
# Performance monitoring decorator
|
||||
def log_performance(operation: str = None):
|
||||
"""Decorator to log function performance"""
|
||||
def decorator(func):
|
||||
async def async_wrapper(*args, **kwargs):
|
||||
op_name = operation or f"{func.__module__}.{func.__name__}"
|
||||
async with AsyncContextLogger(op_name):
|
||||
return await func(*args, **kwargs)
|
||||
|
||||
def sync_wrapper(*args, **kwargs):
|
||||
op_name = operation or f"{func.__module__}.{func.__name__}"
|
||||
start_time = time.time()
|
||||
logger = get_logger(func.__module__)
|
||||
|
||||
try:
|
||||
logger.info("Function started", function=op_name)
|
||||
result = func(*args, **kwargs)
|
||||
duration = time.time() - start_time
|
||||
logger.info(
|
||||
"Function completed",
|
||||
function=op_name,
|
||||
duration_ms=round(duration * 1000, 2)
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
duration = time.time() - start_time
|
||||
logger.error(
|
||||
"Function failed",
|
||||
function=op_name,
|
||||
duration_ms=round(duration * 1000, 2),
|
||||
error=str(e)
|
||||
)
|
||||
raise
|
||||
|
||||
return async_wrapper if asyncio.iscoroutinefunction(func) else sync_wrapper
|
||||
return decorator
|
||||
|
||||
|
||||
# Initialize logging
|
||||
configure_logging()
|
||||
@@ -0,0 +1,566 @@
|
||||
"""Prometheus metrics collection for my-uploader-bot."""
|
||||
|
||||
import logging
|
||||
import time
|
||||
from datetime import datetime
|
||||
from functools import wraps
|
||||
from typing import Dict, Any, Optional, Callable
|
||||
|
||||
from prometheus_client import Counter, Histogram, Gauge, Info, generate_latest, CONTENT_TYPE_LATEST
|
||||
from sanic import Request, Response
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# Application info
|
||||
APP_INFO = Info('myuploader_app_info', 'Application information')
|
||||
APP_INFO.info({
|
||||
'version': '2.0.0',
|
||||
'name': 'my-uploader-bot',
|
||||
'python_version': '3.11+'
|
||||
})
|
||||
|
||||
# HTTP request metrics
|
||||
HTTP_REQUESTS_TOTAL = Counter(
|
||||
'http_requests_total',
|
||||
'Total HTTP requests',
|
||||
['method', 'endpoint', 'status_code']
|
||||
)
|
||||
|
||||
HTTP_REQUEST_DURATION = Histogram(
|
||||
'http_request_duration_seconds',
|
||||
'HTTP request duration in seconds',
|
||||
['method', 'endpoint']
|
||||
)
|
||||
|
||||
HTTP_REQUEST_SIZE = Histogram(
|
||||
'http_request_size_bytes',
|
||||
'HTTP request size in bytes',
|
||||
['method', 'endpoint']
|
||||
)
|
||||
|
||||
HTTP_RESPONSE_SIZE = Histogram(
|
||||
'http_response_size_bytes',
|
||||
'HTTP response size in bytes',
|
||||
['method', 'endpoint']
|
||||
)
|
||||
|
||||
# Authentication metrics
|
||||
AUTH_LOGIN_ATTEMPTS_TOTAL = Counter(
|
||||
'auth_login_attempts_total',
|
||||
'Total login attempts',
|
||||
['status']
|
||||
)
|
||||
|
||||
AUTH_LOGIN_FAILURES_TOTAL = Counter(
|
||||
'auth_login_failures_total',
|
||||
'Total login failures',
|
||||
['reason']
|
||||
)
|
||||
|
||||
AUTH_API_KEY_USAGE_TOTAL = Counter(
|
||||
'auth_api_key_usage_total',
|
||||
'Total API key usage',
|
||||
['key_id', 'status']
|
||||
)
|
||||
|
||||
# File upload metrics
|
||||
UPLOAD_REQUESTS_TOTAL = Counter(
|
||||
'upload_requests_total',
|
||||
'Total upload requests',
|
||||
['status', 'file_type']
|
||||
)
|
||||
|
||||
UPLOAD_SIZE_BYTES = Histogram(
|
||||
'upload_size_bytes',
|
||||
'File upload size in bytes',
|
||||
['file_type']
|
||||
)
|
||||
|
||||
UPLOAD_DURATION_SECONDS = Histogram(
|
||||
'upload_duration_seconds',
|
||||
'File upload duration in seconds',
|
||||
['file_type']
|
||||
)
|
||||
|
||||
UPLOAD_QUEUE_SIZE = Gauge(
|
||||
'upload_queue_size',
|
||||
'Number of files in upload queue'
|
||||
)
|
||||
|
||||
UPLOAD_FAILURES_TOTAL = Counter(
|
||||
'upload_failures_total',
|
||||
'Total upload failures',
|
||||
['reason', 'file_type']
|
||||
)
|
||||
|
||||
# File processing metrics
|
||||
PROCESSING_QUEUE_SIZE = Gauge(
|
||||
'processing_queue_size',
|
||||
'Number of files in processing queue'
|
||||
)
|
||||
|
||||
PROCESSING_DURATION_SECONDS = Histogram(
|
||||
'processing_duration_seconds',
|
||||
'File processing duration in seconds',
|
||||
['file_type', 'operation']
|
||||
)
|
||||
|
||||
PROCESSING_FAILURES_TOTAL = Counter(
|
||||
'processing_failures_total',
|
||||
'Total processing failures',
|
||||
['file_type', 'operation']
|
||||
)
|
||||
|
||||
# Database metrics
|
||||
DB_CONNECTIONS_ACTIVE = Gauge(
|
||||
'db_connections_active',
|
||||
'Number of active database connections'
|
||||
)
|
||||
|
||||
DB_CONNECTIONS_IDLE = Gauge(
|
||||
'db_connections_idle',
|
||||
'Number of idle database connections'
|
||||
)
|
||||
|
||||
DB_QUERY_DURATION_SECONDS = Histogram(
|
||||
'db_query_duration_seconds',
|
||||
'Database query duration in seconds',
|
||||
['operation']
|
||||
)
|
||||
|
||||
DB_TRANSACTIONS_TOTAL = Counter(
|
||||
'db_transactions_total',
|
||||
'Total database transactions',
|
||||
['status']
|
||||
)
|
||||
|
||||
# Cache metrics
|
||||
CACHE_OPERATIONS_TOTAL = Counter(
|
||||
'cache_operations_total',
|
||||
'Total cache operations',
|
||||
['operation', 'status']
|
||||
)
|
||||
|
||||
CACHE_HIT_RATIO = Gauge(
|
||||
'cache_hit_ratio',
|
||||
'Cache hit ratio'
|
||||
)
|
||||
|
||||
CACHE_KEYS_TOTAL = Gauge(
|
||||
'cache_keys_total',
|
||||
'Total number of cache keys'
|
||||
)
|
||||
|
||||
CACHE_MEMORY_USAGE_BYTES = Gauge(
|
||||
'cache_memory_usage_bytes',
|
||||
'Cache memory usage in bytes'
|
||||
)
|
||||
|
||||
# Storage metrics
|
||||
STORAGE_OPERATIONS_TOTAL = Counter(
|
||||
'storage_operations_total',
|
||||
'Total storage operations',
|
||||
['operation', 'backend', 'status']
|
||||
)
|
||||
|
||||
STORAGE_AVAILABLE_BYTES = Gauge(
|
||||
'storage_available_bytes',
|
||||
'Available storage space in bytes',
|
||||
['backend']
|
||||
)
|
||||
|
||||
STORAGE_TOTAL_BYTES = Gauge(
|
||||
'storage_total_bytes',
|
||||
'Total storage space in bytes',
|
||||
['backend']
|
||||
)
|
||||
|
||||
STORAGE_FILES_TOTAL = Gauge(
|
||||
'storage_files_total',
|
||||
'Total number of stored files',
|
||||
['backend']
|
||||
)
|
||||
|
||||
# Blockchain metrics
|
||||
BLOCKCHAIN_TRANSACTIONS_TOTAL = Counter(
|
||||
'blockchain_transactions_total',
|
||||
'Total blockchain transactions',
|
||||
['status', 'network']
|
||||
)
|
||||
|
||||
BLOCKCHAIN_TRANSACTION_FEES = Histogram(
|
||||
'blockchain_transaction_fees',
|
||||
'Blockchain transaction fees',
|
||||
['network']
|
||||
)
|
||||
|
||||
BLOCKCHAIN_PENDING_TRANSACTIONS = Gauge(
|
||||
'blockchain_pending_transactions',
|
||||
'Number of pending blockchain transactions'
|
||||
)
|
||||
|
||||
BLOCKCHAIN_WALLET_BALANCES = Gauge(
|
||||
'blockchain_wallet_balances',
|
||||
'Wallet balances',
|
||||
['wallet_id', 'currency']
|
||||
)
|
||||
|
||||
TON_SERVICE_UP = Gauge(
|
||||
'ton_service_up',
|
||||
'TON service availability (1 = up, 0 = down)'
|
||||
)
|
||||
|
||||
# Security metrics
|
||||
RATE_LIMIT_HITS_TOTAL = Counter(
|
||||
'rate_limit_hits_total',
|
||||
'Total rate limit hits',
|
||||
['endpoint', 'user_id']
|
||||
)
|
||||
|
||||
SECURITY_EVENTS_TOTAL = Counter(
|
||||
'security_events_total',
|
||||
'Total security events',
|
||||
['event_type', 'severity']
|
||||
)
|
||||
|
||||
SECURITY_SUSPICIOUS_EVENTS = Gauge(
|
||||
'security_suspicious_events',
|
||||
'Number of suspicious security events in the last hour'
|
||||
)
|
||||
|
||||
FAILED_LOGIN_ATTEMPTS = Counter(
|
||||
'failed_login_attempts_total',
|
||||
'Total failed login attempts',
|
||||
['ip_address', 'reason']
|
||||
)
|
||||
|
||||
# System metrics
|
||||
SYSTEM_UPTIME_SECONDS = Gauge(
|
||||
'system_uptime_seconds',
|
||||
'System uptime in seconds'
|
||||
)
|
||||
|
||||
BACKGROUND_TASKS_ACTIVE = Gauge(
|
||||
'background_tasks_active',
|
||||
'Number of active background tasks',
|
||||
['service']
|
||||
)
|
||||
|
||||
BACKGROUND_TASKS_COMPLETED = Counter(
|
||||
'background_tasks_completed_total',
|
||||
'Total completed background tasks',
|
||||
['service', 'status']
|
||||
)
|
||||
|
||||
# Error metrics
|
||||
ERROR_RATE = Gauge(
|
||||
'error_rate',
|
||||
'Application error rate'
|
||||
)
|
||||
|
||||
EXCEPTIONS_TOTAL = Counter(
|
||||
'exceptions_total',
|
||||
'Total exceptions',
|
||||
['exception_type', 'handler']
|
||||
)
|
||||
|
||||
|
||||
class MetricsCollector:
|
||||
"""Centralized metrics collection and management."""
|
||||
|
||||
def __init__(self):
|
||||
self.start_time = time.time()
|
||||
self._cache_stats = {
|
||||
'hits': 0,
|
||||
'misses': 0,
|
||||
'operations': 0
|
||||
}
|
||||
|
||||
def record_http_request(
|
||||
self,
|
||||
method: str,
|
||||
endpoint: str,
|
||||
status_code: int,
|
||||
duration: float,
|
||||
request_size: int = 0,
|
||||
response_size: int = 0
|
||||
):
|
||||
"""Record HTTP request metrics."""
|
||||
HTTP_REQUESTS_TOTAL.labels(
|
||||
method=method,
|
||||
endpoint=endpoint,
|
||||
status_code=status_code
|
||||
).inc()
|
||||
|
||||
HTTP_REQUEST_DURATION.labels(
|
||||
method=method,
|
||||
endpoint=endpoint
|
||||
).observe(duration)
|
||||
|
||||
if request_size > 0:
|
||||
HTTP_REQUEST_SIZE.labels(
|
||||
method=method,
|
||||
endpoint=endpoint
|
||||
).observe(request_size)
|
||||
|
||||
if response_size > 0:
|
||||
HTTP_RESPONSE_SIZE.labels(
|
||||
method=method,
|
||||
endpoint=endpoint
|
||||
).observe(response_size)
|
||||
|
||||
def record_auth_event(self, event_type: str, status: str, **labels):
|
||||
"""Record authentication events."""
|
||||
if event_type == 'login':
|
||||
AUTH_LOGIN_ATTEMPTS_TOTAL.labels(status=status).inc()
|
||||
if status == 'failed':
|
||||
reason = labels.get('reason', 'unknown')
|
||||
AUTH_LOGIN_FAILURES_TOTAL.labels(reason=reason).inc()
|
||||
elif event_type == 'api_key':
|
||||
key_id = labels.get('key_id', 'unknown')
|
||||
AUTH_API_KEY_USAGE_TOTAL.labels(key_id=key_id, status=status).inc()
|
||||
|
||||
def record_upload_event(
|
||||
self,
|
||||
status: str,
|
||||
file_type: str,
|
||||
file_size: int = 0,
|
||||
duration: float = 0,
|
||||
**kwargs
|
||||
):
|
||||
"""Record file upload events."""
|
||||
UPLOAD_REQUESTS_TOTAL.labels(status=status, file_type=file_type).inc()
|
||||
|
||||
if file_size > 0:
|
||||
UPLOAD_SIZE_BYTES.labels(file_type=file_type).observe(file_size)
|
||||
|
||||
if duration > 0:
|
||||
UPLOAD_DURATION_SECONDS.labels(file_type=file_type).observe(duration)
|
||||
|
||||
if status == 'failed':
|
||||
reason = kwargs.get('reason', 'unknown')
|
||||
UPLOAD_FAILURES_TOTAL.labels(reason=reason, file_type=file_type).inc()
|
||||
|
||||
def record_processing_event(
|
||||
self,
|
||||
file_type: str,
|
||||
operation: str,
|
||||
duration: float = 0,
|
||||
status: str = 'success'
|
||||
):
|
||||
"""Record file processing events."""
|
||||
if duration > 0:
|
||||
PROCESSING_DURATION_SECONDS.labels(
|
||||
file_type=file_type,
|
||||
operation=operation
|
||||
).observe(duration)
|
||||
|
||||
if status == 'failed':
|
||||
PROCESSING_FAILURES_TOTAL.labels(
|
||||
file_type=file_type,
|
||||
operation=operation
|
||||
).inc()
|
||||
|
||||
def record_db_event(self, operation: str, duration: float = 0, status: str = 'success'):
|
||||
"""Record database events."""
|
||||
if duration > 0:
|
||||
DB_QUERY_DURATION_SECONDS.labels(operation=operation).observe(duration)
|
||||
|
||||
DB_TRANSACTIONS_TOTAL.labels(status=status).inc()
|
||||
|
||||
def record_cache_event(self, operation: str, status: str):
|
||||
"""Record cache events."""
|
||||
CACHE_OPERATIONS_TOTAL.labels(operation=operation, status=status).inc()
|
||||
|
||||
# Update cache stats
|
||||
self._cache_stats['operations'] += 1
|
||||
if status == 'hit':
|
||||
self._cache_stats['hits'] += 1
|
||||
elif status == 'miss':
|
||||
self._cache_stats['misses'] += 1
|
||||
|
||||
# Update hit ratio
|
||||
if self._cache_stats['operations'] > 0:
|
||||
hit_ratio = self._cache_stats['hits'] / self._cache_stats['operations']
|
||||
CACHE_HIT_RATIO.set(hit_ratio)
|
||||
|
||||
def record_blockchain_event(
|
||||
self,
|
||||
event_type: str,
|
||||
status: str,
|
||||
network: str = 'mainnet',
|
||||
**kwargs
|
||||
):
|
||||
"""Record blockchain events."""
|
||||
if event_type == 'transaction':
|
||||
BLOCKCHAIN_TRANSACTIONS_TOTAL.labels(status=status, network=network).inc()
|
||||
|
||||
if 'fee' in kwargs:
|
||||
BLOCKCHAIN_TRANSACTION_FEES.labels(network=network).observe(kwargs['fee'])
|
||||
|
||||
def record_security_event(self, event_type: str, severity: str = 'info', **kwargs):
|
||||
"""Record security events."""
|
||||
SECURITY_EVENTS_TOTAL.labels(event_type=event_type, severity=severity).inc()
|
||||
|
||||
if event_type == 'rate_limit':
|
||||
endpoint = kwargs.get('endpoint', 'unknown')
|
||||
user_id = kwargs.get('user_id', 'anonymous')
|
||||
RATE_LIMIT_HITS_TOTAL.labels(endpoint=endpoint, user_id=user_id).inc()
|
||||
|
||||
elif event_type == 'failed_login':
|
||||
ip_address = kwargs.get('ip_address', 'unknown')
|
||||
reason = kwargs.get('reason', 'unknown')
|
||||
FAILED_LOGIN_ATTEMPTS.labels(ip_address=ip_address, reason=reason).inc()
|
||||
|
||||
def update_system_metrics(self):
|
||||
"""Update system-level metrics."""
|
||||
uptime = time.time() - self.start_time
|
||||
SYSTEM_UPTIME_SECONDS.set(uptime)
|
||||
|
||||
def update_gauge_metrics(self, metrics_data: Dict[str, Any]):
|
||||
"""Update gauge metrics from external data."""
|
||||
# Database metrics
|
||||
if 'db_connections' in metrics_data:
|
||||
db_conn = metrics_data['db_connections']
|
||||
DB_CONNECTIONS_ACTIVE.set(db_conn.get('active', 0))
|
||||
DB_CONNECTIONS_IDLE.set(db_conn.get('idle', 0))
|
||||
|
||||
# Cache metrics
|
||||
if 'cache' in metrics_data:
|
||||
cache_data = metrics_data['cache']
|
||||
CACHE_KEYS_TOTAL.set(cache_data.get('keys', 0))
|
||||
CACHE_MEMORY_USAGE_BYTES.set(cache_data.get('memory_usage', 0))
|
||||
|
||||
# Storage metrics
|
||||
if 'storage' in metrics_data:
|
||||
storage_data = metrics_data['storage']
|
||||
for backend, data in storage_data.items():
|
||||
STORAGE_AVAILABLE_BYTES.labels(backend=backend).set(data.get('available', 0))
|
||||
STORAGE_TOTAL_BYTES.labels(backend=backend).set(data.get('total', 0))
|
||||
STORAGE_FILES_TOTAL.labels(backend=backend).set(data.get('files', 0))
|
||||
|
||||
# Queue metrics
|
||||
if 'queues' in metrics_data:
|
||||
queues = metrics_data['queues']
|
||||
UPLOAD_QUEUE_SIZE.set(queues.get('upload', 0))
|
||||
PROCESSING_QUEUE_SIZE.set(queues.get('processing', 0))
|
||||
|
||||
# Blockchain metrics
|
||||
if 'blockchain' in metrics_data:
|
||||
blockchain_data = metrics_data['blockchain']
|
||||
BLOCKCHAIN_PENDING_TRANSACTIONS.set(blockchain_data.get('pending_transactions', 0))
|
||||
TON_SERVICE_UP.set(1 if blockchain_data.get('ton_service_up') else 0)
|
||||
|
||||
# Wallet balances
|
||||
for wallet_id, balance_data in blockchain_data.get('wallet_balances', {}).items():
|
||||
for currency, balance in balance_data.items():
|
||||
BLOCKCHAIN_WALLET_BALANCES.labels(
|
||||
wallet_id=wallet_id,
|
||||
currency=currency
|
||||
).set(balance)
|
||||
|
||||
# Background tasks
|
||||
if 'background_tasks' in metrics_data:
|
||||
tasks_data = metrics_data['background_tasks']
|
||||
for service, count in tasks_data.items():
|
||||
BACKGROUND_TASKS_ACTIVE.labels(service=service).set(count)
|
||||
|
||||
|
||||
# Global metrics collector instance
|
||||
metrics_collector = MetricsCollector()
|
||||
|
||||
|
||||
def metrics_middleware(request: Request, response: Response):
|
||||
"""Middleware to collect HTTP metrics."""
|
||||
start_time = time.time()
|
||||
|
||||
# After request processing
|
||||
duration = time.time() - start_time
|
||||
|
||||
# Get endpoint info
|
||||
endpoint = request.path
|
||||
method = request.method
|
||||
status_code = response.status
|
||||
|
||||
# Get request/response sizes
|
||||
request_size = len(request.body) if request.body else 0
|
||||
response_size = len(response.body) if hasattr(response, 'body') and response.body else 0
|
||||
|
||||
# Record metrics
|
||||
metrics_collector.record_http_request(
|
||||
method=method,
|
||||
endpoint=endpoint,
|
||||
status_code=status_code,
|
||||
duration=duration,
|
||||
request_size=request_size,
|
||||
response_size=response_size
|
||||
)
|
||||
|
||||
|
||||
def track_function_calls(func_name: str, labels: Optional[Dict[str, str]] = None):
|
||||
"""Decorator to track function call metrics."""
|
||||
def decorator(func: Callable) -> Callable:
|
||||
@wraps(func)
|
||||
async def async_wrapper(*args, **kwargs):
|
||||
start_time = time.time()
|
||||
status = 'success'
|
||||
|
||||
try:
|
||||
result = await func(*args, **kwargs)
|
||||
return result
|
||||
except Exception as e:
|
||||
status = 'error'
|
||||
EXCEPTIONS_TOTAL.labels(
|
||||
exception_type=type(e).__name__,
|
||||
handler=func_name
|
||||
).inc()
|
||||
raise
|
||||
finally:
|
||||
duration = time.time() - start_time
|
||||
# Record custom metrics based on function type
|
||||
if func_name.startswith('db_'):
|
||||
metrics_collector.record_db_event(func_name, duration, status)
|
||||
elif func_name.startswith('cache_'):
|
||||
metrics_collector.record_cache_event(func_name, status)
|
||||
|
||||
@wraps(func)
|
||||
def sync_wrapper(*args, **kwargs):
|
||||
start_time = time.time()
|
||||
status = 'success'
|
||||
|
||||
try:
|
||||
result = func(*args, **kwargs)
|
||||
return result
|
||||
except Exception as e:
|
||||
status = 'error'
|
||||
EXCEPTIONS_TOTAL.labels(
|
||||
exception_type=type(e).__name__,
|
||||
handler=func_name
|
||||
).inc()
|
||||
raise
|
||||
finally:
|
||||
duration = time.time() - start_time
|
||||
# Record custom metrics based on function type
|
||||
if func_name.startswith('db_'):
|
||||
metrics_collector.record_db_event(func_name, duration, status)
|
||||
elif func_name.startswith('cache_'):
|
||||
metrics_collector.record_cache_event(func_name, status)
|
||||
|
||||
return async_wrapper if asyncio.iscoroutinefunction(func) else sync_wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
async def get_metrics():
|
||||
"""Get Prometheus metrics."""
|
||||
# Update system metrics before generating output
|
||||
metrics_collector.update_system_metrics()
|
||||
|
||||
# Generate metrics in Prometheus format
|
||||
return generate_latest()
|
||||
|
||||
|
||||
def get_metrics_content_type():
|
||||
"""Get the content type for metrics."""
|
||||
return CONTENT_TYPE_LATEST
|
||||
+276
-2
@@ -1,3 +1,277 @@
|
||||
from sqlalchemy.ext.declarative import declarative_base
|
||||
"""
|
||||
Base model classes with async SQLAlchemy support
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional, Type, TypeVar, Union
|
||||
|
||||
AlchemyBase = declarative_base()
|
||||
from sqlalchemy import Column, DateTime, String, Boolean, Integer, Text, JSON
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.ext.declarative import declarative_base
|
||||
from sqlalchemy.future import select
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from pydantic import BaseModel
|
||||
import structlog
|
||||
|
||||
logger = structlog.get_logger(__name__)
|
||||
|
||||
# Create declarative base
|
||||
Base = declarative_base()
|
||||
|
||||
# Type variable for model classes
|
||||
ModelType = TypeVar("ModelType", bound="BaseModel")
|
||||
|
||||
|
||||
class TimestampMixin:
|
||||
"""Mixin for automatic timestamp fields"""
|
||||
|
||||
created_at = Column(
|
||||
DateTime,
|
||||
nullable=False,
|
||||
default=datetime.utcnow,
|
||||
comment="Record creation timestamp"
|
||||
)
|
||||
updated_at = Column(
|
||||
DateTime,
|
||||
nullable=False,
|
||||
default=datetime.utcnow,
|
||||
onupdate=datetime.utcnow,
|
||||
comment="Record last update timestamp"
|
||||
)
|
||||
|
||||
|
||||
class UUIDMixin:
|
||||
"""Mixin for UUID primary key"""
|
||||
|
||||
id = Column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=uuid.uuid4,
|
||||
comment="Unique identifier"
|
||||
)
|
||||
|
||||
|
||||
class SoftDeleteMixin:
|
||||
"""Mixin for soft delete functionality"""
|
||||
|
||||
deleted_at = Column(
|
||||
DateTime,
|
||||
nullable=True,
|
||||
comment="Soft delete timestamp"
|
||||
)
|
||||
|
||||
@property
|
||||
def is_deleted(self) -> bool:
|
||||
"""Check if record is soft deleted"""
|
||||
return self.deleted_at is not None
|
||||
|
||||
def soft_delete(self):
|
||||
"""Mark record as soft deleted"""
|
||||
self.deleted_at = datetime.utcnow()
|
||||
|
||||
def restore(self):
|
||||
"""Restore soft deleted record"""
|
||||
self.deleted_at = None
|
||||
|
||||
|
||||
class MetadataMixin:
|
||||
"""Mixin for flexible metadata storage"""
|
||||
|
||||
metadata = Column(
|
||||
JSON,
|
||||
nullable=False,
|
||||
default=dict,
|
||||
comment="Flexible metadata storage"
|
||||
)
|
||||
|
||||
def set_meta(self, key: str, value: Any) -> None:
|
||||
"""Set metadata value"""
|
||||
if self.metadata is None:
|
||||
self.metadata = {}
|
||||
self.metadata[key] = value
|
||||
|
||||
def get_meta(self, key: str, default: Any = None) -> Any:
|
||||
"""Get metadata value"""
|
||||
if self.metadata is None:
|
||||
return default
|
||||
return self.metadata.get(key, default)
|
||||
|
||||
def update_meta(self, updates: Dict[str, Any]) -> None:
|
||||
"""Update multiple metadata values"""
|
||||
if self.metadata is None:
|
||||
self.metadata = {}
|
||||
self.metadata.update(updates)
|
||||
|
||||
|
||||
class StatusMixin:
|
||||
"""Mixin for status tracking"""
|
||||
|
||||
status = Column(
|
||||
String(64),
|
||||
nullable=False,
|
||||
default="active",
|
||||
index=True,
|
||||
comment="Record status"
|
||||
)
|
||||
|
||||
def set_status(self, status: str, reason: Optional[str] = None):
|
||||
"""Set status with optional reason"""
|
||||
self.status = status
|
||||
if reason:
|
||||
self.set_meta("status_reason", reason)
|
||||
self.set_meta("status_changed_at", datetime.utcnow().isoformat())
|
||||
|
||||
|
||||
class BaseModelMixin:
|
||||
"""Base mixin with common functionality"""
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Convert model to dictionary"""
|
||||
result = {}
|
||||
for column in self.__table__.columns:
|
||||
value = getattr(self, column.name)
|
||||
if isinstance(value, datetime):
|
||||
value = value.isoformat()
|
||||
elif hasattr(value, '__dict__'):
|
||||
value = str(value)
|
||||
result[column.name] = value
|
||||
return result
|
||||
|
||||
def update_from_dict(self, data: Dict[str, Any]) -> None:
|
||||
"""Update model from dictionary"""
|
||||
for key, value in data.items():
|
||||
if hasattr(self, key):
|
||||
setattr(self, key, value)
|
||||
|
||||
@classmethod
|
||||
async def get_by_id(
|
||||
cls: Type[ModelType],
|
||||
session: AsyncSession,
|
||||
id_value: Union[int, str, uuid.UUID]
|
||||
) -> Optional[ModelType]:
|
||||
"""Get record by ID"""
|
||||
try:
|
||||
stmt = select(cls).where(cls.id == id_value)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
except Exception as e:
|
||||
logger.error("Error getting record by ID", model=cls.__name__, id=id_value, error=str(e))
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
async def get_all(
|
||||
cls: Type[ModelType],
|
||||
session: AsyncSession,
|
||||
limit: Optional[int] = None,
|
||||
offset: Optional[int] = None
|
||||
) -> list[ModelType]:
|
||||
"""Get all records with optional pagination"""
|
||||
try:
|
||||
stmt = select(cls)
|
||||
if offset:
|
||||
stmt = stmt.offset(offset)
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalars().all()
|
||||
except Exception as e:
|
||||
logger.error("Error getting all records", model=cls.__name__, error=str(e))
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
async def count(cls: Type[ModelType], session: AsyncSession) -> int:
|
||||
"""Get total count of records"""
|
||||
try:
|
||||
from sqlalchemy import func
|
||||
stmt = select(func.count(cls.id))
|
||||
result = await session.execute(stmt)
|
||||
return result.scalar() or 0
|
||||
except Exception as e:
|
||||
logger.error("Error counting records", model=cls.__name__, error=str(e))
|
||||
return 0
|
||||
|
||||
async def save(self, session: AsyncSession) -> None:
|
||||
"""Save model to database"""
|
||||
try:
|
||||
session.add(self)
|
||||
await session.commit()
|
||||
await session.refresh(self)
|
||||
except Exception as e:
|
||||
await session.rollback()
|
||||
logger.error("Error saving model", model=self.__class__.__name__, error=str(e))
|
||||
raise
|
||||
|
||||
async def delete(self, session: AsyncSession) -> None:
|
||||
"""Delete model from database"""
|
||||
try:
|
||||
await session.delete(self)
|
||||
await session.commit()
|
||||
except Exception as e:
|
||||
await session.rollback()
|
||||
logger.error("Error deleting model", model=self.__class__.__name__, error=str(e))
|
||||
raise
|
||||
|
||||
|
||||
class AuditMixin:
|
||||
"""Mixin for audit trail"""
|
||||
|
||||
created_by = Column(
|
||||
UUID(as_uuid=True),
|
||||
nullable=True,
|
||||
comment="User who created the record"
|
||||
)
|
||||
updated_by = Column(
|
||||
UUID(as_uuid=True),
|
||||
nullable=True,
|
||||
comment="User who last updated the record"
|
||||
)
|
||||
|
||||
def set_audit_info(self, user_id: Optional[uuid.UUID] = None):
|
||||
"""Set audit information"""
|
||||
if user_id:
|
||||
if not hasattr(self, 'created_at') or not self.created_at:
|
||||
self.created_by = user_id
|
||||
self.updated_by = user_id
|
||||
|
||||
|
||||
class CacheableMixin:
|
||||
"""Mixin for cacheable models"""
|
||||
|
||||
@property
|
||||
def cache_key(self) -> str:
|
||||
"""Generate cache key for this model"""
|
||||
return f"{self.__class__.__name__.lower()}:{self.id}"
|
||||
|
||||
@property
|
||||
def cache_ttl(self) -> int:
|
||||
"""Default cache TTL in seconds"""
|
||||
return 3600 # 1 hour
|
||||
|
||||
def get_cache_data(self) -> Dict[str, Any]:
|
||||
"""Get data for caching"""
|
||||
return self.to_dict()
|
||||
|
||||
|
||||
# Combined base model class
|
||||
class BaseModel(
|
||||
Base,
|
||||
BaseModelMixin,
|
||||
TimestampMixin,
|
||||
UUIDMixin,
|
||||
SoftDeleteMixin,
|
||||
MetadataMixin,
|
||||
StatusMixin,
|
||||
AuditMixin,
|
||||
CacheableMixin
|
||||
):
|
||||
"""Base model with all mixins"""
|
||||
__abstract__ = True
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""String representation of model"""
|
||||
return f"<{self.__class__.__name__}(id={self.id})>"
|
||||
|
||||
|
||||
# Compatibility with old model base
|
||||
AlchemyBase = Base
|
||||
@@ -0,0 +1,88 @@
|
||||
"""Compatible SQLAlchemy base models for MariaDB."""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Optional, Dict, Any
|
||||
from sqlalchemy import Column, Integer, DateTime, text
|
||||
from sqlalchemy.ext.declarative import declarative_base, declared_attr
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
# Create base class
|
||||
Base = declarative_base()
|
||||
|
||||
|
||||
class TimestampMixin:
|
||||
"""Mixin for adding timestamp fields."""
|
||||
|
||||
@declared_attr
|
||||
def created_at(cls):
|
||||
return Column(
|
||||
DateTime,
|
||||
nullable=False,
|
||||
default=datetime.utcnow,
|
||||
server_default=text('CURRENT_TIMESTAMP')
|
||||
)
|
||||
|
||||
@declared_attr
|
||||
def updated_at(cls):
|
||||
return Column(
|
||||
DateTime,
|
||||
nullable=False,
|
||||
default=datetime.utcnow,
|
||||
onupdate=datetime.utcnow,
|
||||
server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP')
|
||||
)
|
||||
|
||||
|
||||
class BaseModel(Base, TimestampMixin):
|
||||
"""Base model with common fields for all entities."""
|
||||
|
||||
__abstract__ = True
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
def to_dict(self, exclude: Optional[set] = None) -> Dict[str, Any]:
|
||||
"""Convert model instance to dictionary."""
|
||||
exclude = exclude or set()
|
||||
result = {}
|
||||
|
||||
for column in self.__table__.columns:
|
||||
if column.name not in exclude:
|
||||
value = getattr(self, column.name)
|
||||
# Handle datetime serialization
|
||||
if isinstance(value, datetime):
|
||||
result[column.name] = value.isoformat()
|
||||
else:
|
||||
result[column.name] = value
|
||||
|
||||
return result
|
||||
|
||||
def update_from_dict(self, data: Dict[str, Any], exclude: Optional[set] = None) -> None:
|
||||
"""Update model instance from dictionary."""
|
||||
exclude = exclude or {"id", "created_at", "updated_at"}
|
||||
|
||||
for key, value in data.items():
|
||||
if key not in exclude and hasattr(self, key):
|
||||
setattr(self, key, value)
|
||||
|
||||
@classmethod
|
||||
def get_table_name(cls) -> str:
|
||||
"""Get table name."""
|
||||
return cls.__tablename__
|
||||
|
||||
@classmethod
|
||||
def get_columns(cls) -> list:
|
||||
"""Get list of column names."""
|
||||
return [column.name for column in cls.__table__.columns]
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""String representation of model."""
|
||||
return f"<{self.__class__.__name__}(id={getattr(self, 'id', None)})>"
|
||||
|
||||
|
||||
# Legacy session factory for backward compatibility
|
||||
SessionLocal = sessionmaker()
|
||||
|
||||
|
||||
def get_session():
|
||||
"""Get database session (legacy function for compatibility)."""
|
||||
return SessionLocal()
|
||||
@@ -0,0 +1,445 @@
|
||||
"""
|
||||
Blockchain-related models for TON network integration.
|
||||
Handles transaction records, wallet management, and smart contract interactions.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import Dict, List, Optional, Any
|
||||
from uuid import UUID
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import Column, String, Integer, DateTime, Boolean, Text, JSON, ForeignKey, Index
|
||||
from sqlalchemy.orm import relationship, validates
|
||||
from sqlalchemy.dialects.postgresql import UUID as PostgreSQLUUID
|
||||
|
||||
from app.core.models.base import Base, TimestampMixin, UUIDMixin
|
||||
|
||||
class BlockchainTransaction(Base, UUIDMixin, TimestampMixin):
|
||||
"""Model for storing blockchain transaction records."""
|
||||
__tablename__ = "blockchain_transactions"
|
||||
|
||||
# User relationship
|
||||
user_id = Column(PostgreSQLUUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
|
||||
user = relationship("User", back_populates="blockchain_transactions")
|
||||
|
||||
# Transaction details
|
||||
transaction_hash = Column(String(64), unique=True, nullable=False, index=True)
|
||||
transaction_type = Column(String(20), nullable=False) # transfer, mint, burn, stake, etc.
|
||||
status = Column(String(20), nullable=False, default="pending") # pending, confirmed, failed
|
||||
|
||||
# Amount and fees
|
||||
amount = Column(sa.BIGINT, nullable=False, default=0) # Amount in nanotons
|
||||
network_fee = Column(sa.BIGINT, nullable=False, default=0) # Network fee in nanotons
|
||||
|
||||
# Addresses
|
||||
sender_address = Column(String(48), nullable=True, index=True)
|
||||
recipient_address = Column(String(48), nullable=True, index=True)
|
||||
|
||||
# Message and metadata
|
||||
message = Column(Text, nullable=True)
|
||||
metadata = Column(JSON, nullable=True)
|
||||
|
||||
# Blockchain specific fields
|
||||
block_hash = Column(String(64), nullable=True)
|
||||
logical_time = Column(sa.BIGINT, nullable=True) # TON logical time
|
||||
confirmations = Column(Integer, nullable=False, default=0)
|
||||
|
||||
# Timing
|
||||
confirmed_at = Column(DateTime, nullable=True)
|
||||
failed_at = Column(DateTime, nullable=True)
|
||||
|
||||
# Smart contract interaction
|
||||
contract_address = Column(String(48), nullable=True)
|
||||
contract_method = Column(String(100), nullable=True)
|
||||
contract_data = Column(JSON, nullable=True)
|
||||
|
||||
# Internal tracking
|
||||
retry_count = Column(Integer, nullable=False, default=0)
|
||||
last_retry_at = Column(DateTime, nullable=True)
|
||||
error_message = Column(Text, nullable=True)
|
||||
|
||||
# Indexes for performance
|
||||
__table_args__ = (
|
||||
Index("idx_blockchain_tx_user_status", "user_id", "status"),
|
||||
Index("idx_blockchain_tx_hash", "transaction_hash"),
|
||||
Index("idx_blockchain_tx_addresses", "sender_address", "recipient_address"),
|
||||
Index("idx_blockchain_tx_created", "created_at"),
|
||||
Index("idx_blockchain_tx_type_status", "transaction_type", "status"),
|
||||
)
|
||||
|
||||
@validates('transaction_type')
|
||||
def validate_transaction_type(self, key, transaction_type):
|
||||
"""Validate transaction type."""
|
||||
allowed_types = {
|
||||
'transfer', 'mint', 'burn', 'stake', 'unstake',
|
||||
'contract_call', 'deploy', 'withdraw', 'deposit'
|
||||
}
|
||||
if transaction_type not in allowed_types:
|
||||
raise ValueError(f"Invalid transaction type: {transaction_type}")
|
||||
return transaction_type
|
||||
|
||||
@validates('status')
|
||||
def validate_status(self, key, status):
|
||||
"""Validate transaction status."""
|
||||
allowed_statuses = {'pending', 'confirmed', 'failed', 'cancelled'}
|
||||
if status not in allowed_statuses:
|
||||
raise ValueError(f"Invalid status: {status}")
|
||||
return status
|
||||
|
||||
@property
|
||||
def amount_tons(self) -> Decimal:
|
||||
"""Convert nanotons to TON."""
|
||||
return Decimal(self.amount) / Decimal("1000000000")
|
||||
|
||||
@property
|
||||
def fee_tons(self) -> Decimal:
|
||||
"""Convert fee nanotons to TON."""
|
||||
return Decimal(self.network_fee) / Decimal("1000000000")
|
||||
|
||||
@property
|
||||
def is_incoming(self) -> bool:
|
||||
"""Check if transaction is incoming to user's wallet."""
|
||||
return self.transaction_type in {'transfer', 'mint', 'deposit'} and self.recipient_address
|
||||
|
||||
@property
|
||||
def is_outgoing(self) -> bool:
|
||||
"""Check if transaction is outgoing from user's wallet."""
|
||||
return self.transaction_type in {'transfer', 'burn', 'withdraw'} and self.sender_address
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Convert transaction to dictionary."""
|
||||
return {
|
||||
"id": str(self.id),
|
||||
"hash": self.transaction_hash,
|
||||
"type": self.transaction_type,
|
||||
"status": self.status,
|
||||
"amount": self.amount,
|
||||
"amount_tons": str(self.amount_tons),
|
||||
"fee": self.network_fee,
|
||||
"fee_tons": str(self.fee_tons),
|
||||
"sender": self.sender_address,
|
||||
"recipient": self.recipient_address,
|
||||
"message": self.message,
|
||||
"block_hash": self.block_hash,
|
||||
"confirmations": self.confirmations,
|
||||
"created_at": self.created_at.isoformat() if self.created_at else None,
|
||||
"confirmed_at": self.confirmed_at.isoformat() if self.confirmed_at else None,
|
||||
"is_incoming": self.is_incoming,
|
||||
"is_outgoing": self.is_outgoing
|
||||
}
|
||||
|
||||
class SmartContract(Base, UUIDMixin, TimestampMixin):
|
||||
"""Model for smart contract management."""
|
||||
__tablename__ = "smart_contracts"
|
||||
|
||||
# Contract details
|
||||
address = Column(String(48), unique=True, nullable=False, index=True)
|
||||
name = Column(String(100), nullable=False)
|
||||
description = Column(Text, nullable=True)
|
||||
contract_type = Column(String(50), nullable=False) # nft, token, defi, etc.
|
||||
|
||||
# Contract metadata
|
||||
abi = Column(JSON, nullable=True) # Contract ABI if available
|
||||
source_code = Column(Text, nullable=True)
|
||||
compiler_version = Column(String(20), nullable=True)
|
||||
|
||||
# Deployment info
|
||||
deployer_address = Column(String(48), nullable=True)
|
||||
deployment_tx_hash = Column(String(64), nullable=True)
|
||||
deployment_block = Column(sa.BIGINT, nullable=True)
|
||||
|
||||
# Status and verification
|
||||
is_verified = Column(Boolean, nullable=False, default=False)
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
verification_date = Column(DateTime, nullable=True)
|
||||
|
||||
# Usage statistics
|
||||
interaction_count = Column(Integer, nullable=False, default=0)
|
||||
last_interaction_at = Column(DateTime, nullable=True)
|
||||
|
||||
# Relationships
|
||||
transactions = relationship(
|
||||
"BlockchainTransaction",
|
||||
foreign_keys="BlockchainTransaction.contract_address",
|
||||
primaryjoin="SmartContract.address == BlockchainTransaction.contract_address",
|
||||
back_populates=None
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
Index("idx_smart_contract_address", "address"),
|
||||
Index("idx_smart_contract_type", "contract_type"),
|
||||
Index("idx_smart_contract_active", "is_active"),
|
||||
)
|
||||
|
||||
@validates('contract_type')
|
||||
def validate_contract_type(self, key, contract_type):
|
||||
"""Validate contract type."""
|
||||
allowed_types = {
|
||||
'nft', 'token', 'defi', 'game', 'dao', 'bridge',
|
||||
'oracle', 'multisig', 'custom'
|
||||
}
|
||||
if contract_type not in allowed_types:
|
||||
raise ValueError(f"Invalid contract type: {contract_type}")
|
||||
return contract_type
|
||||
|
||||
class TokenBalance(Base, UUIDMixin, TimestampMixin):
|
||||
"""Model for tracking user token balances."""
|
||||
__tablename__ = "token_balances"
|
||||
|
||||
# User relationship
|
||||
user_id = Column(PostgreSQLUUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
|
||||
user = relationship("User", back_populates="token_balances")
|
||||
|
||||
# Token details
|
||||
token_address = Column(String(48), nullable=False, index=True)
|
||||
token_name = Column(String(100), nullable=True)
|
||||
token_symbol = Column(String(10), nullable=True)
|
||||
token_decimals = Column(Integer, nullable=False, default=9)
|
||||
|
||||
# Balance information
|
||||
balance = Column(sa.BIGINT, nullable=False, default=0) # Raw balance
|
||||
locked_balance = Column(sa.BIGINT, nullable=False, default=0) # Locked in contracts
|
||||
|
||||
# Metadata
|
||||
last_update_block = Column(sa.BIGINT, nullable=True)
|
||||
last_update_tx = Column(String(64), nullable=True)
|
||||
|
||||
# Unique constraint
|
||||
__table_args__ = (
|
||||
sa.UniqueConstraint("user_id", "token_address", name="uq_user_token"),
|
||||
Index("idx_token_balance_user", "user_id"),
|
||||
Index("idx_token_balance_token", "token_address"),
|
||||
Index("idx_token_balance_updated", "updated_at"),
|
||||
)
|
||||
|
||||
@property
|
||||
def available_balance(self) -> int:
|
||||
"""Get available (unlocked) balance."""
|
||||
return max(0, self.balance - self.locked_balance)
|
||||
|
||||
@property
|
||||
def formatted_balance(self) -> Decimal:
|
||||
"""Get balance formatted with decimals."""
|
||||
return Decimal(self.balance) / Decimal(10 ** self.token_decimals)
|
||||
|
||||
@property
|
||||
def formatted_available_balance(self) -> Decimal:
|
||||
"""Get available balance formatted with decimals."""
|
||||
return Decimal(self.available_balance) / Decimal(10 ** self.token_decimals)
|
||||
|
||||
class StakingPosition(Base, UUIDMixin, TimestampMixin):
|
||||
"""Model for staking positions."""
|
||||
__tablename__ = "staking_positions"
|
||||
|
||||
# User relationship
|
||||
user_id = Column(PostgreSQLUUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
|
||||
user = relationship("User", back_populates="staking_positions")
|
||||
|
||||
# Staking details
|
||||
validator_address = Column(String(48), nullable=False, index=True)
|
||||
pool_address = Column(String(48), nullable=True)
|
||||
|
||||
# Amount and timing
|
||||
staked_amount = Column(sa.BIGINT, nullable=False) # Amount in nanotons
|
||||
stake_tx_hash = Column(String(64), nullable=False)
|
||||
stake_block = Column(sa.BIGINT, nullable=True)
|
||||
|
||||
# Status
|
||||
status = Column(String(20), nullable=False, default="active") # active, unstaking, withdrawn
|
||||
unstake_tx_hash = Column(String(64), nullable=True)
|
||||
unstake_requested_at = Column(DateTime, nullable=True)
|
||||
withdrawn_at = Column(DateTime, nullable=True)
|
||||
|
||||
# Rewards
|
||||
rewards_earned = Column(sa.BIGINT, nullable=False, default=0)
|
||||
last_reward_claim = Column(DateTime, nullable=True)
|
||||
last_reward_tx = Column(String(64), nullable=True)
|
||||
|
||||
# Lock period
|
||||
lock_period_days = Column(Integer, nullable=False, default=0)
|
||||
unlock_date = Column(DateTime, nullable=True)
|
||||
|
||||
__table_args__ = (
|
||||
Index("idx_staking_user_status", "user_id", "status"),
|
||||
Index("idx_staking_validator", "validator_address"),
|
||||
Index("idx_staking_unlock", "unlock_date"),
|
||||
)
|
||||
|
||||
@validates('status')
|
||||
def validate_status(self, key, status):
|
||||
"""Validate staking status."""
|
||||
allowed_statuses = {'active', 'unstaking', 'withdrawn', 'slashed'}
|
||||
if status not in allowed_statuses:
|
||||
raise ValueError(f"Invalid staking status: {status}")
|
||||
return status
|
||||
|
||||
@property
|
||||
def staked_tons(self) -> Decimal:
|
||||
"""Get staked amount in TON."""
|
||||
return Decimal(self.staked_amount) / Decimal("1000000000")
|
||||
|
||||
@property
|
||||
def rewards_tons(self) -> Decimal:
|
||||
"""Get rewards amount in TON."""
|
||||
return Decimal(self.rewards_earned) / Decimal("1000000000")
|
||||
|
||||
@property
|
||||
def is_locked(self) -> bool:
|
||||
"""Check if staking position is still locked."""
|
||||
if not self.unlock_date:
|
||||
return False
|
||||
return datetime.utcnow() < self.unlock_date
|
||||
|
||||
class NFTCollection(Base, UUIDMixin, TimestampMixin):
|
||||
"""Model for NFT collections."""
|
||||
__tablename__ = "nft_collections"
|
||||
|
||||
# Collection details
|
||||
contract_address = Column(String(48), unique=True, nullable=False, index=True)
|
||||
name = Column(String(100), nullable=False)
|
||||
description = Column(Text, nullable=True)
|
||||
symbol = Column(String(10), nullable=True)
|
||||
|
||||
# Creator and metadata
|
||||
creator_address = Column(String(48), nullable=False)
|
||||
metadata_uri = Column(String(500), nullable=True)
|
||||
base_uri = Column(String(500), nullable=True)
|
||||
|
||||
# Collection stats
|
||||
total_supply = Column(Integer, nullable=False, default=0)
|
||||
max_supply = Column(Integer, nullable=True)
|
||||
floor_price = Column(sa.BIGINT, nullable=True) # In nanotons
|
||||
|
||||
# Status
|
||||
is_verified = Column(Boolean, nullable=False, default=False)
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
# Relationships
|
||||
nfts = relationship("NFTToken", back_populates="collection")
|
||||
|
||||
__table_args__ = (
|
||||
Index("idx_nft_collection_address", "contract_address"),
|
||||
Index("idx_nft_collection_creator", "creator_address"),
|
||||
Index("idx_nft_collection_verified", "is_verified"),
|
||||
)
|
||||
|
||||
class NFTToken(Base, UUIDMixin, TimestampMixin):
|
||||
"""Model for individual NFT tokens."""
|
||||
__tablename__ = "nft_tokens"
|
||||
|
||||
# Token identification
|
||||
collection_id = Column(PostgreSQLUUID(as_uuid=True), ForeignKey("nft_collections.id"), nullable=False)
|
||||
collection = relationship("NFTCollection", back_populates="nfts")
|
||||
|
||||
token_id = Column(String(100), nullable=False) # Token ID within collection
|
||||
token_address = Column(String(48), unique=True, nullable=False, index=True)
|
||||
|
||||
# Ownership
|
||||
owner_address = Column(String(48), nullable=False, index=True)
|
||||
|
||||
# Metadata
|
||||
name = Column(String(200), nullable=True)
|
||||
description = Column(Text, nullable=True)
|
||||
image_uri = Column(String(500), nullable=True)
|
||||
metadata_uri = Column(String(500), nullable=True)
|
||||
attributes = Column(JSON, nullable=True)
|
||||
|
||||
# Trading
|
||||
last_sale_price = Column(sa.BIGINT, nullable=True) # In nanotons
|
||||
last_sale_tx = Column(String(64), nullable=True)
|
||||
last_sale_date = Column(DateTime, nullable=True)
|
||||
|
||||
# Status
|
||||
is_burned = Column(Boolean, nullable=False, default=False)
|
||||
burned_at = Column(DateTime, nullable=True)
|
||||
|
||||
__table_args__ = (
|
||||
sa.UniqueConstraint("collection_id", "token_id", name="uq_collection_token"),
|
||||
Index("idx_nft_token_address", "token_address"),
|
||||
Index("idx_nft_token_owner", "owner_address"),
|
||||
Index("idx_nft_token_collection", "collection_id"),
|
||||
)
|
||||
|
||||
@property
|
||||
def last_sale_tons(self) -> Optional[Decimal]:
|
||||
"""Get last sale price in TON."""
|
||||
if self.last_sale_price is None:
|
||||
return None
|
||||
return Decimal(self.last_sale_price) / Decimal("1000000000")
|
||||
|
||||
class DeFiPosition(Base, UUIDMixin, TimestampMixin):
|
||||
"""Model for DeFi protocol positions."""
|
||||
__tablename__ = "defi_positions"
|
||||
|
||||
# User relationship
|
||||
user_id = Column(PostgreSQLUUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
|
||||
user = relationship("User", back_populates="defi_positions")
|
||||
|
||||
# Protocol details
|
||||
protocol_name = Column(String(100), nullable=False)
|
||||
protocol_address = Column(String(48), nullable=False)
|
||||
position_type = Column(String(50), nullable=False) # liquidity, lending, borrowing, etc.
|
||||
|
||||
# Position details
|
||||
token_a_address = Column(String(48), nullable=True)
|
||||
token_a_amount = Column(sa.BIGINT, nullable=False, default=0)
|
||||
token_b_address = Column(String(48), nullable=True)
|
||||
token_b_amount = Column(sa.BIGINT, nullable=False, default=0)
|
||||
|
||||
# Value tracking
|
||||
initial_value = Column(sa.BIGINT, nullable=False, default=0) # In nanotons
|
||||
current_value = Column(sa.BIGINT, nullable=False, default=0)
|
||||
last_value_update = Column(DateTime, nullable=True)
|
||||
|
||||
# Rewards and fees
|
||||
rewards_earned = Column(sa.BIGINT, nullable=False, default=0)
|
||||
fees_paid = Column(sa.BIGINT, nullable=False, default=0)
|
||||
|
||||
# Status
|
||||
status = Column(String(20), nullable=False, default="active") # active, closed, liquidated
|
||||
opened_tx = Column(String(64), nullable=False)
|
||||
closed_tx = Column(String(64), nullable=True)
|
||||
closed_at = Column(DateTime, nullable=True)
|
||||
|
||||
__table_args__ = (
|
||||
Index("idx_defi_user_protocol", "user_id", "protocol_name"),
|
||||
Index("idx_defi_position_type", "position_type"),
|
||||
Index("idx_defi_status", "status"),
|
||||
)
|
||||
|
||||
@validates('position_type')
|
||||
def validate_position_type(self, key, position_type):
|
||||
"""Validate position type."""
|
||||
allowed_types = {
|
||||
'liquidity', 'lending', 'borrowing', 'farming',
|
||||
'staking', 'options', 'futures', 'insurance'
|
||||
}
|
||||
if position_type not in allowed_types:
|
||||
raise ValueError(f"Invalid position type: {position_type}")
|
||||
return position_type
|
||||
|
||||
@validates('status')
|
||||
def validate_status(self, key, status):
|
||||
"""Validate position status."""
|
||||
allowed_statuses = {'active', 'closed', 'liquidated', 'expired'}
|
||||
if status not in allowed_statuses:
|
||||
raise ValueError(f"Invalid position status: {status}")
|
||||
return status
|
||||
|
||||
@property
|
||||
def current_value_tons(self) -> Decimal:
|
||||
"""Get current value in TON."""
|
||||
return Decimal(self.current_value) / Decimal("1000000000")
|
||||
|
||||
@property
|
||||
def pnl_tons(self) -> Decimal:
|
||||
"""Get profit/loss in TON."""
|
||||
return Decimal(self.current_value - self.initial_value) / Decimal("1000000000")
|
||||
|
||||
@property
|
||||
def pnl_percentage(self) -> Decimal:
|
||||
"""Get profit/loss percentage."""
|
||||
if self.initial_value == 0:
|
||||
return Decimal("0")
|
||||
return (Decimal(self.current_value - self.initial_value) / Decimal(self.initial_value)) * 100
|
||||
@@ -0,0 +1,731 @@
|
||||
"""
|
||||
Content models with async support and enhanced features
|
||||
"""
|
||||
import hashlib
|
||||
import mimetypes
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from typing import Optional, List, Dict, Any, Union
|
||||
from urllib.parse import urljoin
|
||||
|
||||
from sqlalchemy import Column, String, Integer, BigInteger, Boolean, Text, ForeignKey, Index, text
|
||||
from sqlalchemy.dialects.postgresql import JSONB, ARRAY
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.future import select
|
||||
from sqlalchemy.orm import relationship
|
||||
import structlog
|
||||
|
||||
from app.core.models.base import BaseModel
|
||||
from app.core.config import settings, PROJECT_HOST
|
||||
|
||||
logger = structlog.get_logger(__name__)
|
||||
|
||||
|
||||
class ContentType(str, Enum):
|
||||
"""Content type enumeration"""
|
||||
AUDIO = "audio"
|
||||
VIDEO = "video"
|
||||
IMAGE = "image"
|
||||
TEXT = "text"
|
||||
DOCUMENT = "document"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
class ContentStatus(str, Enum):
|
||||
"""Content status enumeration"""
|
||||
UPLOADING = "uploading"
|
||||
PROCESSING = "processing"
|
||||
READY = "ready"
|
||||
FAILED = "failed"
|
||||
DISABLED = "disabled"
|
||||
DELETED = "deleted"
|
||||
|
||||
|
||||
class StorageType(str, Enum):
|
||||
"""Storage type enumeration"""
|
||||
LOCAL = "local"
|
||||
ONCHAIN = "onchain"
|
||||
IPFS = "ipfs"
|
||||
HYBRID = "hybrid"
|
||||
|
||||
|
||||
class LicenseType(str, Enum):
|
||||
"""License type enumeration"""
|
||||
LISTEN = "listen"
|
||||
USE = "use"
|
||||
RESALE = "resale"
|
||||
EXCLUSIVE = "exclusive"
|
||||
|
||||
|
||||
class StoredContent(BaseModel):
|
||||
"""Enhanced content storage model"""
|
||||
|
||||
__tablename__ = 'stored_content'
|
||||
|
||||
# Content identification
|
||||
hash = Column(
|
||||
String(128),
|
||||
nullable=False,
|
||||
unique=True,
|
||||
index=True,
|
||||
comment="Content hash (SHA-256 or custom)"
|
||||
)
|
||||
content_id = Column(
|
||||
String(256),
|
||||
nullable=True,
|
||||
index=True,
|
||||
comment="Content identifier (CID for IPFS)"
|
||||
)
|
||||
|
||||
# File information
|
||||
filename = Column(
|
||||
String(512),
|
||||
nullable=False,
|
||||
comment="Original filename"
|
||||
)
|
||||
file_size = Column(
|
||||
BigInteger,
|
||||
nullable=False,
|
||||
default=0,
|
||||
comment="File size in bytes"
|
||||
)
|
||||
mime_type = Column(
|
||||
String(128),
|
||||
nullable=True,
|
||||
comment="MIME type of the content"
|
||||
)
|
||||
|
||||
# Content type and storage
|
||||
content_type = Column(
|
||||
String(32),
|
||||
nullable=False,
|
||||
default=ContentType.UNKNOWN.value,
|
||||
index=True,
|
||||
comment="Content type category"
|
||||
)
|
||||
storage_type = Column(
|
||||
String(32),
|
||||
nullable=False,
|
||||
default=StorageType.LOCAL.value,
|
||||
index=True,
|
||||
comment="Storage type"
|
||||
)
|
||||
|
||||
# File path and URLs
|
||||
file_path = Column(
|
||||
String(1024),
|
||||
nullable=True,
|
||||
comment="Local file path"
|
||||
)
|
||||
external_url = Column(
|
||||
String(2048),
|
||||
nullable=True,
|
||||
comment="External URL for remote content"
|
||||
)
|
||||
|
||||
# Blockchain related
|
||||
onchain_index = Column(
|
||||
Integer,
|
||||
nullable=True,
|
||||
index=True,
|
||||
comment="On-chain index number"
|
||||
)
|
||||
owner_address = Column(
|
||||
String(256),
|
||||
nullable=True,
|
||||
index=True,
|
||||
comment="Blockchain owner address"
|
||||
)
|
||||
|
||||
# User and access
|
||||
user_id = Column(
|
||||
String(36), # UUID
|
||||
ForeignKey('users.id'),
|
||||
nullable=True,
|
||||
index=True,
|
||||
comment="User who uploaded the content"
|
||||
)
|
||||
|
||||
# Encryption and security
|
||||
encrypted = Column(
|
||||
Boolean,
|
||||
nullable=False,
|
||||
default=False,
|
||||
comment="Whether content is encrypted"
|
||||
)
|
||||
encryption_key_id = Column(
|
||||
String(36), # UUID
|
||||
ForeignKey('encryption_keys.id'),
|
||||
nullable=True,
|
||||
comment="Encryption key reference"
|
||||
)
|
||||
|
||||
# Processing status
|
||||
disabled = Column(
|
||||
Boolean,
|
||||
nullable=False,
|
||||
default=False,
|
||||
index=True,
|
||||
comment="Whether content is disabled"
|
||||
)
|
||||
|
||||
# Content metadata
|
||||
title = Column(
|
||||
String(512),
|
||||
nullable=True,
|
||||
comment="Content title"
|
||||
)
|
||||
description = Column(
|
||||
Text,
|
||||
nullable=True,
|
||||
comment="Content description"
|
||||
)
|
||||
tags = Column(
|
||||
ARRAY(String),
|
||||
nullable=False,
|
||||
default=list,
|
||||
comment="Content tags"
|
||||
)
|
||||
|
||||
# Media-specific metadata
|
||||
duration = Column(
|
||||
Integer,
|
||||
nullable=True,
|
||||
comment="Duration in seconds (for audio/video)"
|
||||
)
|
||||
width = Column(
|
||||
Integer,
|
||||
nullable=True,
|
||||
comment="Width in pixels (for images/video)"
|
||||
)
|
||||
height = Column(
|
||||
Integer,
|
||||
nullable=True,
|
||||
comment="Height in pixels (for images/video)"
|
||||
)
|
||||
bitrate = Column(
|
||||
Integer,
|
||||
nullable=True,
|
||||
comment="Bitrate (for audio/video)"
|
||||
)
|
||||
|
||||
# Conversion and processing
|
||||
processing_status = Column(
|
||||
String(32),
|
||||
nullable=False,
|
||||
default=ContentStatus.READY.value,
|
||||
index=True,
|
||||
comment="Processing status"
|
||||
)
|
||||
conversion_data = Column(
|
||||
JSONB,
|
||||
nullable=False,
|
||||
default=dict,
|
||||
comment="Conversion and processing data"
|
||||
)
|
||||
|
||||
# Statistics
|
||||
download_count = Column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
comment="Number of downloads"
|
||||
)
|
||||
view_count = Column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
comment="Number of views"
|
||||
)
|
||||
|
||||
# Relationships
|
||||
user = relationship('User', back_populates='content_items')
|
||||
encryption_key = relationship('EncryptionKey', back_populates='content_items')
|
||||
user_contents = relationship('UserContent', back_populates='content')
|
||||
user_actions = relationship('UserAction', back_populates='content')
|
||||
|
||||
# Indexes for performance
|
||||
__table_args__ = (
|
||||
Index('idx_content_hash', 'hash'),
|
||||
Index('idx_content_user_type', 'user_id', 'content_type'),
|
||||
Index('idx_content_storage_status', 'storage_type', 'status'),
|
||||
Index('idx_content_onchain', 'onchain_index'),
|
||||
Index('idx_content_created', 'created_at'),
|
||||
Index('idx_content_disabled', 'disabled'),
|
||||
)
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""String representation"""
|
||||
return f"StoredContent({self.id}, hash={self.hash[:8]}..., filename={self.filename})"
|
||||
|
||||
@property
|
||||
def file_extension(self) -> str:
|
||||
"""Get file extension"""
|
||||
return Path(self.filename).suffix.lower()
|
||||
|
||||
@property
|
||||
def web_url(self) -> str:
|
||||
"""Get web accessible URL"""
|
||||
if self.external_url:
|
||||
return self.external_url
|
||||
if self.hash:
|
||||
return urljoin(str(PROJECT_HOST), f"/api/v1.5/storage/{self.hash}")
|
||||
return ""
|
||||
|
||||
@property
|
||||
def download_url(self) -> str:
|
||||
"""Get download URL"""
|
||||
if self.hash:
|
||||
return urljoin(str(PROJECT_HOST), f"/api/v1/storage/{self.hash}")
|
||||
return ""
|
||||
|
||||
@property
|
||||
def is_media(self) -> bool:
|
||||
"""Check if content is media (audio/video/image)"""
|
||||
return self.content_type in [ContentType.AUDIO, ContentType.VIDEO, ContentType.IMAGE]
|
||||
|
||||
@property
|
||||
def is_processed(self) -> bool:
|
||||
"""Check if content is fully processed"""
|
||||
return self.processing_status == ContentStatus.READY.value
|
||||
|
||||
@property
|
||||
def cache_key(self) -> str:
|
||||
"""Override cache key to use hash"""
|
||||
return f"content:hash:{self.hash}"
|
||||
|
||||
def detect_content_type(self) -> ContentType:
|
||||
"""Detect content type from MIME type"""
|
||||
if not self.mime_type:
|
||||
# Try to guess from extension
|
||||
mime_type, _ = mimetypes.guess_type(self.filename)
|
||||
self.mime_type = mime_type
|
||||
|
||||
if self.mime_type:
|
||||
if self.mime_type.startswith('audio/'):
|
||||
return ContentType.AUDIO
|
||||
elif self.mime_type.startswith('video/'):
|
||||
return ContentType.VIDEO
|
||||
elif self.mime_type.startswith('image/'):
|
||||
return ContentType.IMAGE
|
||||
elif self.mime_type.startswith('text/'):
|
||||
return ContentType.TEXT
|
||||
elif 'application/' in self.mime_type:
|
||||
return ContentType.DOCUMENT
|
||||
|
||||
return ContentType.UNKNOWN
|
||||
|
||||
def calculate_hash(self, file_data: bytes) -> str:
|
||||
"""Calculate hash for file data"""
|
||||
return hashlib.sha256(file_data).hexdigest()
|
||||
|
||||
def set_conversion_data(self, key: str, value: Any) -> None:
|
||||
"""Set conversion data"""
|
||||
if not self.conversion_data:
|
||||
self.conversion_data = {}
|
||||
self.conversion_data[key] = value
|
||||
|
||||
def get_conversion_data(self, key: str, default: Any = None) -> Any:
|
||||
"""Get conversion data"""
|
||||
if not self.conversion_data:
|
||||
return default
|
||||
return self.conversion_data.get(key, default)
|
||||
|
||||
def add_tag(self, tag: str) -> None:
|
||||
"""Add tag to content"""
|
||||
if not self.tags:
|
||||
self.tags = []
|
||||
tag = tag.strip().lower()
|
||||
if tag and tag not in self.tags:
|
||||
self.tags.append(tag)
|
||||
|
||||
def remove_tag(self, tag: str) -> None:
|
||||
"""Remove tag from content"""
|
||||
if self.tags:
|
||||
tag = tag.strip().lower()
|
||||
if tag in self.tags:
|
||||
self.tags.remove(tag)
|
||||
|
||||
def increment_download_count(self) -> None:
|
||||
"""Increment download counter"""
|
||||
self.download_count += 1
|
||||
|
||||
def increment_view_count(self) -> None:
|
||||
"""Increment view counter"""
|
||||
self.view_count += 1
|
||||
|
||||
@classmethod
|
||||
async def get_by_hash(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
content_hash: str
|
||||
) -> Optional['StoredContent']:
|
||||
"""Get content by hash"""
|
||||
try:
|
||||
stmt = select(cls).where(cls.hash == content_hash)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
except Exception as e:
|
||||
logger.error("Error getting content by hash", hash=content_hash, error=str(e))
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
async def get_by_user(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
user_id: str,
|
||||
content_type: Optional[ContentType] = None,
|
||||
limit: Optional[int] = None,
|
||||
offset: Optional[int] = None
|
||||
) -> List['StoredContent']:
|
||||
"""Get content by user"""
|
||||
try:
|
||||
stmt = select(cls).where(cls.user_id == user_id)
|
||||
|
||||
if content_type:
|
||||
stmt = stmt.where(cls.content_type == content_type.value)
|
||||
|
||||
stmt = stmt.order_by(cls.created_at.desc())
|
||||
|
||||
if offset:
|
||||
stmt = stmt.offset(offset)
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await session.execute(stmt)
|
||||
return result.scalars().all()
|
||||
except Exception as e:
|
||||
logger.error("Error getting content by user", user_id=user_id, error=str(e))
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
async def get_recent(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
days: int = 7,
|
||||
content_type: Optional[ContentType] = None,
|
||||
limit: Optional[int] = None
|
||||
) -> List['StoredContent']:
|
||||
"""Get recent content"""
|
||||
try:
|
||||
cutoff_date = datetime.utcnow() - timedelta(days=days)
|
||||
stmt = select(cls).where(
|
||||
cls.created_at >= cutoff_date,
|
||||
cls.disabled == False,
|
||||
cls.processing_status == ContentStatus.READY.value
|
||||
)
|
||||
|
||||
if content_type:
|
||||
stmt = stmt.where(cls.content_type == content_type.value)
|
||||
|
||||
stmt = stmt.order_by(cls.created_at.desc())
|
||||
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await session.execute(stmt)
|
||||
return result.scalars().all()
|
||||
except Exception as e:
|
||||
logger.error("Error getting recent content", days=days, error=str(e))
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
async def search(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
query: str,
|
||||
content_type: Optional[ContentType] = None,
|
||||
limit: Optional[int] = None,
|
||||
offset: Optional[int] = None
|
||||
) -> List['StoredContent']:
|
||||
"""Search content by title and description"""
|
||||
try:
|
||||
search_pattern = f"%{query.lower()}%"
|
||||
stmt = select(cls).where(
|
||||
(cls.title.ilike(search_pattern)) |
|
||||
(cls.description.ilike(search_pattern)) |
|
||||
(cls.filename.ilike(search_pattern)),
|
||||
cls.disabled == False,
|
||||
cls.processing_status == ContentStatus.READY.value
|
||||
)
|
||||
|
||||
if content_type:
|
||||
stmt = stmt.where(cls.content_type == content_type.value)
|
||||
|
||||
stmt = stmt.order_by(cls.created_at.desc())
|
||||
|
||||
if offset:
|
||||
stmt = stmt.offset(offset)
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await session.execute(stmt)
|
||||
return result.scalars().all()
|
||||
except Exception as e:
|
||||
logger.error("Error searching content", query=query, error=str(e))
|
||||
return []
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Convert to dictionary with additional computed fields"""
|
||||
data = super().to_dict()
|
||||
data.update({
|
||||
'web_url': self.web_url,
|
||||
'download_url': self.download_url,
|
||||
'file_extension': self.file_extension,
|
||||
'is_media': self.is_media,
|
||||
'is_processed': self.is_processed
|
||||
})
|
||||
return data
|
||||
|
||||
|
||||
class UserContent(BaseModel):
|
||||
"""User content ownership and licensing"""
|
||||
|
||||
__tablename__ = 'user_content'
|
||||
|
||||
# Content relationship
|
||||
content_id = Column(
|
||||
String(36), # UUID
|
||||
ForeignKey('stored_content.id'),
|
||||
nullable=False,
|
||||
index=True,
|
||||
comment="Reference to stored content"
|
||||
)
|
||||
user_id = Column(
|
||||
String(36), # UUID
|
||||
ForeignKey('users.id'),
|
||||
nullable=False,
|
||||
index=True,
|
||||
comment="User who owns this content"
|
||||
)
|
||||
|
||||
# License information
|
||||
license_type = Column(
|
||||
String(32),
|
||||
nullable=False,
|
||||
default=LicenseType.LISTEN.value,
|
||||
comment="Type of license"
|
||||
)
|
||||
|
||||
# Blockchain data
|
||||
onchain_address = Column(
|
||||
String(256),
|
||||
nullable=True,
|
||||
index=True,
|
||||
comment="On-chain contract address"
|
||||
)
|
||||
owner_address = Column(
|
||||
String(256),
|
||||
nullable=True,
|
||||
index=True,
|
||||
comment="Blockchain owner address"
|
||||
)
|
||||
|
||||
# Transaction data
|
||||
purchase_transaction = Column(
|
||||
String(128),
|
||||
nullable=True,
|
||||
comment="Purchase transaction hash"
|
||||
)
|
||||
purchase_amount = Column(
|
||||
BigInteger,
|
||||
nullable=True,
|
||||
comment="Purchase amount in minimal units"
|
||||
)
|
||||
|
||||
# Wallet connection
|
||||
wallet_connection_id = Column(
|
||||
String(36), # UUID
|
||||
ForeignKey('wallet_connections.id'),
|
||||
nullable=True,
|
||||
comment="Wallet connection used for purchase"
|
||||
)
|
||||
|
||||
# Access control
|
||||
access_granted = Column(
|
||||
Boolean,
|
||||
nullable=False,
|
||||
default=False,
|
||||
comment="Whether access is granted"
|
||||
)
|
||||
access_expires_at = Column(
|
||||
DateTime,
|
||||
nullable=True,
|
||||
comment="When access expires (for temporary licenses)"
|
||||
)
|
||||
|
||||
# Usage tracking
|
||||
download_count = Column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
comment="Number of downloads by this user"
|
||||
)
|
||||
last_accessed = Column(
|
||||
DateTime,
|
||||
nullable=True,
|
||||
comment="Last access timestamp"
|
||||
)
|
||||
|
||||
# Relationships
|
||||
user = relationship('User', back_populates='content_items')
|
||||
content = relationship('StoredContent', back_populates='user_contents')
|
||||
wallet_connection = relationship('WalletConnection', back_populates='user_contents')
|
||||
|
||||
# Indexes
|
||||
__table_args__ = (
|
||||
Index('idx_user_content_user', 'user_id'),
|
||||
Index('idx_user_content_content', 'content_id'),
|
||||
Index('idx_user_content_onchain', 'onchain_address'),
|
||||
Index('idx_user_content_owner', 'owner_address'),
|
||||
Index('idx_user_content_status', 'status'),
|
||||
)
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""String representation"""
|
||||
return f"UserContent({self.id}, user={self.user_id}, content={self.content_id})"
|
||||
|
||||
@property
|
||||
def is_expired(self) -> bool:
|
||||
"""Check if access has expired"""
|
||||
if not self.access_expires_at:
|
||||
return False
|
||||
return datetime.utcnow() > self.access_expires_at
|
||||
|
||||
@property
|
||||
def is_accessible(self) -> bool:
|
||||
"""Check if content is accessible"""
|
||||
return self.access_granted and not self.is_expired and self.status == 'active'
|
||||
|
||||
def grant_access(self, expires_at: Optional[datetime] = None) -> None:
|
||||
"""Grant access to content"""
|
||||
self.access_granted = True
|
||||
self.access_expires_at = expires_at
|
||||
self.last_accessed = datetime.utcnow()
|
||||
|
||||
def revoke_access(self) -> None:
|
||||
"""Revoke access to content"""
|
||||
self.access_granted = False
|
||||
|
||||
def record_download(self) -> None:
|
||||
"""Record a download"""
|
||||
self.download_count += 1
|
||||
self.last_accessed = datetime.utcnow()
|
||||
|
||||
@classmethod
|
||||
async def get_user_access(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
user_id: str,
|
||||
content_id: str
|
||||
) -> Optional['UserContent']:
|
||||
"""Get user access to specific content"""
|
||||
try:
|
||||
stmt = select(cls).where(
|
||||
cls.user_id == user_id,
|
||||
cls.content_id == content_id,
|
||||
cls.status == 'active'
|
||||
)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
except Exception as e:
|
||||
logger.error("Error getting user access", user_id=user_id, content_id=content_id, error=str(e))
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
async def get_user_content(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
user_id: str,
|
||||
limit: Optional[int] = None,
|
||||
offset: Optional[int] = None
|
||||
) -> List['UserContent']:
|
||||
"""Get all content accessible by user"""
|
||||
try:
|
||||
stmt = select(cls).where(
|
||||
cls.user_id == user_id,
|
||||
cls.status == 'active',
|
||||
cls.access_granted == True
|
||||
).order_by(cls.created_at.desc())
|
||||
|
||||
if offset:
|
||||
stmt = stmt.offset(offset)
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await session.execute(stmt)
|
||||
return result.scalars().all()
|
||||
except Exception as e:
|
||||
logger.error("Error getting user content", user_id=user_id, error=str(e))
|
||||
return []
|
||||
|
||||
|
||||
class EncryptionKey(BaseModel):
|
||||
"""Encryption key management"""
|
||||
|
||||
__tablename__ = 'encryption_keys'
|
||||
|
||||
# Key identification
|
||||
key_hash = Column(
|
||||
String(128),
|
||||
nullable=False,
|
||||
unique=True,
|
||||
index=True,
|
||||
comment="Hash of the encryption key"
|
||||
)
|
||||
algorithm = Column(
|
||||
String(32),
|
||||
nullable=False,
|
||||
default="AES-256-GCM",
|
||||
comment="Encryption algorithm used"
|
||||
)
|
||||
|
||||
# Key metadata
|
||||
purpose = Column(
|
||||
String(64),
|
||||
nullable=False,
|
||||
comment="Purpose of the key (content, user_data, etc.)"
|
||||
)
|
||||
|
||||
# Access control
|
||||
owner_id = Column(
|
||||
String(36), # UUID
|
||||
ForeignKey('users.id'),
|
||||
nullable=True,
|
||||
comment="Key owner (if user-specific)"
|
||||
)
|
||||
|
||||
# Key lifecycle
|
||||
expires_at = Column(
|
||||
DateTime,
|
||||
nullable=True,
|
||||
comment="Key expiration timestamp"
|
||||
)
|
||||
revoked_at = Column(
|
||||
DateTime,
|
||||
nullable=True,
|
||||
comment="Key revocation timestamp"
|
||||
)
|
||||
|
||||
# Relationships
|
||||
owner = relationship('User', back_populates='encryption_keys')
|
||||
content_items = relationship('StoredContent', back_populates='encryption_key')
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""String representation"""
|
||||
return f"EncryptionKey({self.id}, hash={self.key_hash[:8]}...)"
|
||||
|
||||
@property
|
||||
def is_valid(self) -> bool:
|
||||
"""Check if key is valid (not expired or revoked)"""
|
||||
now = datetime.utcnow()
|
||||
if self.revoked_at and self.revoked_at <= now:
|
||||
return False
|
||||
if self.expires_at and self.expires_at <= now:
|
||||
return False
|
||||
return True
|
||||
|
||||
def revoke(self) -> None:
|
||||
"""Revoke the key"""
|
||||
self.revoked_at = datetime.utcnow()
|
||||
@@ -0,0 +1,388 @@
|
||||
"""Compatible content models for MariaDB."""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Optional, List, Dict, Any
|
||||
from sqlalchemy import Column, String, Boolean, Text, Integer, DateTime, BigInteger, Index, ForeignKey
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.core.models.base_compatible import BaseModel
|
||||
|
||||
|
||||
class Content(BaseModel):
|
||||
"""Content model compatible with existing MariaDB schema."""
|
||||
|
||||
__tablename__ = "content"
|
||||
|
||||
# Basic content information
|
||||
user_id = Column(Integer, ForeignKey('users.id'), nullable=False, index=True)
|
||||
filename = Column(String(255), nullable=False)
|
||||
original_filename = Column(String(255), nullable=False)
|
||||
file_path = Column(String(500), nullable=False)
|
||||
|
||||
# File metadata
|
||||
file_size = Column(BigInteger, nullable=False) # bytes
|
||||
file_type = Column(String(100), nullable=False)
|
||||
mime_type = Column(String(100), nullable=False)
|
||||
file_extension = Column(String(10), nullable=False)
|
||||
|
||||
# Content metadata
|
||||
title = Column(String(255), nullable=True)
|
||||
description = Column(Text, nullable=True)
|
||||
tags = Column(Text, nullable=True) # JSON or comma-separated
|
||||
|
||||
# Status and visibility
|
||||
is_public = Column(Boolean, default=False, nullable=False)
|
||||
is_active = Column(Boolean, default=True, nullable=False)
|
||||
is_indexed = Column(Boolean, default=False, nullable=False)
|
||||
is_converted = Column(Boolean, default=False, nullable=False)
|
||||
|
||||
# Access and security
|
||||
access_password = Column(String(255), nullable=True)
|
||||
download_count = Column(Integer, default=0, nullable=False)
|
||||
view_count = Column(Integer, default=0, nullable=False)
|
||||
|
||||
# Processing status
|
||||
processing_status = Column(String(50), default="pending", nullable=False)
|
||||
processing_error = Column(Text, nullable=True)
|
||||
processing_started = Column(DateTime, nullable=True)
|
||||
processing_completed = Column(DateTime, nullable=True)
|
||||
|
||||
# File hashes for integrity
|
||||
md5_hash = Column(String(32), nullable=True, index=True)
|
||||
sha256_hash = Column(String(64), nullable=True, index=True)
|
||||
|
||||
# Thumbnails and previews
|
||||
thumbnail_path = Column(String(500), nullable=True)
|
||||
preview_path = Column(String(500), nullable=True)
|
||||
|
||||
# TON Blockchain integration
|
||||
ton_transaction_hash = Column(String(100), nullable=True, index=True)
|
||||
ton_storage_proof = Column(Text, nullable=True)
|
||||
ton_storage_fee = Column(BigInteger, default=0, nullable=False) # nanotons
|
||||
|
||||
# Expiration and cleanup
|
||||
expires_at = Column(DateTime, nullable=True)
|
||||
auto_delete = Column(Boolean, default=False, nullable=False)
|
||||
|
||||
# Relationships
|
||||
user = relationship("User", back_populates="content")
|
||||
|
||||
# Table indexes for performance
|
||||
__table_args__ = (
|
||||
Index('idx_content_user_active', 'user_id', 'is_active'),
|
||||
Index('idx_content_public_indexed', 'is_public', 'is_indexed'),
|
||||
Index('idx_content_file_type', 'file_type', 'mime_type'),
|
||||
Index('idx_content_created', 'created_at'),
|
||||
Index('idx_content_size', 'file_size'),
|
||||
Index('idx_content_processing', 'processing_status'),
|
||||
Index('idx_content_ton_tx', 'ton_transaction_hash'),
|
||||
Index('idx_content_expires', 'expires_at', 'auto_delete'),
|
||||
)
|
||||
|
||||
def is_expired(self) -> bool:
|
||||
"""Check if content is expired."""
|
||||
if not self.expires_at:
|
||||
return False
|
||||
return datetime.utcnow() > self.expires_at
|
||||
|
||||
def is_image(self) -> bool:
|
||||
"""Check if content is an image."""
|
||||
return self.file_type.lower() in ['image', 'img'] or \
|
||||
self.mime_type.startswith('image/')
|
||||
|
||||
def is_video(self) -> bool:
|
||||
"""Check if content is a video."""
|
||||
return self.file_type.lower() == 'video' or \
|
||||
self.mime_type.startswith('video/')
|
||||
|
||||
def is_document(self) -> bool:
|
||||
"""Check if content is a document."""
|
||||
return self.file_type.lower() in ['document', 'doc', 'pdf'] or \
|
||||
self.mime_type in ['application/pdf', 'application/msword', 'text/plain']
|
||||
|
||||
def get_file_size_human(self) -> str:
|
||||
"""Get human-readable file size."""
|
||||
size = self.file_size
|
||||
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
|
||||
if size < 1024.0:
|
||||
return f"{size:.1f} {unit}"
|
||||
size /= 1024.0
|
||||
return f"{size:.1f} PB"
|
||||
|
||||
def increment_download_count(self) -> None:
|
||||
"""Increment download counter."""
|
||||
self.download_count += 1
|
||||
|
||||
def increment_view_count(self) -> None:
|
||||
"""Increment view counter."""
|
||||
self.view_count += 1
|
||||
|
||||
def mark_as_indexed(self) -> None:
|
||||
"""Mark content as indexed."""
|
||||
self.is_indexed = True
|
||||
|
||||
def mark_as_converted(self) -> None:
|
||||
"""Mark content as converted."""
|
||||
self.is_converted = True
|
||||
self.processing_status = "completed"
|
||||
self.processing_completed = datetime.utcnow()
|
||||
|
||||
def set_processing_error(self, error: str) -> None:
|
||||
"""Set processing error."""
|
||||
self.processing_status = "error"
|
||||
self.processing_error = error
|
||||
self.processing_completed = datetime.utcnow()
|
||||
|
||||
def start_processing(self) -> None:
|
||||
"""Mark processing as started."""
|
||||
self.processing_status = "processing"
|
||||
self.processing_started = datetime.utcnow()
|
||||
self.processing_error = None
|
||||
|
||||
def get_tags_list(self) -> List[str]:
|
||||
"""Get tags as list."""
|
||||
if not self.tags:
|
||||
return []
|
||||
|
||||
# Try to parse as JSON first, fallback to comma-separated
|
||||
try:
|
||||
import json
|
||||
return json.loads(self.tags)
|
||||
except:
|
||||
return [tag.strip() for tag in self.tags.split(',') if tag.strip()]
|
||||
|
||||
def set_tags_list(self, tags: List[str]) -> None:
|
||||
"""Set tags from list."""
|
||||
import json
|
||||
self.tags = json.dumps(tags) if tags else None
|
||||
|
||||
def to_dict(self, include_sensitive: bool = False) -> Dict[str, Any]:
|
||||
"""Convert to dictionary with option to exclude sensitive data."""
|
||||
exclude = set()
|
||||
if not include_sensitive:
|
||||
exclude.update({"access_password", "file_path", "processing_error"})
|
||||
|
||||
data = super().to_dict(exclude=exclude)
|
||||
|
||||
# Add computed fields
|
||||
data.update({
|
||||
"file_size_human": self.get_file_size_human(),
|
||||
"is_image": self.is_image(),
|
||||
"is_video": self.is_video(),
|
||||
"is_document": self.is_document(),
|
||||
"is_expired": self.is_expired(),
|
||||
"tags_list": self.get_tags_list(),
|
||||
})
|
||||
|
||||
return data
|
||||
|
||||
def to_public_dict(self) -> Dict[str, Any]:
|
||||
"""Convert to public dictionary (minimal content info)."""
|
||||
return {
|
||||
"id": self.id,
|
||||
"filename": self.filename,
|
||||
"title": self.title,
|
||||
"description": self.description,
|
||||
"file_type": self.file_type,
|
||||
"file_size": self.file_size,
|
||||
"file_size_human": self.get_file_size_human(),
|
||||
"is_image": self.is_image(),
|
||||
"is_video": self.is_video(),
|
||||
"is_document": self.is_document(),
|
||||
"download_count": self.download_count,
|
||||
"view_count": self.view_count,
|
||||
"tags_list": self.get_tags_list(),
|
||||
"created_at": self.created_at.isoformat() if self.created_at else None,
|
||||
}
|
||||
|
||||
|
||||
class ContentShare(BaseModel):
|
||||
"""Content sharing model for tracking shared content."""
|
||||
|
||||
__tablename__ = "content_shares"
|
||||
|
||||
content_id = Column(Integer, ForeignKey('content.id'), nullable=False, index=True)
|
||||
user_id = Column(Integer, ForeignKey('users.id'), nullable=True, index=True) # Can be null for anonymous shares
|
||||
|
||||
# Share metadata
|
||||
share_token = Column(String(100), unique=True, nullable=False, index=True)
|
||||
share_url = Column(String(500), nullable=False)
|
||||
|
||||
# Share settings
|
||||
is_active = Column(Boolean, default=True, nullable=False)
|
||||
is_password_protected = Column(Boolean, default=False, nullable=False)
|
||||
share_password = Column(String(255), nullable=True)
|
||||
|
||||
# Access control
|
||||
max_downloads = Column(Integer, nullable=True) # Null = unlimited
|
||||
download_count = Column(Integer, default=0, nullable=False)
|
||||
view_count = Column(Integer, default=0, nullable=False)
|
||||
|
||||
# Time limits
|
||||
expires_at = Column(DateTime, nullable=True)
|
||||
|
||||
# Tracking
|
||||
ip_address = Column(String(45), nullable=True)
|
||||
user_agent = Column(Text, nullable=True)
|
||||
|
||||
# Relationships
|
||||
content = relationship("Content")
|
||||
user = relationship("User")
|
||||
|
||||
__table_args__ = (
|
||||
Index('idx_shares_content_active', 'content_id', 'is_active'),
|
||||
Index('idx_shares_token', 'share_token'),
|
||||
Index('idx_shares_expires', 'expires_at'),
|
||||
)
|
||||
|
||||
def is_expired(self) -> bool:
|
||||
"""Check if share is expired."""
|
||||
if not self.expires_at:
|
||||
return False
|
||||
return datetime.utcnow() > self.expires_at
|
||||
|
||||
def is_download_limit_reached(self) -> bool:
|
||||
"""Check if download limit is reached."""
|
||||
if not self.max_downloads:
|
||||
return False
|
||||
return self.download_count >= self.max_downloads
|
||||
|
||||
def is_valid(self) -> bool:
|
||||
"""Check if share is valid."""
|
||||
return (self.is_active and
|
||||
not self.is_expired() and
|
||||
not self.is_download_limit_reached())
|
||||
|
||||
def increment_download(self) -> bool:
|
||||
"""Increment download count and return if still valid."""
|
||||
if not self.is_valid():
|
||||
return False
|
||||
|
||||
self.download_count += 1
|
||||
return self.is_valid()
|
||||
|
||||
def increment_view(self) -> None:
|
||||
"""Increment view count."""
|
||||
self.view_count += 1
|
||||
|
||||
|
||||
class ContentMetadata(BaseModel):
|
||||
"""Extended metadata for content files."""
|
||||
|
||||
__tablename__ = "content_metadata"
|
||||
|
||||
content_id = Column(Integer, ForeignKey('content.id'), unique=True, nullable=False, index=True)
|
||||
|
||||
# Image metadata
|
||||
image_width = Column(Integer, nullable=True)
|
||||
image_height = Column(Integer, nullable=True)
|
||||
image_dpi = Column(Integer, nullable=True)
|
||||
image_color_space = Column(String(50), nullable=True)
|
||||
|
||||
# Video metadata
|
||||
video_duration = Column(Integer, nullable=True) # seconds
|
||||
video_bitrate = Column(Integer, nullable=True)
|
||||
video_fps = Column(Integer, nullable=True)
|
||||
video_resolution = Column(String(20), nullable=True) # e.g., "1920x1080"
|
||||
video_codec = Column(String(50), nullable=True)
|
||||
|
||||
# Audio metadata
|
||||
audio_duration = Column(Integer, nullable=True) # seconds
|
||||
audio_bitrate = Column(Integer, nullable=True)
|
||||
audio_sample_rate = Column(Integer, nullable=True)
|
||||
audio_channels = Column(Integer, nullable=True)
|
||||
audio_codec = Column(String(50), nullable=True)
|
||||
|
||||
# Document metadata
|
||||
document_pages = Column(Integer, nullable=True)
|
||||
document_words = Column(Integer, nullable=True)
|
||||
document_language = Column(String(10), nullable=True)
|
||||
document_author = Column(String(255), nullable=True)
|
||||
|
||||
# EXIF data (JSON)
|
||||
exif_data = Column(Text, nullable=True)
|
||||
|
||||
# GPS coordinates
|
||||
gps_latitude = Column(String(50), nullable=True)
|
||||
gps_longitude = Column(String(50), nullable=True)
|
||||
gps_altitude = Column(String(50), nullable=True)
|
||||
|
||||
# Technical metadata
|
||||
compression_ratio = Column(String(20), nullable=True)
|
||||
quality_score = Column(Integer, nullable=True) # 0-100
|
||||
|
||||
# Relationships
|
||||
content = relationship("Content")
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Convert metadata to dictionary."""
|
||||
data = super().to_dict(exclude={"content_id"})
|
||||
|
||||
# Parse EXIF data if present
|
||||
if self.exif_data:
|
||||
try:
|
||||
import json
|
||||
data["exif_data"] = json.loads(self.exif_data)
|
||||
except:
|
||||
data["exif_data"] = None
|
||||
|
||||
return data
|
||||
|
||||
def set_exif_data(self, exif_dict: Dict[str, Any]) -> None:
|
||||
"""Set EXIF data from dictionary."""
|
||||
if exif_dict:
|
||||
import json
|
||||
self.exif_data = json.dumps(exif_dict)
|
||||
else:
|
||||
self.exif_data = None
|
||||
|
||||
def get_exif_data(self) -> Optional[Dict[str, Any]]:
|
||||
"""Get EXIF data as dictionary."""
|
||||
if not self.exif_data:
|
||||
return None
|
||||
|
||||
try:
|
||||
import json
|
||||
return json.loads(self.exif_data)
|
||||
except:
|
||||
return None
|
||||
|
||||
|
||||
class ContentVersion(BaseModel):
|
||||
"""Content version history for tracking changes."""
|
||||
|
||||
__tablename__ = "content_versions"
|
||||
|
||||
content_id = Column(Integer, ForeignKey('content.id'), nullable=False, index=True)
|
||||
user_id = Column(Integer, ForeignKey('users.id'), nullable=False, index=True)
|
||||
|
||||
# Version information
|
||||
version_number = Column(Integer, nullable=False)
|
||||
version_name = Column(String(100), nullable=True)
|
||||
change_description = Column(Text, nullable=True)
|
||||
|
||||
# File information
|
||||
file_path = Column(String(500), nullable=False)
|
||||
file_size = Column(BigInteger, nullable=False)
|
||||
file_hash = Column(String(64), nullable=False)
|
||||
|
||||
# Status
|
||||
is_current = Column(Boolean, default=False, nullable=False)
|
||||
|
||||
# Relationships
|
||||
content = relationship("Content")
|
||||
user = relationship("User")
|
||||
|
||||
__table_args__ = (
|
||||
Index('idx_versions_content_number', 'content_id', 'version_number'),
|
||||
Index('idx_versions_current', 'content_id', 'is_current'),
|
||||
)
|
||||
|
||||
def mark_as_current(self) -> None:
|
||||
"""Mark this version as current."""
|
||||
self.is_current = True
|
||||
|
||||
|
||||
# Add relationship to User model
|
||||
# This would be added to the User model:
|
||||
# content = relationship("Content", back_populates="user")
|
||||
@@ -0,0 +1,420 @@
|
||||
"""
|
||||
User model with async support and enhanced security
|
||||
"""
|
||||
import hashlib
|
||||
import secrets
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Optional, List, Dict, Any
|
||||
from enum import Enum
|
||||
|
||||
from sqlalchemy import Column, String, BigInteger, Boolean, Integer, Index, text
|
||||
from sqlalchemy.dialects.postgresql import ARRAY, JSONB
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.future import select
|
||||
from sqlalchemy.orm import relationship
|
||||
import structlog
|
||||
|
||||
from app.core.models.base import BaseModel
|
||||
from app.core.config import settings
|
||||
|
||||
logger = structlog.get_logger(__name__)
|
||||
|
||||
|
||||
class UserRole(str, Enum):
|
||||
"""User role enumeration"""
|
||||
USER = "user"
|
||||
MODERATOR = "moderator"
|
||||
ADMIN = "admin"
|
||||
SUPER_ADMIN = "super_admin"
|
||||
|
||||
|
||||
class UserStatus(str, Enum):
|
||||
"""User status enumeration"""
|
||||
ACTIVE = "active"
|
||||
SUSPENDED = "suspended"
|
||||
BANNED = "banned"
|
||||
PENDING = "pending"
|
||||
|
||||
|
||||
class User(BaseModel):
|
||||
"""Enhanced User model with security and async support"""
|
||||
|
||||
__tablename__ = 'users'
|
||||
|
||||
# Telegram specific fields
|
||||
telegram_id = Column(
|
||||
BigInteger,
|
||||
nullable=False,
|
||||
unique=True,
|
||||
index=True,
|
||||
comment="Telegram user ID"
|
||||
)
|
||||
username = Column(
|
||||
String(512),
|
||||
nullable=True,
|
||||
index=True,
|
||||
comment="Telegram username"
|
||||
)
|
||||
first_name = Column(
|
||||
String(256),
|
||||
nullable=True,
|
||||
comment="User first name"
|
||||
)
|
||||
last_name = Column(
|
||||
String(256),
|
||||
nullable=True,
|
||||
comment="User last name"
|
||||
)
|
||||
|
||||
# Localization
|
||||
language_code = Column(
|
||||
String(8),
|
||||
nullable=False,
|
||||
default="en",
|
||||
comment="User language code"
|
||||
)
|
||||
|
||||
# Security and access control
|
||||
role = Column(
|
||||
String(32),
|
||||
nullable=False,
|
||||
default=UserRole.USER.value,
|
||||
index=True,
|
||||
comment="User role"
|
||||
)
|
||||
permissions = Column(
|
||||
ARRAY(String),
|
||||
nullable=False,
|
||||
default=list,
|
||||
comment="User permissions list"
|
||||
)
|
||||
|
||||
# Activity tracking
|
||||
last_activity = Column(
|
||||
"last_use", # Keep old column name for compatibility
|
||||
DateTime,
|
||||
nullable=False,
|
||||
default=datetime.utcnow,
|
||||
index=True,
|
||||
comment="Last user activity timestamp"
|
||||
)
|
||||
login_count = Column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
comment="Total login count"
|
||||
)
|
||||
|
||||
# Account status
|
||||
is_verified = Column(
|
||||
Boolean,
|
||||
nullable=False,
|
||||
default=False,
|
||||
comment="Whether user is verified"
|
||||
)
|
||||
is_premium = Column(
|
||||
Boolean,
|
||||
nullable=False,
|
||||
default=False,
|
||||
comment="Whether user has premium access"
|
||||
)
|
||||
|
||||
# Security settings
|
||||
two_factor_enabled = Column(
|
||||
Boolean,
|
||||
nullable=False,
|
||||
default=False,
|
||||
comment="Whether 2FA is enabled"
|
||||
)
|
||||
security_settings = Column(
|
||||
JSONB,
|
||||
nullable=False,
|
||||
default=dict,
|
||||
comment="User security settings"
|
||||
)
|
||||
|
||||
# Preferences
|
||||
preferences = Column(
|
||||
JSONB,
|
||||
nullable=False,
|
||||
default=dict,
|
||||
comment="User preferences and settings"
|
||||
)
|
||||
|
||||
# Statistics
|
||||
content_uploaded_count = Column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
comment="Number of content items uploaded"
|
||||
)
|
||||
content_purchased_count = Column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
comment="Number of content items purchased"
|
||||
)
|
||||
|
||||
# Rate limiting
|
||||
rate_limit_reset = Column(
|
||||
DateTime,
|
||||
nullable=True,
|
||||
comment="Rate limit reset timestamp"
|
||||
)
|
||||
rate_limit_count = Column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
comment="Current rate limit count"
|
||||
)
|
||||
|
||||
# Relationships
|
||||
balances = relationship('UserBalance', back_populates='user', cascade="all, delete-orphan")
|
||||
transactions = relationship('InternalTransaction', back_populates='user', cascade="all, delete-orphan")
|
||||
wallet_connections = relationship('WalletConnection', back_populates='user', cascade="all, delete-orphan")
|
||||
content_items = relationship('UserContent', back_populates='user', cascade="all, delete-orphan")
|
||||
actions = relationship('UserAction', back_populates='user', cascade="all, delete-orphan")
|
||||
activities = relationship('UserActivity', back_populates='user', cascade="all, delete-orphan")
|
||||
|
||||
# Indexes for performance
|
||||
__table_args__ = (
|
||||
Index('idx_users_telegram_id', 'telegram_id'),
|
||||
Index('idx_users_username', 'username'),
|
||||
Index('idx_users_role_status', 'role', 'status'),
|
||||
Index('idx_users_last_activity', 'last_activity'),
|
||||
Index('idx_users_created_at', 'created_at'),
|
||||
)
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""String representation"""
|
||||
return f"User({self.id}, telegram_id={self.telegram_id}, username={self.username})"
|
||||
|
||||
@property
|
||||
def full_name(self) -> str:
|
||||
"""Get user's full name"""
|
||||
parts = [self.first_name, self.last_name]
|
||||
return " ".join(filter(None, parts)) or self.username or f"User_{self.telegram_id}"
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Get user's display name"""
|
||||
return self.username or self.full_name
|
||||
|
||||
@property
|
||||
def is_admin(self) -> bool:
|
||||
"""Check if user is admin"""
|
||||
return self.role in [UserRole.ADMIN.value, UserRole.SUPER_ADMIN.value]
|
||||
|
||||
@property
|
||||
def is_moderator(self) -> bool:
|
||||
"""Check if user is moderator or higher"""
|
||||
return self.role in [UserRole.MODERATOR.value, UserRole.ADMIN.value, UserRole.SUPER_ADMIN.value]
|
||||
|
||||
@property
|
||||
def cache_key(self) -> str:
|
||||
"""Override cache key to include telegram_id"""
|
||||
return f"user:telegram:{self.telegram_id}"
|
||||
|
||||
def has_permission(self, permission: str) -> bool:
|
||||
"""Check if user has specific permission"""
|
||||
if self.is_admin:
|
||||
return True
|
||||
return permission in (self.permissions or [])
|
||||
|
||||
def add_permission(self, permission: str) -> None:
|
||||
"""Add permission to user"""
|
||||
if not self.permissions:
|
||||
self.permissions = []
|
||||
if permission not in self.permissions:
|
||||
self.permissions.append(permission)
|
||||
|
||||
def remove_permission(self, permission: str) -> None:
|
||||
"""Remove permission from user"""
|
||||
if self.permissions and permission in self.permissions:
|
||||
self.permissions.remove(permission)
|
||||
|
||||
def update_activity(self) -> None:
|
||||
"""Update user activity timestamp"""
|
||||
self.last_activity = datetime.utcnow()
|
||||
self.login_count += 1
|
||||
|
||||
def check_rate_limit(self, limit: int = None, window: int = None) -> bool:
|
||||
"""Check if user is within rate limits"""
|
||||
if self.is_admin:
|
||||
return True
|
||||
|
||||
limit = limit or settings.RATE_LIMIT_REQUESTS
|
||||
window = window or settings.RATE_LIMIT_WINDOW
|
||||
|
||||
now = datetime.utcnow()
|
||||
|
||||
# Reset counter if window has passed
|
||||
if not self.rate_limit_reset or now > self.rate_limit_reset:
|
||||
self.rate_limit_reset = now + timedelta(seconds=window)
|
||||
self.rate_limit_count = 0
|
||||
|
||||
return self.rate_limit_count < limit
|
||||
|
||||
def increment_rate_limit(self) -> None:
|
||||
"""Increment rate limit counter"""
|
||||
if not self.is_admin:
|
||||
self.rate_limit_count += 1
|
||||
|
||||
def set_preference(self, key: str, value: Any) -> None:
|
||||
"""Set user preference"""
|
||||
if not self.preferences:
|
||||
self.preferences = {}
|
||||
self.preferences[key] = value
|
||||
|
||||
def get_preference(self, key: str, default: Any = None) -> Any:
|
||||
"""Get user preference"""
|
||||
if not self.preferences:
|
||||
return default
|
||||
return self.preferences.get(key, default)
|
||||
|
||||
def set_security_setting(self, key: str, value: Any) -> None:
|
||||
"""Set security setting"""
|
||||
if not self.security_settings:
|
||||
self.security_settings = {}
|
||||
self.security_settings[key] = value
|
||||
|
||||
def get_security_setting(self, key: str, default: Any = None) -> Any:
|
||||
"""Get security setting"""
|
||||
if not self.security_settings:
|
||||
return default
|
||||
return self.security_settings.get(key, default)
|
||||
|
||||
def generate_api_token(self) -> str:
|
||||
"""Generate secure API token for user"""
|
||||
token_data = f"{self.id}:{self.telegram_id}:{datetime.utcnow().timestamp()}:{secrets.token_hex(16)}"
|
||||
return hashlib.sha256(token_data.encode()).hexdigest()
|
||||
|
||||
def can_upload_content(self) -> bool:
|
||||
"""Check if user can upload content"""
|
||||
if self.status != UserStatus.ACTIVE.value:
|
||||
return False
|
||||
|
||||
if not self.check_rate_limit(limit=10, window=3600): # 10 uploads per hour
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def can_purchase_content(self) -> bool:
|
||||
"""Check if user can purchase content"""
|
||||
return self.status == UserStatus.ACTIVE.value
|
||||
|
||||
@classmethod
|
||||
async def get_by_telegram_id(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
telegram_id: int
|
||||
) -> Optional['User']:
|
||||
"""Get user by Telegram ID"""
|
||||
try:
|
||||
stmt = select(cls).where(cls.telegram_id == telegram_id)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
except Exception as e:
|
||||
logger.error("Error getting user by telegram_id", telegram_id=telegram_id, error=str(e))
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
async def get_by_username(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
username: str
|
||||
) -> Optional['User']:
|
||||
"""Get user by username"""
|
||||
try:
|
||||
stmt = select(cls).where(cls.username == username)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
except Exception as e:
|
||||
logger.error("Error getting user by username", username=username, error=str(e))
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
async def get_active_users(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
days: int = 30,
|
||||
limit: Optional[int] = None
|
||||
) -> List['User']:
|
||||
"""Get active users within specified days"""
|
||||
try:
|
||||
cutoff_date = datetime.utcnow() - timedelta(days=days)
|
||||
stmt = select(cls).where(
|
||||
cls.last_activity >= cutoff_date,
|
||||
cls.status == UserStatus.ACTIVE.value
|
||||
).order_by(cls.last_activity.desc())
|
||||
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await session.execute(stmt)
|
||||
return result.scalars().all()
|
||||
except Exception as e:
|
||||
logger.error("Error getting active users", days=days, error=str(e))
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
async def get_admins(cls, session: AsyncSession) -> List['User']:
|
||||
"""Get all admin users"""
|
||||
try:
|
||||
stmt = select(cls).where(
|
||||
cls.role.in_([UserRole.ADMIN.value, UserRole.SUPER_ADMIN.value])
|
||||
)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalars().all()
|
||||
except Exception as e:
|
||||
logger.error("Error getting admin users", error=str(e))
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
async def create_from_telegram(
|
||||
cls,
|
||||
session: AsyncSession,
|
||||
telegram_id: int,
|
||||
username: Optional[str] = None,
|
||||
first_name: Optional[str] = None,
|
||||
last_name: Optional[str] = None,
|
||||
language_code: str = "en"
|
||||
) -> 'User':
|
||||
"""Create user from Telegram data"""
|
||||
user = cls(
|
||||
telegram_id=telegram_id,
|
||||
username=username,
|
||||
first_name=first_name,
|
||||
last_name=last_name,
|
||||
language_code=language_code,
|
||||
status=UserStatus.ACTIVE.value
|
||||
)
|
||||
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
|
||||
logger.info("User created from Telegram", telegram_id=telegram_id, user_id=user.id)
|
||||
return user
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Convert to dictionary with safe data"""
|
||||
data = super().to_dict()
|
||||
|
||||
# Remove sensitive fields
|
||||
sensitive_fields = ['security_settings', 'permissions']
|
||||
for field in sensitive_fields:
|
||||
data.pop(field, None)
|
||||
|
||||
return data
|
||||
|
||||
def to_public_dict(self) -> Dict[str, Any]:
|
||||
"""Convert to public dictionary with minimal data"""
|
||||
return {
|
||||
'id': str(self.id),
|
||||
'username': self.username,
|
||||
'display_name': self.display_name,
|
||||
'is_verified': self.is_verified,
|
||||
'is_premium': self.is_premium,
|
||||
'created_at': self.created_at.isoformat() if self.created_at else None
|
||||
}
|
||||
@@ -0,0 +1,247 @@
|
||||
"""Compatible user models for MariaDB."""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Optional, List, Dict, Any
|
||||
from sqlalchemy import Column, String, Boolean, Text, Integer, DateTime, Index
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.core.models.base_compatible import BaseModel
|
||||
|
||||
|
||||
class User(BaseModel):
|
||||
"""User model compatible with existing MariaDB schema."""
|
||||
|
||||
__tablename__ = "users"
|
||||
|
||||
# Basic user information
|
||||
username = Column(String(50), unique=True, nullable=False, index=True)
|
||||
email = Column(String(100), unique=True, nullable=True, index=True)
|
||||
password_hash = Column(String(255), nullable=False)
|
||||
|
||||
# User status and flags
|
||||
is_active = Column(Boolean, default=True, nullable=False)
|
||||
is_verified = Column(Boolean, default=False, nullable=False)
|
||||
is_admin = Column(Boolean, default=False, nullable=False)
|
||||
|
||||
# Profile information
|
||||
first_name = Column(String(50), nullable=True)
|
||||
last_name = Column(String(50), nullable=True)
|
||||
bio = Column(Text, nullable=True)
|
||||
avatar_url = Column(String(255), nullable=True)
|
||||
|
||||
# System tracking
|
||||
last_login = Column(DateTime, nullable=True)
|
||||
login_count = Column(Integer, default=0, nullable=False)
|
||||
|
||||
# Storage and limits
|
||||
storage_used = Column(Integer, default=0, nullable=False) # bytes
|
||||
storage_limit = Column(Integer, default=100*1024*1024, nullable=False) # 100MB default
|
||||
|
||||
# TON Blockchain integration
|
||||
ton_wallet_address = Column(String(100), nullable=True, index=True)
|
||||
ton_balance = Column(Integer, default=0, nullable=False) # nanotons
|
||||
|
||||
# License and subscription
|
||||
license_key = Column(String(100), nullable=True, index=True)
|
||||
license_expires = Column(DateTime, nullable=True)
|
||||
subscription_level = Column(String(20), default="free", nullable=False)
|
||||
|
||||
# API access
|
||||
api_key = Column(String(100), nullable=True, unique=True, index=True)
|
||||
api_calls_count = Column(Integer, default=0, nullable=False)
|
||||
api_calls_limit = Column(Integer, default=1000, nullable=False)
|
||||
|
||||
# Relationships will be defined when we create content models
|
||||
|
||||
# Table indexes for performance
|
||||
__table_args__ = (
|
||||
Index('idx_users_username_active', 'username', 'is_active'),
|
||||
Index('idx_users_email_verified', 'email', 'is_verified'),
|
||||
Index('idx_users_ton_wallet', 'ton_wallet_address'),
|
||||
Index('idx_users_license', 'license_key', 'license_expires'),
|
||||
)
|
||||
|
||||
def check_storage_limit(self, file_size: int) -> bool:
|
||||
"""Check if user can upload file of given size."""
|
||||
return (self.storage_used + file_size) <= self.storage_limit
|
||||
|
||||
def update_storage_usage(self, size_change: int) -> None:
|
||||
"""Update user's storage usage."""
|
||||
self.storage_used = max(0, self.storage_used + size_change)
|
||||
|
||||
def is_license_valid(self) -> bool:
|
||||
"""Check if user's license is valid."""
|
||||
if not self.license_key or not self.license_expires:
|
||||
return False
|
||||
return self.license_expires > datetime.utcnow()
|
||||
|
||||
def can_make_api_call(self) -> bool:
|
||||
"""Check if user can make API call."""
|
||||
return self.api_calls_count < self.api_calls_limit
|
||||
|
||||
def increment_api_calls(self) -> None:
|
||||
"""Increment API calls counter."""
|
||||
self.api_calls_count += 1
|
||||
|
||||
def reset_api_calls(self) -> None:
|
||||
"""Reset API calls counter (for monthly reset)."""
|
||||
self.api_calls_count = 0
|
||||
|
||||
def get_storage_usage_percent(self) -> float:
|
||||
"""Get storage usage as percentage."""
|
||||
if self.storage_limit == 0:
|
||||
return 0.0
|
||||
return (self.storage_used / self.storage_limit) * 100
|
||||
|
||||
def get_api_usage_percent(self) -> float:
|
||||
"""Get API usage as percentage."""
|
||||
if self.api_calls_limit == 0:
|
||||
return 0.0
|
||||
return (self.api_calls_count / self.api_calls_limit) * 100
|
||||
|
||||
def get_display_name(self) -> str:
|
||||
"""Get user's display name."""
|
||||
if self.first_name and self.last_name:
|
||||
return f"{self.first_name} {self.last_name}"
|
||||
elif self.first_name:
|
||||
return self.first_name
|
||||
return self.username
|
||||
|
||||
def to_dict(self, include_sensitive: bool = False) -> Dict[str, Any]:
|
||||
"""Convert to dictionary with option to exclude sensitive data."""
|
||||
exclude = set()
|
||||
if not include_sensitive:
|
||||
exclude.update({"password_hash", "api_key", "license_key"})
|
||||
|
||||
data = super().to_dict(exclude=exclude)
|
||||
|
||||
# Add computed fields
|
||||
data.update({
|
||||
"display_name": self.get_display_name(),
|
||||
"storage_usage_percent": self.get_storage_usage_percent(),
|
||||
"api_usage_percent": self.get_api_usage_percent(),
|
||||
"license_valid": self.is_license_valid(),
|
||||
})
|
||||
|
||||
return data
|
||||
|
||||
def to_public_dict(self) -> Dict[str, Any]:
|
||||
"""Convert to public dictionary (minimal user info)."""
|
||||
return {
|
||||
"id": self.id,
|
||||
"username": self.username,
|
||||
"display_name": self.get_display_name(),
|
||||
"avatar_url": self.avatar_url,
|
||||
"is_verified": self.is_verified,
|
||||
"subscription_level": self.subscription_level,
|
||||
"created_at": self.created_at.isoformat() if self.created_at else None,
|
||||
}
|
||||
|
||||
|
||||
class UserSession(BaseModel):
|
||||
"""User session model for tracking active sessions."""
|
||||
|
||||
__tablename__ = "user_sessions"
|
||||
|
||||
user_id = Column(Integer, nullable=False, index=True)
|
||||
session_token = Column(String(255), unique=True, nullable=False, index=True)
|
||||
refresh_token = Column(String(255), unique=True, nullable=True, index=True)
|
||||
|
||||
# Session metadata
|
||||
ip_address = Column(String(45), nullable=True) # IPv6 support
|
||||
user_agent = Column(Text, nullable=True)
|
||||
device_info = Column(Text, nullable=True)
|
||||
|
||||
# Session status
|
||||
is_active = Column(Boolean, default=True, nullable=False)
|
||||
expires_at = Column(DateTime, nullable=False)
|
||||
last_activity = Column(DateTime, default=datetime.utcnow, nullable=False)
|
||||
|
||||
# Security flags
|
||||
is_suspicious = Column(Boolean, default=False, nullable=False)
|
||||
failed_attempts = Column(Integer, default=0, nullable=False)
|
||||
|
||||
__table_args__ = (
|
||||
Index('idx_sessions_user_active', 'user_id', 'is_active'),
|
||||
Index('idx_sessions_token', 'session_token'),
|
||||
Index('idx_sessions_expires', 'expires_at'),
|
||||
)
|
||||
|
||||
def is_expired(self) -> bool:
|
||||
"""Check if session is expired."""
|
||||
return datetime.utcnow() > self.expires_at
|
||||
|
||||
def is_valid(self) -> bool:
|
||||
"""Check if session is valid."""
|
||||
return self.is_active and not self.is_expired()
|
||||
|
||||
def extend_session(self, hours: int = 24) -> None:
|
||||
"""Extend session expiration."""
|
||||
from datetime import timedelta
|
||||
self.expires_at = datetime.utcnow() + timedelta(hours=hours)
|
||||
self.last_activity = datetime.utcnow()
|
||||
|
||||
def mark_suspicious(self) -> None:
|
||||
"""Mark session as suspicious."""
|
||||
self.is_suspicious = True
|
||||
self.failed_attempts += 1
|
||||
|
||||
def deactivate(self) -> None:
|
||||
"""Deactivate session."""
|
||||
self.is_active = False
|
||||
|
||||
|
||||
class UserPreferences(BaseModel):
|
||||
"""User preferences and settings."""
|
||||
|
||||
__tablename__ = "user_preferences"
|
||||
|
||||
user_id = Column(Integer, unique=True, nullable=False, index=True)
|
||||
|
||||
# UI preferences
|
||||
theme = Column(String(20), default="light", nullable=False)
|
||||
language = Column(String(10), default="en", nullable=False)
|
||||
timezone = Column(String(50), default="UTC", nullable=False)
|
||||
|
||||
# Notification preferences
|
||||
email_notifications = Column(Boolean, default=True, nullable=False)
|
||||
upload_notifications = Column(Boolean, default=True, nullable=False)
|
||||
storage_alerts = Column(Boolean, default=True, nullable=False)
|
||||
|
||||
# Privacy settings
|
||||
public_profile = Column(Boolean, default=False, nullable=False)
|
||||
show_email = Column(Boolean, default=False, nullable=False)
|
||||
allow_indexing = Column(Boolean, default=True, nullable=False)
|
||||
|
||||
# Upload preferences
|
||||
auto_optimize_images = Column(Boolean, default=True, nullable=False)
|
||||
default_privacy = Column(String(20), default="private", nullable=False)
|
||||
max_file_size_mb = Column(Integer, default=10, nullable=False)
|
||||
|
||||
# Cache and performance
|
||||
cache_thumbnails = Column(Boolean, default=True, nullable=False)
|
||||
preload_content = Column(Boolean, default=False, nullable=False)
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Convert preferences to dictionary."""
|
||||
return super().to_dict(exclude={"user_id"})
|
||||
|
||||
@classmethod
|
||||
def get_default_preferences(cls) -> Dict[str, Any]:
|
||||
"""Get default user preferences."""
|
||||
return {
|
||||
"theme": "light",
|
||||
"language": "en",
|
||||
"timezone": "UTC",
|
||||
"email_notifications": True,
|
||||
"upload_notifications": True,
|
||||
"storage_alerts": True,
|
||||
"public_profile": False,
|
||||
"show_email": False,
|
||||
"allow_indexing": True,
|
||||
"auto_optimize_images": True,
|
||||
"default_privacy": "private",
|
||||
"max_file_size_mb": 10,
|
||||
"cache_thumbnails": True,
|
||||
"preload_content": False,
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
"""MY Network - Distributed Content Replication System."""
|
||||
|
||||
from .node_service import MyNetworkNodeService
|
||||
from .sync_manager import ContentSyncManager
|
||||
from .peer_manager import PeerManager
|
||||
from .bootstrap_manager import BootstrapManager
|
||||
|
||||
__all__ = [
|
||||
'MyNetworkNodeService',
|
||||
'ContentSyncManager',
|
||||
'PeerManager',
|
||||
'BootstrapManager'
|
||||
]
|
||||
@@ -0,0 +1,312 @@
|
||||
"""Bootstrap Manager - управление bootstrap нодами и начальной конфигурацией."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Optional, Any
|
||||
from datetime import datetime
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BootstrapManager:
|
||||
"""Менеджер для работы с bootstrap конфигурацией."""
|
||||
|
||||
def __init__(self, bootstrap_path: str = "bootstrap.json"):
|
||||
self.bootstrap_path = Path(bootstrap_path)
|
||||
self.config = {}
|
||||
self.nodes_history_path = Path("nodes_history.json")
|
||||
self.nodes_history = {"successful_connections": [], "last_updated": None}
|
||||
|
||||
logger.info(f"Bootstrap Manager initialized with path: {self.bootstrap_path}")
|
||||
|
||||
async def load_bootstrap_config(self) -> Dict[str, Any]:
|
||||
"""Загрузка bootstrap конфигурации."""
|
||||
try:
|
||||
if not self.bootstrap_path.exists():
|
||||
logger.error(f"Bootstrap config not found: {self.bootstrap_path}")
|
||||
raise FileNotFoundError(f"Bootstrap config not found: {self.bootstrap_path}")
|
||||
|
||||
with open(self.bootstrap_path, 'r', encoding='utf-8') as f:
|
||||
self.config = json.load(f)
|
||||
|
||||
logger.info(f"Bootstrap config loaded: {len(self.config.get('bootstrap_nodes', []))} nodes")
|
||||
|
||||
# Загрузить историю нод
|
||||
await self._load_nodes_history()
|
||||
|
||||
return self.config
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error loading bootstrap config: {e}")
|
||||
raise
|
||||
|
||||
async def _load_nodes_history(self) -> None:
|
||||
"""Загрузка истории подключенных нод."""
|
||||
try:
|
||||
if self.nodes_history_path.exists():
|
||||
with open(self.nodes_history_path, 'r', encoding='utf-8') as f:
|
||||
self.nodes_history = json.load(f)
|
||||
|
||||
logger.info(f"Loaded nodes history: {len(self.nodes_history.get('successful_connections', []))} nodes")
|
||||
else:
|
||||
logger.info("No nodes history found, starting fresh")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error loading nodes history: {e}")
|
||||
self.nodes_history = {"successful_connections": [], "last_updated": None}
|
||||
|
||||
async def save_nodes_history(self) -> None:
|
||||
"""Сохранение истории нод."""
|
||||
try:
|
||||
self.nodes_history["last_updated"] = datetime.utcnow().isoformat()
|
||||
|
||||
with open(self.nodes_history_path, 'w', encoding='utf-8') as f:
|
||||
json.dump(self.nodes_history, f, indent=2, ensure_ascii=False)
|
||||
|
||||
logger.debug("Nodes history saved")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error saving nodes history: {e}")
|
||||
|
||||
def get_bootstrap_nodes(self) -> List[Dict[str, Any]]:
|
||||
"""Получить список bootstrap нод."""
|
||||
return self.config.get('bootstrap_nodes', [])
|
||||
|
||||
def get_network_settings(self) -> Dict[str, Any]:
|
||||
"""Получить настройки сети."""
|
||||
return self.config.get('network_settings', {})
|
||||
|
||||
def get_sync_settings(self) -> Dict[str, Any]:
|
||||
"""Получить настройки синхронизации."""
|
||||
return self.config.get('sync_settings', {})
|
||||
|
||||
def get_content_settings(self) -> Dict[str, Any]:
|
||||
"""Получить настройки контента."""
|
||||
return self.config.get('content_settings', {})
|
||||
|
||||
def get_security_settings(self) -> Dict[str, Any]:
|
||||
"""Получить настройки безопасности."""
|
||||
return self.config.get('security_settings', {})
|
||||
|
||||
def get_api_settings(self) -> Dict[str, Any]:
|
||||
"""Получить настройки API."""
|
||||
return self.config.get('api_settings', {})
|
||||
|
||||
def get_monitoring_settings(self) -> Dict[str, Any]:
|
||||
"""Получить настройки мониторинга."""
|
||||
return self.config.get('monitoring_settings', {})
|
||||
|
||||
def get_storage_settings(self) -> Dict[str, Any]:
|
||||
"""Получить настройки хранилища."""
|
||||
return self.config.get('storage_settings', {})
|
||||
|
||||
def get_consensus_settings(self) -> Dict[str, Any]:
|
||||
"""Получить настройки консенсуса."""
|
||||
return self.config.get('consensus', {})
|
||||
|
||||
def get_feature_flags(self) -> Dict[str, Any]:
|
||||
"""Получить флаги функций."""
|
||||
return self.config.get('feature_flags', {})
|
||||
|
||||
def is_feature_enabled(self, feature_name: str) -> bool:
|
||||
"""Проверить, включена ли функция."""
|
||||
return self.get_feature_flags().get(feature_name, False)
|
||||
|
||||
def get_regional_settings(self, region: str = None) -> Dict[str, Any]:
|
||||
"""Получить региональные настройки."""
|
||||
regional_settings = self.config.get('regional_settings', {})
|
||||
|
||||
if region and region in regional_settings:
|
||||
return regional_settings[region]
|
||||
|
||||
return regional_settings
|
||||
|
||||
def get_emergency_settings(self) -> Dict[str, Any]:
|
||||
"""Получить настройки экстренных ситуаций."""
|
||||
return self.config.get('emergency_settings', {})
|
||||
|
||||
def is_emergency_mode(self) -> bool:
|
||||
"""Проверить, включен ли режим экстренной ситуации."""
|
||||
return self.get_emergency_settings().get('emergency_mode', False)
|
||||
|
||||
def get_nodes_from_history(self) -> List[Dict[str, Any]]:
|
||||
"""Получить ноды из истории успешных подключений."""
|
||||
return self.nodes_history.get('successful_connections', [])
|
||||
|
||||
def add_successful_connection(self, node_info: Dict[str, Any]) -> None:
|
||||
"""Добавить информацию об успешном подключении."""
|
||||
try:
|
||||
# Обновить существующую запись или добавить новую
|
||||
existing_node = None
|
||||
for i, node in enumerate(self.nodes_history['successful_connections']):
|
||||
if node['node_id'] == node_info['node_id']:
|
||||
existing_node = i
|
||||
break
|
||||
|
||||
connection_info = {
|
||||
"node_id": node_info['node_id'],
|
||||
"address": node_info['address'],
|
||||
"last_seen": datetime.utcnow().isoformat(),
|
||||
"connection_count": node_info.get('connection_count', 1),
|
||||
"performance_score": node_info.get('performance_score', 1.0),
|
||||
"features": node_info.get('features', []),
|
||||
"region": node_info.get('region', 'unknown'),
|
||||
"metadata": node_info.get('metadata', {})
|
||||
}
|
||||
|
||||
if existing_node is not None:
|
||||
# Обновить существующую запись
|
||||
old_info = self.nodes_history['successful_connections'][existing_node]
|
||||
connection_info['connection_count'] = old_info.get('connection_count', 0) + 1
|
||||
connection_info['first_seen'] = old_info.get('first_seen', connection_info['last_seen'])
|
||||
|
||||
self.nodes_history['successful_connections'][existing_node] = connection_info
|
||||
else:
|
||||
# Добавить новую запись
|
||||
connection_info['first_seen'] = connection_info['last_seen']
|
||||
self.nodes_history['successful_connections'].append(connection_info)
|
||||
|
||||
# Ограничить историю (максимум 100 нод)
|
||||
if len(self.nodes_history['successful_connections']) > 100:
|
||||
# Сортировать по последнему подключению и оставить 100 самых свежих
|
||||
self.nodes_history['successful_connections'].sort(
|
||||
key=lambda x: x['last_seen'],
|
||||
reverse=True
|
||||
)
|
||||
self.nodes_history['successful_connections'] = \
|
||||
self.nodes_history['successful_connections'][:100]
|
||||
|
||||
logger.debug(f"Added successful connection to history: {node_info['node_id']}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error adding successful connection: {e}")
|
||||
|
||||
def remove_failed_connection(self, node_id: str) -> None:
|
||||
"""Удалить ноду из истории при неудачном подключении."""
|
||||
try:
|
||||
self.nodes_history['successful_connections'] = [
|
||||
node for node in self.nodes_history['successful_connections']
|
||||
if node['node_id'] != node_id
|
||||
]
|
||||
|
||||
logger.debug(f"Removed failed connection from history: {node_id}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error removing failed connection: {e}")
|
||||
|
||||
def get_preferred_nodes(self, max_nodes: int = 10) -> List[Dict[str, Any]]:
|
||||
"""Получить предпочтительные ноды для подключения."""
|
||||
try:
|
||||
# Комбинировать bootstrap ноды и ноды из истории
|
||||
all_nodes = []
|
||||
|
||||
# Добавить bootstrap ноды (высокий приоритет)
|
||||
for node in self.get_bootstrap_nodes():
|
||||
all_nodes.append({
|
||||
"node_id": node['id'],
|
||||
"address": node['address'],
|
||||
"priority": 100, # Высокий приоритет для bootstrap
|
||||
"features": node.get('features', []),
|
||||
"region": node.get('region', 'unknown'),
|
||||
"source": "bootstrap"
|
||||
})
|
||||
|
||||
# Добавить ноды из истории
|
||||
for node in self.get_nodes_from_history():
|
||||
# Пропустить, если уже есть в bootstrap
|
||||
if any(n['node_id'] == node['node_id'] for n in all_nodes):
|
||||
continue
|
||||
|
||||
# Рассчитать приоритет на основе performance_score и connection_count
|
||||
priority = min(90, node.get('performance_score', 0.5) * 50 +
|
||||
min(40, node.get('connection_count', 1) * 2))
|
||||
|
||||
all_nodes.append({
|
||||
"node_id": node['node_id'],
|
||||
"address": node['address'],
|
||||
"priority": priority,
|
||||
"features": node.get('features', []),
|
||||
"region": node.get('region', 'unknown'),
|
||||
"source": "history"
|
||||
})
|
||||
|
||||
# Сортировать по приоритету и взять топ
|
||||
all_nodes.sort(key=lambda x: x['priority'], reverse=True)
|
||||
|
||||
return all_nodes[:max_nodes]
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting preferred nodes: {e}")
|
||||
return []
|
||||
|
||||
def validate_config(self) -> bool:
|
||||
"""Валидация конфигурации bootstrap."""
|
||||
try:
|
||||
required_fields = ['version', 'network_id', 'bootstrap_nodes']
|
||||
|
||||
for field in required_fields:
|
||||
if field not in self.config:
|
||||
logger.error(f"Missing required field: {field}")
|
||||
return False
|
||||
|
||||
# Проверить bootstrap ноды
|
||||
bootstrap_nodes = self.config.get('bootstrap_nodes', [])
|
||||
if not bootstrap_nodes:
|
||||
logger.error("No bootstrap nodes configured")
|
||||
return False
|
||||
|
||||
for node in bootstrap_nodes:
|
||||
required_node_fields = ['id', 'address']
|
||||
for field in required_node_fields:
|
||||
if field not in node:
|
||||
logger.error(f"Bootstrap node missing field: {field}")
|
||||
return False
|
||||
|
||||
logger.info("Bootstrap configuration validated successfully")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error validating config: {e}")
|
||||
return False
|
||||
|
||||
def get_config_checksum(self) -> str:
|
||||
"""Получить чек-сумму конфигурации."""
|
||||
return self.config.get('checksum', '')
|
||||
|
||||
def verify_config_signature(self) -> bool:
|
||||
"""Проверить подпись конфигурации."""
|
||||
# Заглушка для проверки подписи
|
||||
# В реальной реализации здесь была бы криптографическая проверка
|
||||
signature = self.config.get('signature', '')
|
||||
return bool(signature)
|
||||
|
||||
async def update_bootstrap_config(self, new_config: Dict[str, Any]) -> bool:
|
||||
"""Обновление bootstrap конфигурации."""
|
||||
try:
|
||||
# Сохранить резервную копию
|
||||
backup_path = self.bootstrap_path.with_suffix('.backup')
|
||||
if self.bootstrap_path.exists():
|
||||
self.bootstrap_path.rename(backup_path)
|
||||
|
||||
# Сохранить новую конфигурацию
|
||||
with open(self.bootstrap_path, 'w', encoding='utf-8') as f:
|
||||
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
||||
|
||||
# Перезагрузить конфигурацию
|
||||
await self.load_bootstrap_config()
|
||||
|
||||
logger.info("Bootstrap configuration updated successfully")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating bootstrap config: {e}")
|
||||
|
||||
# Восстановить из резервной копии
|
||||
try:
|
||||
if backup_path.exists():
|
||||
backup_path.rename(self.bootstrap_path)
|
||||
except:
|
||||
pass
|
||||
|
||||
return False
|
||||
@@ -0,0 +1,386 @@
|
||||
"""MY Network Node Service - основной сервис ноды."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Dict, List, Optional, Set, Any
|
||||
from pathlib import Path
|
||||
|
||||
from app.core.database_compatible import get_async_session
|
||||
from app.core.models.content_compatible import Content
|
||||
from app.core.cache import cache
|
||||
from .bootstrap_manager import BootstrapManager
|
||||
from .peer_manager import PeerManager
|
||||
from .sync_manager import ContentSyncManager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MyNetworkNodeService:
|
||||
"""Основной сервис ноды MY Network."""
|
||||
|
||||
def __init__(self, node_id: str = None, storage_path: str = "./storage/my-network"):
|
||||
self.node_id = node_id or self._generate_node_id()
|
||||
self.storage_path = Path(storage_path)
|
||||
self.storage_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Инициализация менеджеров
|
||||
self.bootstrap_manager = BootstrapManager()
|
||||
self.peer_manager = PeerManager(self.node_id)
|
||||
self.sync_manager = ContentSyncManager(self.node_id)
|
||||
|
||||
# Состояние ноды
|
||||
self.is_running = False
|
||||
self.start_time = None
|
||||
self.last_sync_time = None
|
||||
self.node_metrics = {
|
||||
"requests_30min": 0,
|
||||
"total_requests": 0,
|
||||
"content_synced": 0,
|
||||
"active_peers": 0,
|
||||
"storage_used_mb": 0
|
||||
}
|
||||
|
||||
# История запросов для балансировки нагрузки
|
||||
self.request_history = []
|
||||
|
||||
logger.info(f"MY Network Node Service initialized with ID: {self.node_id}")
|
||||
|
||||
def _generate_node_id(self) -> str:
|
||||
"""Генерация уникального ID ноды."""
|
||||
import uuid
|
||||
return f"node-{uuid.uuid4().hex[:8]}"
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Запуск ноды MY Network."""
|
||||
try:
|
||||
logger.info(f"Starting MY Network Node: {self.node_id}")
|
||||
|
||||
# Загрузка bootstrap конфигурации
|
||||
await self.bootstrap_manager.load_bootstrap_config()
|
||||
|
||||
# Инициализация peer manager
|
||||
await self.peer_manager.initialize()
|
||||
|
||||
# Подключение к bootstrap нодам
|
||||
await self._connect_to_bootstrap_nodes()
|
||||
|
||||
# Обнаружение других нод в сети
|
||||
await self._discover_network_nodes()
|
||||
|
||||
# Запуск синхронизации контента
|
||||
await self.sync_manager.start_sync_process()
|
||||
|
||||
# Запуск фоновых задач
|
||||
asyncio.create_task(self._background_tasks())
|
||||
|
||||
self.is_running = True
|
||||
self.start_time = datetime.utcnow()
|
||||
|
||||
logger.info(f"MY Network Node {self.node_id} started successfully")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to start MY Network Node: {e}")
|
||||
raise
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Остановка ноды MY Network."""
|
||||
try:
|
||||
logger.info(f"Stopping MY Network Node: {self.node_id}")
|
||||
|
||||
self.is_running = False
|
||||
|
||||
# Остановка синхронизации
|
||||
await self.sync_manager.stop_sync_process()
|
||||
|
||||
# Отключение от пиров
|
||||
await self.peer_manager.disconnect_all()
|
||||
|
||||
logger.info(f"MY Network Node {self.node_id} stopped")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error stopping MY Network Node: {e}")
|
||||
|
||||
async def _connect_to_bootstrap_nodes(self) -> None:
|
||||
"""Подключение к bootstrap нодам."""
|
||||
bootstrap_nodes = self.bootstrap_manager.get_bootstrap_nodes()
|
||||
|
||||
for node in bootstrap_nodes:
|
||||
try:
|
||||
# Не подключаться к самому себе
|
||||
if node["id"] == self.node_id:
|
||||
continue
|
||||
|
||||
success = await self.peer_manager.connect_to_peer(
|
||||
node["id"],
|
||||
node["address"]
|
||||
)
|
||||
|
||||
if success:
|
||||
logger.info(f"Connected to bootstrap node: {node['id']}")
|
||||
else:
|
||||
logger.warning(f"Failed to connect to bootstrap node: {node['id']}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error connecting to bootstrap node {node['id']}: {e}")
|
||||
|
||||
async def _discover_network_nodes(self) -> None:
|
||||
"""Обнаружение других нод в сети."""
|
||||
try:
|
||||
# Запрос списка нод у подключенных пиров
|
||||
connected_peers = self.peer_manager.get_connected_peers()
|
||||
|
||||
for peer_id in connected_peers:
|
||||
try:
|
||||
nodes_list = await self.peer_manager.request_nodes_list(peer_id)
|
||||
|
||||
for node_info in nodes_list:
|
||||
# Пропустить себя
|
||||
if node_info["id"] == self.node_id:
|
||||
continue
|
||||
|
||||
# Попытаться подключиться к новой ноде
|
||||
if not self.peer_manager.is_connected(node_info["id"]):
|
||||
await self.peer_manager.connect_to_peer(
|
||||
node_info["id"],
|
||||
node_info["address"]
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error discovering nodes from peer {peer_id}: {e}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in network discovery: {e}")
|
||||
|
||||
async def _background_tasks(self) -> None:
|
||||
"""Фоновые задачи ноды."""
|
||||
while self.is_running:
|
||||
try:
|
||||
# Обновление метрик
|
||||
await self._update_metrics()
|
||||
|
||||
# Очистка истории запросов (оставляем только за последние 30 минут)
|
||||
await self._cleanup_request_history()
|
||||
|
||||
# Проверка состояния пиров
|
||||
await self.peer_manager.check_peers_health()
|
||||
|
||||
# Периодическая синхронизация
|
||||
if self._should_sync():
|
||||
await self.sync_manager.sync_with_network()
|
||||
self.last_sync_time = datetime.utcnow()
|
||||
|
||||
# Обновление кэша статистики
|
||||
await self._update_cache_stats()
|
||||
|
||||
await asyncio.sleep(30) # Проверка каждые 30 секунд
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in background tasks: {e}")
|
||||
await asyncio.sleep(60) # Увеличиваем интервал при ошибке
|
||||
|
||||
async def _update_metrics(self) -> None:
|
||||
"""Обновление метрик ноды."""
|
||||
try:
|
||||
# Подсчет запросов за последние 30 минут
|
||||
cutoff_time = datetime.utcnow() - timedelta(minutes=30)
|
||||
recent_requests = [
|
||||
req for req in self.request_history
|
||||
if req["timestamp"] > cutoff_time
|
||||
]
|
||||
|
||||
self.node_metrics.update({
|
||||
"requests_30min": len(recent_requests),
|
||||
"active_peers": len(self.peer_manager.get_connected_peers()),
|
||||
"storage_used_mb": await self._calculate_storage_usage(),
|
||||
"uptime_hours": self._get_uptime_hours()
|
||||
})
|
||||
|
||||
# Сохранение в кэш для быстрого доступа
|
||||
await cache.set(
|
||||
f"my_network:node:{self.node_id}:metrics",
|
||||
self.node_metrics,
|
||||
ttl=60
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating metrics: {e}")
|
||||
|
||||
async def _cleanup_request_history(self) -> None:
|
||||
"""Очистка истории запросов."""
|
||||
cutoff_time = datetime.utcnow() - timedelta(minutes=30)
|
||||
self.request_history = [
|
||||
req for req in self.request_history
|
||||
if req["timestamp"] > cutoff_time
|
||||
]
|
||||
|
||||
def _should_sync(self) -> bool:
|
||||
"""Проверка, нужно ли запускать синхронизацию."""
|
||||
if not self.last_sync_time:
|
||||
return True
|
||||
|
||||
# Синхронизация каждые 5 минут
|
||||
return datetime.utcnow() - self.last_sync_time > timedelta(minutes=5)
|
||||
|
||||
async def _calculate_storage_usage(self) -> int:
|
||||
"""Подсчет использования хранилища в МБ."""
|
||||
try:
|
||||
total_size = 0
|
||||
if self.storage_path.exists():
|
||||
for file_path in self.storage_path.rglob("*"):
|
||||
if file_path.is_file():
|
||||
total_size += file_path.stat().st_size
|
||||
|
||||
return total_size // (1024 * 1024) # Конвертация в МБ
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error calculating storage usage: {e}")
|
||||
return 0
|
||||
|
||||
def _get_uptime_hours(self) -> float:
|
||||
"""Получение времени работы ноды в часах."""
|
||||
if not self.start_time:
|
||||
return 0.0
|
||||
|
||||
uptime = datetime.utcnow() - self.start_time
|
||||
return uptime.total_seconds() / 3600
|
||||
|
||||
async def _update_cache_stats(self) -> None:
|
||||
"""Обновление статистики в кэше."""
|
||||
try:
|
||||
stats = {
|
||||
"node_id": self.node_id,
|
||||
"is_running": self.is_running,
|
||||
"start_time": self.start_time.isoformat() if self.start_time else None,
|
||||
"last_sync_time": self.last_sync_time.isoformat() if self.last_sync_time else None,
|
||||
"metrics": self.node_metrics,
|
||||
"connected_peers": list(self.peer_manager.get_connected_peers()),
|
||||
"sync_status": await self.sync_manager.get_sync_status()
|
||||
}
|
||||
|
||||
await cache.set(
|
||||
f"my_network:node:{self.node_id}:status",
|
||||
stats,
|
||||
ttl=30
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating cache stats: {e}")
|
||||
|
||||
def record_request(self, request_info: Dict[str, Any]) -> None:
|
||||
"""Записать информацию о запросе для метрик."""
|
||||
self.request_history.append({
|
||||
"timestamp": datetime.utcnow(),
|
||||
"endpoint": request_info.get("endpoint", "unknown"),
|
||||
"method": request_info.get("method", "GET"),
|
||||
"client_ip": request_info.get("client_ip", "unknown")
|
||||
})
|
||||
|
||||
self.node_metrics["total_requests"] += 1
|
||||
|
||||
def get_load_info(self) -> Dict[str, Any]:
|
||||
"""Получить информацию о нагрузке ноды для балансировки."""
|
||||
return {
|
||||
"node_id": self.node_id,
|
||||
"requests_30min": self.node_metrics["requests_30min"],
|
||||
"load_percentage": min(100, (self.node_metrics["requests_30min"] / 1000) * 100),
|
||||
"active_peers": self.node_metrics["active_peers"],
|
||||
"storage_used_mb": self.node_metrics["storage_used_mb"],
|
||||
"uptime_hours": self._get_uptime_hours(),
|
||||
"is_healthy": self.is_running and self.node_metrics["active_peers"] > 0
|
||||
}
|
||||
|
||||
async def replicate_content(self, content_hash: str, target_nodes: List[str] = None) -> Dict[str, Any]:
|
||||
"""Реплицировать контент на другие ноды."""
|
||||
try:
|
||||
logger.info(f"Starting replication of content: {content_hash}")
|
||||
|
||||
# Найти контент в локальной БД
|
||||
async with get_async_session() as session:
|
||||
content = await session.get(Content, {"hash": content_hash})
|
||||
|
||||
if not content:
|
||||
raise ValueError(f"Content not found: {content_hash}")
|
||||
|
||||
# Определить целевые ноды
|
||||
if not target_nodes:
|
||||
target_nodes = self.peer_manager.select_replication_nodes()
|
||||
|
||||
# Запустить репликацию через sync manager
|
||||
result = await self.sync_manager.replicate_content_to_nodes(
|
||||
content_hash,
|
||||
target_nodes
|
||||
)
|
||||
|
||||
logger.info(f"Content replication completed: {content_hash}")
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error replicating content {content_hash}: {e}")
|
||||
raise
|
||||
|
||||
async def get_network_status(self) -> Dict[str, Any]:
|
||||
"""Получить статус всей сети MY Network."""
|
||||
try:
|
||||
connected_peers = self.peer_manager.get_connected_peers()
|
||||
sync_status = await self.sync_manager.get_sync_status()
|
||||
|
||||
# Получить статус от всех подключенных пиров
|
||||
peer_statuses = {}
|
||||
for peer_id in connected_peers:
|
||||
try:
|
||||
peer_status = await self.peer_manager.request_peer_status(peer_id)
|
||||
peer_statuses[peer_id] = peer_status
|
||||
except Exception as e:
|
||||
peer_statuses[peer_id] = {"error": str(e)}
|
||||
|
||||
return {
|
||||
"local_node": {
|
||||
"id": self.node_id,
|
||||
"status": "running" if self.is_running else "stopped",
|
||||
"metrics": self.node_metrics,
|
||||
"uptime_hours": self._get_uptime_hours()
|
||||
},
|
||||
"network": {
|
||||
"connected_peers": len(connected_peers),
|
||||
"total_discovered_nodes": len(peer_statuses) + 1,
|
||||
"sync_status": sync_status,
|
||||
"last_sync": self.last_sync_time.isoformat() if self.last_sync_time else None
|
||||
},
|
||||
"peers": peer_statuses
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting network status: {e}")
|
||||
return {"error": str(e)}
|
||||
|
||||
async def get_content_sync_status(self, content_hash: str) -> Dict[str, Any]:
|
||||
"""Получить статус синхронизации конкретного контента."""
|
||||
return await self.sync_manager.get_content_sync_status(content_hash)
|
||||
|
||||
|
||||
# Глобальный экземпляр сервиса ноды
|
||||
_node_service: Optional[MyNetworkNodeService] = None
|
||||
|
||||
|
||||
def get_node_service() -> MyNetworkNodeService:
|
||||
"""Получить глобальный экземпляр сервиса ноды."""
|
||||
global _node_service
|
||||
if _node_service is None:
|
||||
_node_service = MyNetworkNodeService()
|
||||
return _node_service
|
||||
|
||||
|
||||
async def initialize_my_network() -> None:
|
||||
"""Инициализация MY Network."""
|
||||
node_service = get_node_service()
|
||||
await node_service.start()
|
||||
|
||||
|
||||
async def shutdown_my_network() -> None:
|
||||
"""Остановка MY Network."""
|
||||
global _node_service
|
||||
if _node_service:
|
||||
await _node_service.stop()
|
||||
_node_service = None
|
||||
@@ -0,0 +1,477 @@
|
||||
"""Peer Manager - управление подключениями к другим нодам."""
|
||||
|
||||
import asyncio
|
||||
import aiohttp
|
||||
import logging
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Dict, List, Set, Optional, Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class PeerConnection:
|
||||
"""Представление подключения к пиру."""
|
||||
|
||||
def __init__(self, peer_id: str, address: str):
|
||||
self.peer_id = peer_id
|
||||
self.address = address
|
||||
self.connected_at = datetime.utcnow()
|
||||
self.last_ping = None
|
||||
self.last_pong = None
|
||||
self.is_healthy = True
|
||||
self.ping_failures = 0
|
||||
self.request_count = 0
|
||||
self.features = []
|
||||
self.metadata = {}
|
||||
|
||||
@property
|
||||
def uptime(self) -> timedelta:
|
||||
"""Время подключения."""
|
||||
return datetime.utcnow() - self.connected_at
|
||||
|
||||
@property
|
||||
def ping_latency(self) -> Optional[float]:
|
||||
"""Задержка пинга в миллисекундах."""
|
||||
if self.last_ping and self.last_pong:
|
||||
return (self.last_pong - self.last_ping).total_seconds() * 1000
|
||||
return None
|
||||
|
||||
def mark_ping_sent(self):
|
||||
"""Отметить отправку пинга."""
|
||||
self.last_ping = datetime.utcnow()
|
||||
|
||||
def mark_pong_received(self):
|
||||
"""Отметить получение понга."""
|
||||
self.last_pong = datetime.utcnow()
|
||||
self.ping_failures = 0
|
||||
self.is_healthy = True
|
||||
|
||||
def mark_ping_failed(self):
|
||||
"""Отметить неудачный пинг."""
|
||||
self.ping_failures += 1
|
||||
if self.ping_failures >= 3:
|
||||
self.is_healthy = False
|
||||
|
||||
|
||||
class PeerManager:
|
||||
"""Менеджер для управления подключениями к пирам."""
|
||||
|
||||
def __init__(self, node_id: str):
|
||||
self.node_id = node_id
|
||||
self.connections: Dict[str, PeerConnection] = {}
|
||||
self.blacklisted_peers: Set[str] = set()
|
||||
self.connection_semaphore = asyncio.Semaphore(25) # Макс 25 исходящих подключений
|
||||
self.session: Optional[aiohttp.ClientSession] = None
|
||||
|
||||
logger.info(f"Peer Manager initialized for node: {node_id}")
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Инициализация менеджера пиров."""
|
||||
try:
|
||||
# Создать HTTP сессию для запросов
|
||||
timeout = aiohttp.ClientTimeout(total=30, connect=10)
|
||||
self.session = aiohttp.ClientSession(
|
||||
timeout=timeout,
|
||||
headers={'User-Agent': f'MY-Network-Node/{self.node_id}'}
|
||||
)
|
||||
|
||||
logger.info("Peer Manager initialized successfully")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error initializing Peer Manager: {e}")
|
||||
raise
|
||||
|
||||
async def cleanup(self) -> None:
|
||||
"""Очистка ресурсов."""
|
||||
if self.session:
|
||||
await self.session.close()
|
||||
self.session = None
|
||||
|
||||
self.connections.clear()
|
||||
logger.info("Peer Manager cleaned up")
|
||||
|
||||
async def connect_to_peer(self, peer_id: str, address: str) -> bool:
|
||||
"""Подключение к пиру."""
|
||||
try:
|
||||
# Проверить, что не подключаемся к себе
|
||||
if peer_id == self.node_id:
|
||||
logger.debug(f"Skipping connection to self: {peer_id}")
|
||||
return False
|
||||
|
||||
# Проверить черный список
|
||||
if peer_id in self.blacklisted_peers:
|
||||
logger.debug(f"Peer {peer_id} is blacklisted")
|
||||
return False
|
||||
|
||||
# Проверить, уже подключены ли
|
||||
if peer_id in self.connections:
|
||||
connection = self.connections[peer_id]
|
||||
if connection.is_healthy:
|
||||
logger.debug(f"Already connected to peer: {peer_id}")
|
||||
return True
|
||||
else:
|
||||
# Удалить нездоровое подключение
|
||||
del self.connections[peer_id]
|
||||
|
||||
async with self.connection_semaphore:
|
||||
logger.info(f"Connecting to peer: {peer_id} at {address}")
|
||||
|
||||
# Попытка подключения через handshake
|
||||
success = await self._perform_handshake(peer_id, address)
|
||||
|
||||
if success:
|
||||
# Создать подключение
|
||||
connection = PeerConnection(peer_id, address)
|
||||
self.connections[peer_id] = connection
|
||||
|
||||
logger.info(f"Successfully connected to peer: {peer_id}")
|
||||
return True
|
||||
else:
|
||||
logger.warning(f"Failed to connect to peer: {peer_id}")
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error connecting to peer {peer_id}: {e}")
|
||||
return False
|
||||
|
||||
async def _perform_handshake(self, peer_id: str, address: str) -> bool:
|
||||
"""Выполнить handshake с пиром."""
|
||||
try:
|
||||
if not self.session:
|
||||
return False
|
||||
|
||||
# Парсить адрес
|
||||
parsed_url = self._parse_peer_address(address)
|
||||
if not parsed_url:
|
||||
return False
|
||||
|
||||
handshake_url = f"{parsed_url}/api/my/handshake"
|
||||
|
||||
handshake_data = {
|
||||
"node_id": self.node_id,
|
||||
"protocol_version": "1.0.0",
|
||||
"features": [
|
||||
"content_sync",
|
||||
"consensus",
|
||||
"monitoring"
|
||||
],
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
async with self.session.post(handshake_url, json=handshake_data) as response:
|
||||
if response.status == 200:
|
||||
response_data = await response.json()
|
||||
|
||||
# Проверить ответ
|
||||
if (response_data.get("node_id") == peer_id and
|
||||
response_data.get("status") == "accepted"):
|
||||
|
||||
# Сохранить информацию о пире
|
||||
if peer_id in self.connections:
|
||||
self.connections[peer_id].features = response_data.get("features", [])
|
||||
self.connections[peer_id].metadata = response_data.get("metadata", {})
|
||||
|
||||
return True
|
||||
|
||||
logger.warning(f"Handshake failed with peer {peer_id}: HTTP {response.status}")
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in handshake with peer {peer_id}: {e}")
|
||||
return False
|
||||
|
||||
def _parse_peer_address(self, address: str) -> Optional[str]:
|
||||
"""Парсинг адреса пира."""
|
||||
try:
|
||||
# Поддержка форматов:
|
||||
# my://host:port
|
||||
# http://host:port
|
||||
# https://host:port
|
||||
# host:port
|
||||
|
||||
if address.startswith("my://"):
|
||||
# Конвертировать MY протокол в HTTP
|
||||
address = address.replace("my://", "http://")
|
||||
elif not address.startswith(("http://", "https://")):
|
||||
# Добавить HTTP префикс
|
||||
address = f"http://{address}"
|
||||
|
||||
parsed = urlparse(address)
|
||||
if parsed.hostname:
|
||||
return f"{parsed.scheme}://{parsed.netloc}"
|
||||
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error parsing peer address {address}: {e}")
|
||||
return None
|
||||
|
||||
async def disconnect_from_peer(self, peer_id: str) -> None:
|
||||
"""Отключение от пира."""
|
||||
try:
|
||||
if peer_id in self.connections:
|
||||
connection = self.connections[peer_id]
|
||||
|
||||
# Попытаться отправить уведомление об отключении
|
||||
try:
|
||||
await self._send_disconnect_notification(peer_id)
|
||||
except:
|
||||
pass # Игнорировать ошибки при отключении
|
||||
|
||||
# Удалить подключение
|
||||
del self.connections[peer_id]
|
||||
|
||||
logger.info(f"Disconnected from peer: {peer_id}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error disconnecting from peer {peer_id}: {e}")
|
||||
|
||||
async def _send_disconnect_notification(self, peer_id: str) -> None:
|
||||
"""Отправить уведомление об отключении."""
|
||||
try:
|
||||
if peer_id not in self.connections or not self.session:
|
||||
return
|
||||
|
||||
connection = self.connections[peer_id]
|
||||
parsed_url = self._parse_peer_address(connection.address)
|
||||
|
||||
if parsed_url:
|
||||
disconnect_url = f"{parsed_url}/api/my/disconnect"
|
||||
disconnect_data = {
|
||||
"node_id": self.node_id,
|
||||
"reason": "graceful_shutdown",
|
||||
"timestamp": datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
async with self.session.post(disconnect_url, json=disconnect_data) as response:
|
||||
if response.status == 200:
|
||||
logger.debug(f"Disconnect notification sent to {peer_id}")
|
||||
|
||||
except Exception as e:
|
||||
logger.debug(f"Error sending disconnect notification to {peer_id}: {e}")
|
||||
|
||||
async def disconnect_all(self) -> None:
|
||||
"""Отключение от всех пиров."""
|
||||
disconnect_tasks = []
|
||||
|
||||
for peer_id in list(self.connections.keys()):
|
||||
disconnect_tasks.append(self.disconnect_from_peer(peer_id))
|
||||
|
||||
if disconnect_tasks:
|
||||
await asyncio.gather(*disconnect_tasks, return_exceptions=True)
|
||||
|
||||
logger.info("Disconnected from all peers")
|
||||
|
||||
async def check_peers_health(self) -> None:
|
||||
"""Проверка здоровья всех подключений."""
|
||||
ping_tasks = []
|
||||
|
||||
for peer_id in list(self.connections.keys()):
|
||||
ping_tasks.append(self._ping_peer(peer_id))
|
||||
|
||||
if ping_tasks:
|
||||
await asyncio.gather(*ping_tasks, return_exceptions=True)
|
||||
|
||||
# Удалить нездоровые подключения
|
||||
unhealthy_peers = [
|
||||
peer_id for peer_id, conn in self.connections.items()
|
||||
if not conn.is_healthy
|
||||
]
|
||||
|
||||
for peer_id in unhealthy_peers:
|
||||
logger.warning(f"Removing unhealthy peer: {peer_id}")
|
||||
await self.disconnect_from_peer(peer_id)
|
||||
|
||||
async def _ping_peer(self, peer_id: str) -> None:
|
||||
"""Пинг пира."""
|
||||
try:
|
||||
if peer_id not in self.connections or not self.session:
|
||||
return
|
||||
|
||||
connection = self.connections[peer_id]
|
||||
parsed_url = self._parse_peer_address(connection.address)
|
||||
|
||||
if not parsed_url:
|
||||
connection.mark_ping_failed()
|
||||
return
|
||||
|
||||
ping_url = f"{parsed_url}/api/my/ping"
|
||||
connection.mark_ping_sent()
|
||||
|
||||
async with self.session.get(ping_url) as response:
|
||||
if response.status == 200:
|
||||
connection.mark_pong_received()
|
||||
logger.debug(f"Ping successful to {peer_id}, latency: {connection.ping_latency:.1f}ms")
|
||||
else:
|
||||
connection.mark_ping_failed()
|
||||
logger.debug(f"Ping failed to {peer_id}: HTTP {response.status}")
|
||||
|
||||
except Exception as e:
|
||||
if peer_id in self.connections:
|
||||
self.connections[peer_id].mark_ping_failed()
|
||||
logger.debug(f"Ping error to {peer_id}: {e}")
|
||||
|
||||
def get_connected_peers(self) -> Set[str]:
|
||||
"""Получить множество подключенных пиров."""
|
||||
return {
|
||||
peer_id for peer_id, conn in self.connections.items()
|
||||
if conn.is_healthy
|
||||
}
|
||||
|
||||
def is_connected(self, peer_id: str) -> bool:
|
||||
"""Проверить, подключены ли к пиру."""
|
||||
return (peer_id in self.connections and
|
||||
self.connections[peer_id].is_healthy)
|
||||
|
||||
def get_peer_info(self, peer_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""Получить информацию о пире."""
|
||||
if peer_id not in self.connections:
|
||||
return None
|
||||
|
||||
connection = self.connections[peer_id]
|
||||
return {
|
||||
"peer_id": peer_id,
|
||||
"address": connection.address,
|
||||
"connected_at": connection.connected_at.isoformat(),
|
||||
"uptime_seconds": connection.uptime.total_seconds(),
|
||||
"is_healthy": connection.is_healthy,
|
||||
"ping_latency_ms": connection.ping_latency,
|
||||
"ping_failures": connection.ping_failures,
|
||||
"request_count": connection.request_count,
|
||||
"features": connection.features,
|
||||
"metadata": connection.metadata
|
||||
}
|
||||
|
||||
def get_all_peers_info(self) -> Dict[str, Dict[str, Any]]:
|
||||
"""Получить информацию обо всех пирах."""
|
||||
return {
|
||||
peer_id: self.get_peer_info(peer_id)
|
||||
for peer_id in self.connections.keys()
|
||||
}
|
||||
|
||||
def select_replication_nodes(self, count: int = 3) -> List[str]:
|
||||
"""Выбрать ноды для репликации контента."""
|
||||
healthy_peers = [
|
||||
peer_id for peer_id, conn in self.connections.items()
|
||||
if conn.is_healthy
|
||||
]
|
||||
|
||||
if len(healthy_peers) <= count:
|
||||
return healthy_peers
|
||||
|
||||
# Выбрать ноды с лучшими характеристиками
|
||||
peer_scores = []
|
||||
for peer_id in healthy_peers:
|
||||
connection = self.connections[peer_id]
|
||||
|
||||
# Рассчитать оценку на основе различных факторов
|
||||
latency_score = 1.0
|
||||
if connection.ping_latency:
|
||||
latency_score = max(0.1, 1.0 - (connection.ping_latency / 1000))
|
||||
|
||||
uptime_score = min(1.0, connection.uptime.total_seconds() / 3600) # Время работы в часах
|
||||
failure_score = max(0.1, 1.0 - (connection.ping_failures / 10))
|
||||
|
||||
total_score = (latency_score * 0.4 + uptime_score * 0.3 + failure_score * 0.3)
|
||||
peer_scores.append((peer_id, total_score))
|
||||
|
||||
# Сортировать по оценке и взять топ
|
||||
peer_scores.sort(key=lambda x: x[1], reverse=True)
|
||||
return [peer_id for peer_id, _ in peer_scores[:count]]
|
||||
|
||||
async def request_nodes_list(self, peer_id: str) -> List[Dict[str, Any]]:
|
||||
"""Запросить список нод у пира."""
|
||||
try:
|
||||
if peer_id not in self.connections or not self.session:
|
||||
return []
|
||||
|
||||
connection = self.connections[peer_id]
|
||||
parsed_url = self._parse_peer_address(connection.address)
|
||||
|
||||
if not parsed_url:
|
||||
return []
|
||||
|
||||
nodes_url = f"{parsed_url}/api/my/nodes"
|
||||
|
||||
async with self.session.get(nodes_url) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
return data.get("nodes", [])
|
||||
else:
|
||||
logger.warning(f"Failed to get nodes list from {peer_id}: HTTP {response.status}")
|
||||
return []
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error requesting nodes list from {peer_id}: {e}")
|
||||
return []
|
||||
|
||||
async def request_peer_status(self, peer_id: str) -> Dict[str, Any]:
|
||||
"""Запросить статус пира."""
|
||||
try:
|
||||
if peer_id not in self.connections or not self.session:
|
||||
return {"error": "Not connected"}
|
||||
|
||||
connection = self.connections[peer_id]
|
||||
parsed_url = self._parse_peer_address(connection.address)
|
||||
|
||||
if not parsed_url:
|
||||
return {"error": "Invalid address"}
|
||||
|
||||
status_url = f"{parsed_url}/api/my/status"
|
||||
|
||||
async with self.session.get(status_url) as response:
|
||||
if response.status == 200:
|
||||
return await response.json()
|
||||
else:
|
||||
return {"error": f"HTTP {response.status}"}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error requesting peer status from {peer_id}: {e}")
|
||||
return {"error": str(e)}
|
||||
|
||||
def add_to_blacklist(self, peer_id: str, duration_hours: int = 24) -> None:
|
||||
"""Добавить пира в черный список."""
|
||||
self.blacklisted_peers.add(peer_id)
|
||||
|
||||
# Запланировать удаление из черного списка
|
||||
async def remove_from_blacklist():
|
||||
await asyncio.sleep(duration_hours * 3600)
|
||||
self.blacklisted_peers.discard(peer_id)
|
||||
logger.info(f"Removed {peer_id} from blacklist")
|
||||
|
||||
asyncio.create_task(remove_from_blacklist())
|
||||
|
||||
logger.info(f"Added {peer_id} to blacklist for {duration_hours} hours")
|
||||
|
||||
def get_connection_stats(self) -> Dict[str, Any]:
|
||||
"""Получить статистику подключений."""
|
||||
healthy_connections = sum(1 for conn in self.connections.values() if conn.is_healthy)
|
||||
|
||||
return {
|
||||
"total_connections": len(self.connections),
|
||||
"healthy_connections": healthy_connections,
|
||||
"blacklisted_peers": len(self.blacklisted_peers),
|
||||
"average_latency_ms": self._calculate_average_latency(),
|
||||
"connection_details": [
|
||||
{
|
||||
"peer_id": peer_id,
|
||||
"uptime_hours": conn.uptime.total_seconds() / 3600,
|
||||
"ping_latency_ms": conn.ping_latency,
|
||||
"is_healthy": conn.is_healthy
|
||||
}
|
||||
for peer_id, conn in self.connections.items()
|
||||
]
|
||||
}
|
||||
|
||||
def _calculate_average_latency(self) -> Optional[float]:
|
||||
"""Рассчитать среднюю задержку."""
|
||||
latencies = [
|
||||
conn.ping_latency for conn in self.connections.values()
|
||||
if conn.ping_latency is not None and conn.is_healthy
|
||||
]
|
||||
|
||||
if latencies:
|
||||
return sum(latencies) / len(latencies)
|
||||
|
||||
return None
|
||||
@@ -0,0 +1,698 @@
|
||||
"""Content Sync Manager - синхронизация контента между нодами."""
|
||||
|
||||
import asyncio
|
||||
import aiohttp
|
||||
import hashlib
|
||||
import logging
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Optional, Any, Set
|
||||
from sqlalchemy import select, and_
|
||||
|
||||
from app.core.database_compatible import get_async_session
|
||||
from app.core.models.content_compatible import Content, ContentMetadata
|
||||
from app.core.cache import cache
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ContentSyncStatus:
|
||||
"""Статус синхронизации контента."""
|
||||
|
||||
def __init__(self, content_hash: str):
|
||||
self.content_hash = content_hash
|
||||
self.sync_started = datetime.utcnow()
|
||||
self.sync_completed = None
|
||||
self.nodes_synced = set()
|
||||
self.nodes_failed = set()
|
||||
self.total_nodes = 0
|
||||
self.bytes_synced = 0
|
||||
self.status = "syncing" # syncing, completed, failed, partial
|
||||
self.error_message = None
|
||||
|
||||
@property
|
||||
def is_completed(self) -> bool:
|
||||
return self.status in ["completed", "partial"]
|
||||
|
||||
@property
|
||||
def success_rate(self) -> float:
|
||||
if self.total_nodes == 0:
|
||||
return 0.0
|
||||
return len(self.nodes_synced) / self.total_nodes
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"content_hash": self.content_hash,
|
||||
"status": self.status,
|
||||
"sync_started": self.sync_started.isoformat(),
|
||||
"sync_completed": self.sync_completed.isoformat() if self.sync_completed else None,
|
||||
"nodes_synced": list(self.nodes_synced),
|
||||
"nodes_failed": list(self.nodes_failed),
|
||||
"total_nodes": self.total_nodes,
|
||||
"bytes_synced": self.bytes_synced,
|
||||
"success_rate": self.success_rate,
|
||||
"error_message": self.error_message
|
||||
}
|
||||
|
||||
|
||||
class ContentSyncManager:
|
||||
"""Менеджер синхронизации контента между нодами MY Network."""
|
||||
|
||||
def __init__(self, node_id: str):
|
||||
self.node_id = node_id
|
||||
self.sync_queue: asyncio.Queue = asyncio.Queue()
|
||||
self.active_syncs: Dict[str, ContentSyncStatus] = {}
|
||||
self.sync_history: List[ContentSyncStatus] = []
|
||||
self.is_running = False
|
||||
self.sync_workers: List[asyncio.Task] = []
|
||||
self.session: Optional[aiohttp.ClientSession] = None
|
||||
|
||||
# Настройки синхронизации
|
||||
self.max_concurrent_syncs = 5
|
||||
self.chunk_size = 1024 * 1024 # 1MB chunks
|
||||
self.sync_timeout = 300 # 5 minutes per content
|
||||
self.retry_attempts = 3
|
||||
|
||||
logger.info(f"Content Sync Manager initialized for node: {node_id}")
|
||||
|
||||
async def start_sync_process(self) -> None:
|
||||
"""Запуск процесса синхронизации."""
|
||||
try:
|
||||
# Создать HTTP сессию
|
||||
timeout = aiohttp.ClientTimeout(total=self.sync_timeout)
|
||||
self.session = aiohttp.ClientSession(timeout=timeout)
|
||||
|
||||
# Запустить worker'ы для синхронизации
|
||||
self.is_running = True
|
||||
for i in range(self.max_concurrent_syncs):
|
||||
worker = asyncio.create_task(self._sync_worker(f"worker-{i}"))
|
||||
self.sync_workers.append(worker)
|
||||
|
||||
logger.info(f"Started {len(self.sync_workers)} sync workers")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error starting sync process: {e}")
|
||||
raise
|
||||
|
||||
async def stop_sync_process(self) -> None:
|
||||
"""Остановка процесса синхронизации."""
|
||||
try:
|
||||
self.is_running = False
|
||||
|
||||
# Остановить worker'ы
|
||||
for worker in self.sync_workers:
|
||||
worker.cancel()
|
||||
|
||||
if self.sync_workers:
|
||||
await asyncio.gather(*self.sync_workers, return_exceptions=True)
|
||||
|
||||
# Закрыть HTTP сессию
|
||||
if self.session:
|
||||
await self.session.close()
|
||||
self.session = None
|
||||
|
||||
self.sync_workers.clear()
|
||||
logger.info("Sync process stopped")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error stopping sync process: {e}")
|
||||
|
||||
async def _sync_worker(self, worker_name: str) -> None:
|
||||
"""Worker для обработки очереди синхронизации."""
|
||||
logger.info(f"Sync worker {worker_name} started")
|
||||
|
||||
while self.is_running:
|
||||
try:
|
||||
# Получить задачу из очереди
|
||||
sync_task = await asyncio.wait_for(
|
||||
self.sync_queue.get(),
|
||||
timeout=1.0
|
||||
)
|
||||
|
||||
# Обработать задачу синхронизации
|
||||
await self._process_sync_task(sync_task)
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
continue # Продолжить ожидание
|
||||
except Exception as e:
|
||||
logger.error(f"Error in sync worker {worker_name}: {e}")
|
||||
await asyncio.sleep(5) # Пауза при ошибке
|
||||
|
||||
logger.info(f"Sync worker {worker_name} stopped")
|
||||
|
||||
async def _process_sync_task(self, sync_task: Dict[str, Any]) -> None:
|
||||
"""Обработка задачи синхронизации."""
|
||||
try:
|
||||
task_type = sync_task.get("type")
|
||||
content_hash = sync_task.get("content_hash")
|
||||
target_nodes = sync_task.get("target_nodes", [])
|
||||
|
||||
if task_type == "replicate":
|
||||
await self._replicate_content(content_hash, target_nodes)
|
||||
elif task_type == "download":
|
||||
source_node = sync_task.get("source_node")
|
||||
await self._download_content(content_hash, source_node)
|
||||
elif task_type == "verify":
|
||||
await self._verify_content_integrity(content_hash)
|
||||
else:
|
||||
logger.warning(f"Unknown sync task type: {task_type}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing sync task: {e}")
|
||||
|
||||
async def replicate_content_to_nodes(self, content_hash: str, target_nodes: List[str]) -> Dict[str, Any]:
|
||||
"""Реплицировать контент на указанные ноды."""
|
||||
try:
|
||||
# Создать статус синхронизации
|
||||
sync_status = ContentSyncStatus(content_hash)
|
||||
sync_status.total_nodes = len(target_nodes)
|
||||
self.active_syncs[content_hash] = sync_status
|
||||
|
||||
# Добавить задачу в очередь
|
||||
sync_task = {
|
||||
"type": "replicate",
|
||||
"content_hash": content_hash,
|
||||
"target_nodes": target_nodes
|
||||
}
|
||||
|
||||
await self.sync_queue.put(sync_task)
|
||||
|
||||
logger.info(f"Queued replication of {content_hash} to {len(target_nodes)} nodes")
|
||||
|
||||
return {
|
||||
"status": "queued",
|
||||
"content_hash": content_hash,
|
||||
"target_nodes": target_nodes,
|
||||
"sync_id": content_hash
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error queuing content replication: {e}")
|
||||
raise
|
||||
|
||||
async def _replicate_content(self, content_hash: str, target_nodes: List[str]) -> None:
|
||||
"""Реплицировать контент на целевые ноды."""
|
||||
try:
|
||||
if content_hash not in self.active_syncs:
|
||||
logger.warning(f"No sync status found for content: {content_hash}")
|
||||
return
|
||||
|
||||
sync_status = self.active_syncs[content_hash]
|
||||
|
||||
# Получить контент из локальной БД
|
||||
content_info = await self._get_local_content_info(content_hash)
|
||||
if not content_info:
|
||||
sync_status.status = "failed"
|
||||
sync_status.error_message = "Content not found locally"
|
||||
return
|
||||
|
||||
# Реплицировать на каждую ноду
|
||||
replication_tasks = []
|
||||
for node_id in target_nodes:
|
||||
task = self._replicate_to_single_node(content_hash, node_id, content_info)
|
||||
replication_tasks.append(task)
|
||||
|
||||
# Ждать завершения всех репликаций
|
||||
results = await asyncio.gather(*replication_tasks, return_exceptions=True)
|
||||
|
||||
# Обработать результаты
|
||||
for i, result in enumerate(results):
|
||||
node_id = target_nodes[i]
|
||||
|
||||
if isinstance(result, Exception):
|
||||
sync_status.nodes_failed.add(node_id)
|
||||
logger.error(f"Replication to {node_id} failed: {result}")
|
||||
elif result:
|
||||
sync_status.nodes_synced.add(node_id)
|
||||
sync_status.bytes_synced += content_info.get("file_size", 0)
|
||||
logger.info(f"Successfully replicated to {node_id}")
|
||||
else:
|
||||
sync_status.nodes_failed.add(node_id)
|
||||
|
||||
# Завершить синхронизацию
|
||||
self._complete_sync(sync_status)
|
||||
|
||||
except Exception as e:
|
||||
if content_hash in self.active_syncs:
|
||||
self.active_syncs[content_hash].status = "failed"
|
||||
self.active_syncs[content_hash].error_message = str(e)
|
||||
logger.error(f"Error replicating content {content_hash}: {e}")
|
||||
|
||||
async def _replicate_to_single_node(self, content_hash: str, node_id: str, content_info: Dict[str, Any]) -> bool:
|
||||
"""Реплицировать контент на одну ноду."""
|
||||
try:
|
||||
if not self.session:
|
||||
return False
|
||||
|
||||
# Получить адрес ноды (через peer manager)
|
||||
from .node_service import get_node_service
|
||||
node_service = get_node_service()
|
||||
peer_info = node_service.peer_manager.get_peer_info(node_id)
|
||||
|
||||
if not peer_info:
|
||||
logger.warning(f"No peer info for node: {node_id}")
|
||||
return False
|
||||
|
||||
# Парсить адрес
|
||||
peer_address = node_service.peer_manager._parse_peer_address(peer_info["address"])
|
||||
if not peer_address:
|
||||
return False
|
||||
|
||||
# Проверить, нужна ли репликация
|
||||
check_url = f"{peer_address}/api/my/content/{content_hash}/exists"
|
||||
async with self.session.get(check_url) as response:
|
||||
if response.status == 200:
|
||||
exists_data = await response.json()
|
||||
if exists_data.get("exists", False):
|
||||
logger.debug(f"Content {content_hash} already exists on {node_id}")
|
||||
return True
|
||||
|
||||
# Начать репликацию
|
||||
replicate_url = f"{peer_address}/api/my/content/replicate"
|
||||
|
||||
# Подготовить данные для репликации
|
||||
replication_data = {
|
||||
"content_hash": content_hash,
|
||||
"metadata": content_info,
|
||||
"source_node": self.node_id
|
||||
}
|
||||
|
||||
async with self.session.post(replicate_url, json=replication_data) as response:
|
||||
if response.status == 200:
|
||||
# Передать сам файл
|
||||
success = await self._upload_content_to_node(
|
||||
content_hash,
|
||||
peer_address,
|
||||
content_info
|
||||
)
|
||||
return success
|
||||
else:
|
||||
logger.warning(f"Replication request failed to {node_id}: HTTP {response.status}")
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error replicating to node {node_id}: {e}")
|
||||
return False
|
||||
|
||||
async def _upload_content_to_node(self, content_hash: str, peer_address: str, content_info: Dict[str, Any]) -> bool:
|
||||
"""Загрузить файл контента на ноду."""
|
||||
try:
|
||||
if not self.session:
|
||||
return False
|
||||
|
||||
# Найти файл локально
|
||||
file_path = Path(content_info.get("file_path", ""))
|
||||
if not file_path.exists():
|
||||
logger.error(f"Local file not found: {file_path}")
|
||||
return False
|
||||
|
||||
upload_url = f"{peer_address}/api/my/content/{content_hash}/upload"
|
||||
|
||||
# Создать multipart upload
|
||||
with open(file_path, 'rb') as file:
|
||||
data = aiohttp.FormData()
|
||||
data.add_field('file', file, filename=content_info.get("filename", "unknown"))
|
||||
|
||||
async with self.session.post(upload_url, data=data) as response:
|
||||
if response.status == 200:
|
||||
result = await response.json()
|
||||
return result.get("success", False)
|
||||
else:
|
||||
logger.error(f"File upload failed: HTTP {response.status}")
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error uploading content to node: {e}")
|
||||
return False
|
||||
|
||||
async def _get_local_content_info(self, content_hash: str) -> Optional[Dict[str, Any]]:
|
||||
"""Получить информацию о локальном контенте."""
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
# Найти контент по хешу
|
||||
stmt = select(Content).where(Content.md5_hash == content_hash or Content.sha256_hash == content_hash)
|
||||
result = await session.execute(stmt)
|
||||
content = result.scalar_one_or_none()
|
||||
|
||||
if not content:
|
||||
return None
|
||||
|
||||
# Получить метаданные
|
||||
metadata_stmt = select(ContentMetadata).where(ContentMetadata.content_id == content.id)
|
||||
metadata_result = await session.execute(metadata_stmt)
|
||||
metadata = metadata_result.scalar_one_or_none()
|
||||
|
||||
return {
|
||||
"id": content.id,
|
||||
"hash": content_hash,
|
||||
"filename": content.filename,
|
||||
"original_filename": content.original_filename,
|
||||
"file_path": content.file_path,
|
||||
"file_size": content.file_size,
|
||||
"file_type": content.file_type,
|
||||
"mime_type": content.mime_type,
|
||||
"encrypted": content.encrypted if hasattr(content, 'encrypted') else False,
|
||||
"metadata": metadata.to_dict() if metadata else {}
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting local content info: {e}")
|
||||
return None
|
||||
|
||||
async def download_content_from_network(self, content_hash: str, source_nodes: List[str] = None) -> bool:
|
||||
"""Скачать контент из сети."""
|
||||
try:
|
||||
# Добавить задачу загрузки в очередь
|
||||
for source_node in (source_nodes or []):
|
||||
sync_task = {
|
||||
"type": "download",
|
||||
"content_hash": content_hash,
|
||||
"source_node": source_node
|
||||
}
|
||||
|
||||
await self.sync_queue.put(sync_task)
|
||||
|
||||
logger.info(f"Queued download of {content_hash} from {len(source_nodes or [])} nodes")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error queuing content download: {e}")
|
||||
return False
|
||||
|
||||
async def _download_content(self, content_hash: str, source_node: str) -> bool:
|
||||
"""Скачать контент с конкретной ноды."""
|
||||
try:
|
||||
if not self.session:
|
||||
return False
|
||||
|
||||
# Получить адрес исходной ноды
|
||||
from .node_service import get_node_service
|
||||
node_service = get_node_service()
|
||||
peer_info = node_service.peer_manager.get_peer_info(source_node)
|
||||
|
||||
if not peer_info:
|
||||
logger.warning(f"No peer info for source node: {source_node}")
|
||||
return False
|
||||
|
||||
peer_address = node_service.peer_manager._parse_peer_address(peer_info["address"])
|
||||
if not peer_address:
|
||||
return False
|
||||
|
||||
# Получить метаданные контента
|
||||
metadata_url = f"{peer_address}/api/my/content/{content_hash}/metadata"
|
||||
async with self.session.get(metadata_url) as response:
|
||||
if response.status != 200:
|
||||
logger.error(f"Failed to get content metadata: HTTP {response.status}")
|
||||
return False
|
||||
|
||||
content_metadata = await response.json()
|
||||
|
||||
# Скачать файл
|
||||
download_url = f"{peer_address}/api/my/content/{content_hash}/download"
|
||||
async with self.session.get(download_url) as response:
|
||||
if response.status != 200:
|
||||
logger.error(f"Failed to download content: HTTP {response.status}")
|
||||
return False
|
||||
|
||||
# Сохранить файл локально
|
||||
local_path = await self._save_downloaded_content(
|
||||
content_hash,
|
||||
response,
|
||||
content_metadata
|
||||
)
|
||||
|
||||
if local_path:
|
||||
# Сохранить в базу данных
|
||||
await self._save_content_to_db(content_hash, local_path, content_metadata)
|
||||
logger.info(f"Successfully downloaded content {content_hash} from {source_node}")
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error downloading content from {source_node}: {e}")
|
||||
return False
|
||||
|
||||
async def _save_downloaded_content(self, content_hash: str, response: aiohttp.ClientResponse, metadata: Dict[str, Any]) -> Optional[Path]:
|
||||
"""Сохранить скачанный контент."""
|
||||
try:
|
||||
# Создать путь для сохранения
|
||||
storage_path = Path("./storage/my-network/downloaded")
|
||||
storage_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
filename = metadata.get("filename", f"{content_hash}")
|
||||
file_path = storage_path / filename
|
||||
|
||||
# Сохранить файл
|
||||
with open(file_path, 'wb') as f:
|
||||
async for chunk in response.content.iter_chunked(self.chunk_size):
|
||||
f.write(chunk)
|
||||
|
||||
# Проверить целостность
|
||||
if await self._verify_file_integrity(file_path, content_hash):
|
||||
return file_path
|
||||
else:
|
||||
file_path.unlink() # Удалить поврежденный файл
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error saving downloaded content: {e}")
|
||||
return None
|
||||
|
||||
async def _verify_file_integrity(self, file_path: Path, expected_hash: str) -> bool:
|
||||
"""Проверить целостность файла."""
|
||||
try:
|
||||
# Вычислить хеш файла
|
||||
hash_md5 = hashlib.md5()
|
||||
hash_sha256 = hashlib.sha256()
|
||||
|
||||
with open(file_path, 'rb') as f:
|
||||
for chunk in iter(lambda: f.read(self.chunk_size), b""):
|
||||
hash_md5.update(chunk)
|
||||
hash_sha256.update(chunk)
|
||||
|
||||
file_md5 = hash_md5.hexdigest()
|
||||
file_sha256 = hash_sha256.hexdigest()
|
||||
|
||||
# Проверить соответствие
|
||||
return expected_hash in [file_md5, file_sha256]
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error verifying file integrity: {e}")
|
||||
return False
|
||||
|
||||
async def _save_content_to_db(self, content_hash: str, file_path: Path, metadata: Dict[str, Any]) -> None:
|
||||
"""Сохранить информацию о контенте в базу данных."""
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
# Создать запись контента
|
||||
content = Content(
|
||||
filename=metadata.get("filename", file_path.name),
|
||||
original_filename=metadata.get("original_filename", file_path.name),
|
||||
file_path=str(file_path),
|
||||
file_size=file_path.stat().st_size,
|
||||
file_type=metadata.get("file_type", "unknown"),
|
||||
mime_type=metadata.get("mime_type", "application/octet-stream"),
|
||||
md5_hash=content_hash if len(content_hash) == 32 else None,
|
||||
sha256_hash=content_hash if len(content_hash) == 64 else None,
|
||||
is_active=True,
|
||||
processing_status="completed"
|
||||
)
|
||||
|
||||
session.add(content)
|
||||
await session.flush()
|
||||
|
||||
# Сохранить метаданные если есть
|
||||
if metadata.get("metadata"):
|
||||
content_metadata = ContentMetadata(
|
||||
content_id=content.id,
|
||||
**metadata["metadata"]
|
||||
)
|
||||
session.add(content_metadata)
|
||||
|
||||
await session.commit()
|
||||
logger.info(f"Saved content {content_hash} to database")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error saving content to database: {e}")
|
||||
|
||||
def _complete_sync(self, sync_status: ContentSyncStatus) -> None:
|
||||
"""Завершить синхронизацию."""
|
||||
sync_status.sync_completed = datetime.utcnow()
|
||||
|
||||
# Определить итоговый статус
|
||||
if len(sync_status.nodes_synced) == sync_status.total_nodes:
|
||||
sync_status.status = "completed"
|
||||
elif len(sync_status.nodes_synced) > 0:
|
||||
sync_status.status = "partial"
|
||||
else:
|
||||
sync_status.status = "failed"
|
||||
|
||||
# Переместить в историю
|
||||
self.sync_history.append(sync_status)
|
||||
del self.active_syncs[sync_status.content_hash]
|
||||
|
||||
# Ограничить историю
|
||||
if len(self.sync_history) > 100:
|
||||
self.sync_history = self.sync_history[-100:]
|
||||
|
||||
logger.info(f"Sync completed for {sync_status.content_hash}: {sync_status.status}")
|
||||
|
||||
async def sync_with_network(self) -> Dict[str, Any]:
|
||||
"""Синхронизация с сетью - обнаружение и загрузка нового контента."""
|
||||
try:
|
||||
from .node_service import get_node_service
|
||||
node_service = get_node_service()
|
||||
|
||||
connected_peers = node_service.peer_manager.get_connected_peers()
|
||||
if not connected_peers:
|
||||
return {"status": "no_peers", "message": "No connected peers for sync"}
|
||||
|
||||
# Получить списки контента от всех пиров
|
||||
network_content = {}
|
||||
for peer_id in connected_peers:
|
||||
try:
|
||||
peer_content = await self._get_peer_content_list(peer_id)
|
||||
network_content[peer_id] = peer_content
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting content list from {peer_id}: {e}")
|
||||
|
||||
# Найти новый контент для загрузки
|
||||
new_content = await self._identify_new_content(network_content)
|
||||
|
||||
# Запустить загрузку нового контента
|
||||
download_tasks = []
|
||||
for content_hash, source_nodes in new_content.items():
|
||||
download_tasks.append(
|
||||
self.download_content_from_network(content_hash, source_nodes)
|
||||
)
|
||||
|
||||
if download_tasks:
|
||||
results = await asyncio.gather(*download_tasks, return_exceptions=True)
|
||||
successful_downloads = sum(1 for r in results if r is True)
|
||||
|
||||
return {
|
||||
"status": "sync_completed",
|
||||
"new_content_found": len(new_content),
|
||||
"downloads_queued": len(download_tasks),
|
||||
"immediate_successes": successful_downloads
|
||||
}
|
||||
else:
|
||||
return {
|
||||
"status": "up_to_date",
|
||||
"message": "No new content found"
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in network sync: {e}")
|
||||
return {"status": "error", "message": str(e)}
|
||||
|
||||
async def _get_peer_content_list(self, peer_id: str) -> List[Dict[str, Any]]:
|
||||
"""Получить список контента от пира."""
|
||||
try:
|
||||
if not self.session:
|
||||
return []
|
||||
|
||||
from .node_service import get_node_service
|
||||
node_service = get_node_service()
|
||||
peer_info = node_service.peer_manager.get_peer_info(peer_id)
|
||||
|
||||
if not peer_info:
|
||||
return []
|
||||
|
||||
peer_address = node_service.peer_manager._parse_peer_address(peer_info["address"])
|
||||
if not peer_address:
|
||||
return []
|
||||
|
||||
content_list_url = f"{peer_address}/api/my/content/list"
|
||||
async with self.session.get(content_list_url) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
return data.get("content", [])
|
||||
else:
|
||||
logger.warning(f"Failed to get content list from {peer_id}: HTTP {response.status}")
|
||||
return []
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting content list from {peer_id}: {e}")
|
||||
return []
|
||||
|
||||
async def _identify_new_content(self, network_content: Dict[str, List[Dict[str, Any]]]) -> Dict[str, List[str]]:
|
||||
"""Определить новый контент для загрузки."""
|
||||
try:
|
||||
# Получить список локального контента
|
||||
local_hashes = await self._get_local_content_hashes()
|
||||
|
||||
# Найти новый контент
|
||||
new_content = {}
|
||||
|
||||
for peer_id, content_list in network_content.items():
|
||||
for content_info in content_list:
|
||||
content_hash = content_info.get("hash")
|
||||
if not content_hash:
|
||||
continue
|
||||
|
||||
# Проверить, есть ли у нас этот контент
|
||||
if content_hash not in local_hashes:
|
||||
if content_hash not in new_content:
|
||||
new_content[content_hash] = []
|
||||
new_content[content_hash].append(peer_id)
|
||||
|
||||
return new_content
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error identifying new content: {e}")
|
||||
return {}
|
||||
|
||||
async def _get_local_content_hashes(self) -> Set[str]:
|
||||
"""Получить множество хешей локального контента."""
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
stmt = select(Content.md5_hash, Content.sha256_hash).where(Content.is_active == True)
|
||||
result = await session.execute(stmt)
|
||||
|
||||
hashes = set()
|
||||
for row in result:
|
||||
if row.md5_hash:
|
||||
hashes.add(row.md5_hash)
|
||||
if row.sha256_hash:
|
||||
hashes.add(row.sha256_hash)
|
||||
|
||||
return hashes
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting local content hashes: {e}")
|
||||
return set()
|
||||
|
||||
async def get_sync_status(self) -> Dict[str, Any]:
|
||||
"""Получить статус синхронизации."""
|
||||
return {
|
||||
"is_running": self.is_running,
|
||||
"active_syncs": len(self.active_syncs),
|
||||
"queue_size": self.sync_queue.qsize(),
|
||||
"workers_count": len(self.sync_workers),
|
||||
"recent_syncs": [
|
||||
sync.to_dict() for sync in self.sync_history[-10:]
|
||||
],
|
||||
"current_syncs": {
|
||||
content_hash: sync.to_dict()
|
||||
for content_hash, sync in self.active_syncs.items()
|
||||
}
|
||||
}
|
||||
|
||||
async def get_content_sync_status(self, content_hash: str) -> Dict[str, Any]:
|
||||
"""Получить статус синхронизации конкретного контента."""
|
||||
# Проверить активные синхронизации
|
||||
if content_hash in self.active_syncs:
|
||||
return self.active_syncs[content_hash].to_dict()
|
||||
|
||||
# Проверить историю
|
||||
for sync in reversed(self.sync_history):
|
||||
if sync.content_hash == content_hash:
|
||||
return sync.to_dict()
|
||||
|
||||
return {
|
||||
"content_hash": content_hash,
|
||||
"status": "not_found",
|
||||
"message": "No sync information found for this content"
|
||||
}
|
||||
@@ -0,0 +1,571 @@
|
||||
"""
|
||||
Comprehensive security module with encryption, JWT tokens, password hashing, and access control.
|
||||
Provides secure file encryption, token management, and authentication utilities.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import secrets
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Dict, List, Optional, Any, Union
|
||||
from uuid import UUID
|
||||
|
||||
import bcrypt
|
||||
import jwt
|
||||
from cryptography.fernet import Fernet
|
||||
from cryptography.hazmat.primitives import hashes
|
||||
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
|
||||
import base64
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
class SecurityManager:
|
||||
"""Main security manager for encryption, tokens, and authentication."""
|
||||
|
||||
def __init__(self):
|
||||
self.fernet_key = self._get_or_create_fernet_key()
|
||||
self.fernet = Fernet(self.fernet_key)
|
||||
|
||||
def _get_or_create_fernet_key(self) -> bytes:
|
||||
"""Get or create Fernet encryption key from settings."""
|
||||
if hasattr(settings, 'ENCRYPTION_KEY') and settings.ENCRYPTION_KEY:
|
||||
# Derive key from settings
|
||||
kdf = PBKDF2HMAC(
|
||||
algorithm=hashes.SHA256(),
|
||||
length=32,
|
||||
salt=settings.SECRET_KEY.encode()[:16],
|
||||
iterations=100000,
|
||||
)
|
||||
key = base64.urlsafe_b64encode(kdf.derive(settings.ENCRYPTION_KEY.encode()))
|
||||
return key
|
||||
else:
|
||||
# Generate random key (for development only)
|
||||
return Fernet.generate_key()
|
||||
|
||||
# Global security manager instance
|
||||
_security_manager = SecurityManager()
|
||||
|
||||
def hash_password(password: str) -> str:
|
||||
"""
|
||||
Hash password using bcrypt with salt.
|
||||
|
||||
Args:
|
||||
password: Plain text password
|
||||
|
||||
Returns:
|
||||
str: Hashed password
|
||||
"""
|
||||
try:
|
||||
salt = bcrypt.gensalt(rounds=12)
|
||||
hashed = bcrypt.hashpw(password.encode('utf-8'), salt)
|
||||
return hashed.decode('utf-8')
|
||||
except Exception as e:
|
||||
logger.error("Failed to hash password", error=str(e))
|
||||
raise
|
||||
|
||||
def verify_password(password: str, hashed_password: str) -> bool:
|
||||
"""
|
||||
Verify password against hash.
|
||||
|
||||
Args:
|
||||
password: Plain text password
|
||||
hashed_password: Bcrypt hashed password
|
||||
|
||||
Returns:
|
||||
bool: True if password matches
|
||||
"""
|
||||
try:
|
||||
return bcrypt.checkpw(password.encode('utf-8'), hashed_password.encode('utf-8'))
|
||||
except Exception as e:
|
||||
logger.error("Failed to verify password", error=str(e))
|
||||
return False
|
||||
|
||||
def generate_access_token(
|
||||
payload: Dict[str, Any],
|
||||
expires_in: int = 3600,
|
||||
token_type: str = "access"
|
||||
) -> str:
|
||||
"""
|
||||
Generate JWT access token.
|
||||
|
||||
Args:
|
||||
payload: Token payload data
|
||||
expires_in: Token expiration time in seconds
|
||||
token_type: Type of token (access, refresh, api)
|
||||
|
||||
Returns:
|
||||
str: JWT token
|
||||
"""
|
||||
try:
|
||||
now = datetime.utcnow()
|
||||
token_payload = {
|
||||
"iat": now,
|
||||
"exp": now + timedelta(seconds=expires_in),
|
||||
"type": token_type,
|
||||
"jti": secrets.token_urlsafe(16), # Unique token ID
|
||||
**payload
|
||||
}
|
||||
|
||||
token = jwt.encode(
|
||||
token_payload,
|
||||
settings.SECRET_KEY,
|
||||
algorithm="HS256"
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
"Access token generated",
|
||||
token_type=token_type,
|
||||
expires_in=expires_in,
|
||||
user_id=payload.get("user_id")
|
||||
)
|
||||
|
||||
return token
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to generate access token", error=str(e))
|
||||
raise
|
||||
|
||||
def verify_access_token(token: str, token_type: str = "access") -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Verify and decode JWT token.
|
||||
|
||||
Args:
|
||||
token: JWT token string
|
||||
token_type: Expected token type
|
||||
|
||||
Returns:
|
||||
Optional[Dict]: Decoded payload or None if invalid
|
||||
"""
|
||||
try:
|
||||
payload = jwt.decode(
|
||||
token,
|
||||
settings.SECRET_KEY,
|
||||
algorithms=["HS256"]
|
||||
)
|
||||
|
||||
# Verify token type
|
||||
if payload.get("type") != token_type:
|
||||
logger.warning("Token type mismatch", expected=token_type, actual=payload.get("type"))
|
||||
return None
|
||||
|
||||
# Check expiration
|
||||
if datetime.utcnow() > datetime.fromtimestamp(payload["exp"]):
|
||||
logger.warning("Token expired", exp=payload["exp"])
|
||||
return None
|
||||
|
||||
return payload
|
||||
|
||||
except jwt.ExpiredSignatureError:
|
||||
logger.warning("Token expired")
|
||||
return None
|
||||
except jwt.InvalidTokenError as e:
|
||||
logger.warning("Invalid token", error=str(e))
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.error("Failed to verify token", error=str(e))
|
||||
return None
|
||||
|
||||
def generate_refresh_token(user_id: UUID, device_id: Optional[str] = None) -> str:
|
||||
"""
|
||||
Generate long-lived refresh token.
|
||||
|
||||
Args:
|
||||
user_id: User UUID
|
||||
device_id: Optional device identifier
|
||||
|
||||
Returns:
|
||||
str: Refresh token
|
||||
"""
|
||||
payload = {
|
||||
"user_id": str(user_id),
|
||||
"device_id": device_id,
|
||||
"token_family": secrets.token_urlsafe(16) # For token rotation
|
||||
}
|
||||
|
||||
return generate_access_token(
|
||||
payload,
|
||||
expires_in=settings.REFRESH_TOKEN_EXPIRE_DAYS * 24 * 3600,
|
||||
token_type="refresh"
|
||||
)
|
||||
|
||||
def generate_api_key(
|
||||
user_id: UUID,
|
||||
permissions: List[str],
|
||||
name: str,
|
||||
expires_in: Optional[int] = None
|
||||
) -> str:
|
||||
"""
|
||||
Generate API key with specific permissions.
|
||||
|
||||
Args:
|
||||
user_id: User UUID
|
||||
permissions: List of permissions
|
||||
name: API key name
|
||||
expires_in: Optional expiration time in seconds
|
||||
|
||||
Returns:
|
||||
str: API key token
|
||||
"""
|
||||
payload = {
|
||||
"user_id": str(user_id),
|
||||
"permissions": permissions,
|
||||
"name": name,
|
||||
"key_id": secrets.token_urlsafe(16)
|
||||
}
|
||||
|
||||
expires = expires_in or (365 * 24 * 3600) # Default 1 year
|
||||
|
||||
return generate_access_token(payload, expires_in=expires, token_type="api")
|
||||
|
||||
def encrypt_data(data: Union[str, bytes], context: str = "") -> str:
|
||||
"""
|
||||
Encrypt data using Fernet symmetric encryption.
|
||||
|
||||
Args:
|
||||
data: Data to encrypt
|
||||
context: Optional context for additional security
|
||||
|
||||
Returns:
|
||||
str: Base64 encoded encrypted data
|
||||
"""
|
||||
try:
|
||||
if isinstance(data, str):
|
||||
data = data.encode('utf-8')
|
||||
|
||||
# Add context to data for additional security
|
||||
if context:
|
||||
data = f"{context}:{len(data)}:".encode('utf-8') + data
|
||||
|
||||
encrypted = _security_manager.fernet.encrypt(data)
|
||||
return base64.urlsafe_b64encode(encrypted).decode('utf-8')
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to encrypt data", error=str(e))
|
||||
raise
|
||||
|
||||
def decrypt_data(encrypted_data: str, context: str = "") -> Union[str, bytes]:
|
||||
"""
|
||||
Decrypt data using Fernet symmetric encryption.
|
||||
|
||||
Args:
|
||||
encrypted_data: Base64 encoded encrypted data
|
||||
context: Optional context for verification
|
||||
|
||||
Returns:
|
||||
Union[str, bytes]: Decrypted data
|
||||
"""
|
||||
try:
|
||||
encrypted_bytes = base64.urlsafe_b64decode(encrypted_data.encode('utf-8'))
|
||||
decrypted = _security_manager.fernet.decrypt(encrypted_bytes)
|
||||
|
||||
# Verify and remove context if provided
|
||||
if context:
|
||||
context_prefix = f"{context}:".encode('utf-8')
|
||||
if not decrypted.startswith(context_prefix):
|
||||
raise ValueError("Context mismatch during decryption")
|
||||
|
||||
# Extract length and data
|
||||
remaining = decrypted[len(context_prefix):]
|
||||
length_end = remaining.find(b':')
|
||||
if length_end == -1:
|
||||
raise ValueError("Invalid encrypted data format")
|
||||
|
||||
expected_length = int(remaining[:length_end].decode('utf-8'))
|
||||
data = remaining[length_end + 1:]
|
||||
|
||||
if len(data) != expected_length:
|
||||
raise ValueError("Data length mismatch")
|
||||
|
||||
return data
|
||||
|
||||
return decrypted
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to decrypt data", error=str(e))
|
||||
raise
|
||||
|
||||
def encrypt_file(file_data: bytes, file_id: str) -> bytes:
|
||||
"""
|
||||
Encrypt file data with file-specific context.
|
||||
|
||||
Args:
|
||||
file_data: File bytes to encrypt
|
||||
file_id: Unique file identifier
|
||||
|
||||
Returns:
|
||||
bytes: Encrypted file data
|
||||
"""
|
||||
try:
|
||||
encrypted_str = encrypt_data(file_data, context=f"file:{file_id}")
|
||||
return base64.urlsafe_b64decode(encrypted_str.encode('utf-8'))
|
||||
except Exception as e:
|
||||
logger.error("Failed to encrypt file", file_id=file_id, error=str(e))
|
||||
raise
|
||||
|
||||
def decrypt_file(encrypted_data: bytes, file_id: str) -> bytes:
|
||||
"""
|
||||
Decrypt file data with file-specific context.
|
||||
|
||||
Args:
|
||||
encrypted_data: Encrypted file bytes
|
||||
file_id: Unique file identifier
|
||||
|
||||
Returns:
|
||||
bytes: Decrypted file data
|
||||
"""
|
||||
try:
|
||||
encrypted_str = base64.urlsafe_b64encode(encrypted_data).decode('utf-8')
|
||||
decrypted = decrypt_data(encrypted_str, context=f"file:{file_id}")
|
||||
return decrypted if isinstance(decrypted, bytes) else decrypted.encode('utf-8')
|
||||
except Exception as e:
|
||||
logger.error("Failed to decrypt file", file_id=file_id, error=str(e))
|
||||
raise
|
||||
|
||||
def generate_secure_filename(original_filename: str, user_id: UUID) -> str:
|
||||
"""
|
||||
Generate secure filename to prevent path traversal and collisions.
|
||||
|
||||
Args:
|
||||
original_filename: Original filename
|
||||
user_id: User UUID
|
||||
|
||||
Returns:
|
||||
str: Secure filename
|
||||
"""
|
||||
# Extract extension
|
||||
parts = original_filename.rsplit('.', 1)
|
||||
extension = parts[1] if len(parts) > 1 else ''
|
||||
|
||||
# Generate secure base name
|
||||
timestamp = datetime.utcnow().strftime('%Y%m%d_%H%M%S')
|
||||
random_part = secrets.token_urlsafe(8)
|
||||
user_hash = hashlib.sha256(str(user_id).encode()).hexdigest()[:8]
|
||||
|
||||
secure_name = f"{timestamp}_{user_hash}_{random_part}"
|
||||
|
||||
if extension:
|
||||
# Validate extension
|
||||
allowed_extensions = {
|
||||
'txt', 'pdf', 'doc', 'docx', 'xls', 'xlsx', 'ppt', 'pptx',
|
||||
'jpg', 'jpeg', 'png', 'gif', 'bmp', 'webp', 'svg',
|
||||
'mp3', 'wav', 'flac', 'ogg', 'mp4', 'avi', 'mkv', 'webm',
|
||||
'zip', 'rar', '7z', 'tar', 'gz', 'json', 'xml', 'csv'
|
||||
}
|
||||
|
||||
clean_extension = extension.lower().strip()
|
||||
if clean_extension in allowed_extensions:
|
||||
secure_name += f".{clean_extension}"
|
||||
|
||||
return secure_name
|
||||
|
||||
def validate_file_signature(file_data: bytes, claimed_type: str) -> bool:
|
||||
"""
|
||||
Validate file signature against claimed MIME type.
|
||||
|
||||
Args:
|
||||
file_data: File bytes to validate
|
||||
claimed_type: Claimed MIME type
|
||||
|
||||
Returns:
|
||||
bool: True if signature matches type
|
||||
"""
|
||||
if len(file_data) < 8:
|
||||
return False
|
||||
|
||||
# File signatures (magic numbers)
|
||||
signatures = {
|
||||
'image/jpeg': [b'\xFF\xD8\xFF'],
|
||||
'image/png': [b'\x89PNG\r\n\x1a\n'],
|
||||
'image/gif': [b'GIF87a', b'GIF89a'],
|
||||
'image/webp': [b'RIFF', b'WEBP'],
|
||||
'application/pdf': [b'%PDF-'],
|
||||
'application/zip': [b'PK\x03\x04', b'PK\x05\x06', b'PK\x07\x08'],
|
||||
'audio/mpeg': [b'ID3', b'\xFF\xFB', b'\xFF\xF3', b'\xFF\xF2'],
|
||||
'video/mp4': [b'\x00\x00\x00\x18ftypmp4', b'\x00\x00\x00\x20ftypmp4'],
|
||||
'text/plain': [], # Text files don't have reliable signatures
|
||||
}
|
||||
|
||||
expected_sigs = signatures.get(claimed_type, [])
|
||||
|
||||
# If no signatures defined, allow (like text files)
|
||||
if not expected_sigs:
|
||||
return True
|
||||
|
||||
# Check if file starts with any expected signature
|
||||
file_start = file_data[:32] # Check first 32 bytes
|
||||
|
||||
for sig in expected_sigs:
|
||||
if file_start.startswith(sig):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def generate_csrf_token(user_id: UUID, session_id: str) -> str:
|
||||
"""
|
||||
Generate CSRF token for form protection.
|
||||
|
||||
Args:
|
||||
user_id: User UUID
|
||||
session_id: Session identifier
|
||||
|
||||
Returns:
|
||||
str: CSRF token
|
||||
"""
|
||||
timestamp = str(int(datetime.utcnow().timestamp()))
|
||||
data = f"{user_id}:{session_id}:{timestamp}"
|
||||
|
||||
signature = hmac.new(
|
||||
settings.SECRET_KEY.encode(),
|
||||
data.encode(),
|
||||
hashlib.sha256
|
||||
).hexdigest()
|
||||
|
||||
token_data = f"{data}:{signature}"
|
||||
return base64.urlsafe_b64encode(token_data.encode()).decode()
|
||||
|
||||
def verify_csrf_token(token: str, user_id: UUID, session_id: str, max_age: int = 3600) -> bool:
|
||||
"""
|
||||
Verify CSRF token.
|
||||
|
||||
Args:
|
||||
token: CSRF token to verify
|
||||
user_id: User UUID
|
||||
session_id: Session identifier
|
||||
max_age: Maximum token age in seconds
|
||||
|
||||
Returns:
|
||||
bool: True if token is valid
|
||||
"""
|
||||
try:
|
||||
token_data = base64.urlsafe_b64decode(token.encode()).decode()
|
||||
parts = token_data.split(':')
|
||||
|
||||
if len(parts) != 4:
|
||||
return False
|
||||
|
||||
token_user_id, token_session_id, timestamp, signature = parts
|
||||
|
||||
# Verify components
|
||||
if token_user_id != str(user_id) or token_session_id != session_id:
|
||||
return False
|
||||
|
||||
# Check age
|
||||
token_time = int(timestamp)
|
||||
current_time = int(datetime.utcnow().timestamp())
|
||||
|
||||
if current_time - token_time > max_age:
|
||||
return False
|
||||
|
||||
# Verify signature
|
||||
data = f"{token_user_id}:{token_session_id}:{timestamp}"
|
||||
expected_signature = hmac.new(
|
||||
settings.SECRET_KEY.encode(),
|
||||
data.encode(),
|
||||
hashlib.sha256
|
||||
).hexdigest()
|
||||
|
||||
return hmac.compare_digest(signature, expected_signature)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning("Failed to verify CSRF token", error=str(e))
|
||||
return False
|
||||
|
||||
def sanitize_input(input_data: str, max_length: int = 1000) -> str:
|
||||
"""
|
||||
Sanitize user input to prevent XSS and injection attacks.
|
||||
|
||||
Args:
|
||||
input_data: Input string to sanitize
|
||||
max_length: Maximum allowed length
|
||||
|
||||
Returns:
|
||||
str: Sanitized input
|
||||
"""
|
||||
if not input_data:
|
||||
return ""
|
||||
|
||||
# Truncate if too long
|
||||
if len(input_data) > max_length:
|
||||
input_data = input_data[:max_length]
|
||||
|
||||
# Remove/escape dangerous characters
|
||||
dangerous_chars = ['<', '>', '"', "'", '&', '\x00', '\r', '\n']
|
||||
|
||||
for char in dangerous_chars:
|
||||
if char in input_data:
|
||||
input_data = input_data.replace(char, '')
|
||||
|
||||
# Strip whitespace
|
||||
return input_data.strip()
|
||||
|
||||
def check_permission(user_permissions: List[str], required_permission: str) -> bool:
|
||||
"""
|
||||
Check if user has required permission.
|
||||
|
||||
Args:
|
||||
user_permissions: List of user permissions
|
||||
required_permission: Required permission string
|
||||
|
||||
Returns:
|
||||
bool: True if user has permission
|
||||
"""
|
||||
# Admin has all permissions
|
||||
if 'admin' in user_permissions:
|
||||
return True
|
||||
|
||||
# Check exact permission
|
||||
if required_permission in user_permissions:
|
||||
return True
|
||||
|
||||
# Check wildcard permissions
|
||||
permission_parts = required_permission.split('.')
|
||||
for i in range(len(permission_parts)):
|
||||
wildcard_perm = '.'.join(permission_parts[:i+1]) + '.*'
|
||||
if wildcard_perm in user_permissions:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def rate_limit_key(identifier: str, action: str, window: str = "default") -> str:
|
||||
"""
|
||||
Generate rate limiting key.
|
||||
|
||||
Args:
|
||||
identifier: User/IP identifier
|
||||
action: Action being rate limited
|
||||
window: Time window identifier
|
||||
|
||||
Returns:
|
||||
str: Rate limit cache key
|
||||
"""
|
||||
key_data = f"rate_limit:{action}:{window}:{identifier}"
|
||||
return hashlib.sha256(key_data.encode()).hexdigest()
|
||||
|
||||
def generate_otp(length: int = 6) -> str:
|
||||
"""
|
||||
Generate one-time password.
|
||||
|
||||
Args:
|
||||
length: Length of OTP
|
||||
|
||||
Returns:
|
||||
str: Numeric OTP
|
||||
"""
|
||||
return ''.join(secrets.choice('0123456789') for _ in range(length))
|
||||
|
||||
def constant_time_compare(a: str, b: str) -> bool:
|
||||
"""
|
||||
Constant time string comparison to prevent timing attacks.
|
||||
|
||||
Args:
|
||||
a: First string
|
||||
b: Second string
|
||||
|
||||
Returns:
|
||||
bool: True if strings are equal
|
||||
"""
|
||||
return hmac.compare_digest(a.encode('utf-8'), b.encode('utf-8'))
|
||||
+565
-36
@@ -1,45 +1,574 @@
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
"""
|
||||
Comprehensive storage management with chunked uploads, multiple backends, and security.
|
||||
Supports local storage, S3-compatible storage, and async operations with Redis caching.
|
||||
"""
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from sqlalchemy.sql import text
|
||||
import asyncio
|
||||
import hashlib
|
||||
import mimetypes
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Optional, AsyncGenerator, Any, Tuple
|
||||
from uuid import UUID, uuid4
|
||||
|
||||
from app.core._config import MYSQL_URI, MYSQL_DATABASE
|
||||
from app.core.logger import make_log
|
||||
from sqlalchemy.pool import NullPool
|
||||
import aiofiles
|
||||
import aiofiles.os
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
engine = create_engine(MYSQL_URI, poolclass=NullPool) #, echo=True)
|
||||
Session = sessionmaker(bind=engine)
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_async_session, get_cache_manager
|
||||
from app.core.logging import get_logger
|
||||
from app.core.models.content import Content, ContentChunk
|
||||
from app.core.security import encrypt_file, decrypt_file, generate_access_token
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
database_initialized = False
|
||||
while not database_initialized:
|
||||
try:
|
||||
with Session() as session:
|
||||
databases_list = session.execute(text("SHOW DATABASES;"))
|
||||
databases_list = [row[0] for row in databases_list]
|
||||
make_log("SQL", 'Database list: ' + str(databases_list), level='debug')
|
||||
assert MYSQL_DATABASE in databases_list, 'Database not found'
|
||||
database_initialized = True
|
||||
except Exception as e:
|
||||
make_log("SQL", 'MariaDB is not ready yet: ' + str(e), level='debug')
|
||||
time.sleep(1)
|
||||
class StorageBackend:
|
||||
"""Abstract base class for storage backends."""
|
||||
|
||||
async def store_chunk(self, upload_id: UUID, chunk_index: int, data: bytes) -> str:
|
||||
"""Store a file chunk and return its identifier."""
|
||||
raise NotImplementedError
|
||||
|
||||
async def retrieve_chunk(self, chunk_id: str) -> bytes:
|
||||
"""Retrieve a file chunk by its identifier."""
|
||||
raise NotImplementedError
|
||||
|
||||
async def delete_chunk(self, chunk_id: str) -> bool:
|
||||
"""Delete a file chunk."""
|
||||
raise NotImplementedError
|
||||
|
||||
async def assemble_file(self, upload_id: UUID, chunks: List[str]) -> str:
|
||||
"""Assemble chunks into final file and return file path."""
|
||||
raise NotImplementedError
|
||||
|
||||
async def delete_file(self, file_path: str) -> bool:
|
||||
"""Delete a complete file."""
|
||||
raise NotImplementedError
|
||||
|
||||
async def get_file_stream(self, file_path: str) -> AsyncGenerator[bytes, None]:
|
||||
"""Get async file stream for download."""
|
||||
raise NotImplementedError
|
||||
|
||||
engine = create_engine(f"{MYSQL_URI}/{MYSQL_DATABASE}", poolclass=NullPool)
|
||||
Session = sessionmaker(bind=engine)
|
||||
class LocalStorageBackend(StorageBackend):
|
||||
"""Local filesystem storage backend with encryption support."""
|
||||
|
||||
def __init__(self):
|
||||
self.base_path = Path(settings.STORAGE_PATH)
|
||||
self.chunks_path = self.base_path / "chunks"
|
||||
self.files_path = self.base_path / "files"
|
||||
|
||||
# Create directories if they don't exist
|
||||
self.chunks_path.mkdir(parents=True, exist_ok=True)
|
||||
self.files_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
async def store_chunk(self, upload_id: UUID, chunk_index: int, data: bytes) -> str:
|
||||
"""Store chunk to local filesystem with optional encryption."""
|
||||
try:
|
||||
chunk_id = f"{upload_id}_{chunk_index:06d}"
|
||||
chunk_path = self.chunks_path / f"{chunk_id}.chunk"
|
||||
|
||||
# Encrypt chunk if encryption is enabled
|
||||
if settings.ENCRYPT_FILES:
|
||||
data = encrypt_file(data, str(upload_id))
|
||||
|
||||
async with aiofiles.open(chunk_path, 'wb') as f:
|
||||
await f.write(data)
|
||||
|
||||
await logger.adebug(
|
||||
"Chunk stored successfully",
|
||||
upload_id=str(upload_id),
|
||||
chunk_index=chunk_index,
|
||||
chunk_size=len(data)
|
||||
)
|
||||
|
||||
return chunk_id
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Failed to store chunk",
|
||||
upload_id=str(upload_id),
|
||||
chunk_index=chunk_index,
|
||||
error=str(e)
|
||||
)
|
||||
raise
|
||||
|
||||
async def retrieve_chunk(self, chunk_id: str) -> bytes:
|
||||
"""Retrieve and optionally decrypt chunk from local filesystem."""
|
||||
try:
|
||||
chunk_path = self.chunks_path / f"{chunk_id}.chunk"
|
||||
|
||||
if not chunk_path.exists():
|
||||
raise FileNotFoundError(f"Chunk {chunk_id} not found")
|
||||
|
||||
async with aiofiles.open(chunk_path, 'rb') as f:
|
||||
data = await f.read()
|
||||
|
||||
# Decrypt chunk if encryption is enabled
|
||||
if settings.ENCRYPT_FILES:
|
||||
upload_id = chunk_id.split('_')[0]
|
||||
data = decrypt_file(data, upload_id)
|
||||
|
||||
return data
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Failed to retrieve chunk", chunk_id=chunk_id, error=str(e))
|
||||
raise
|
||||
|
||||
async def delete_chunk(self, chunk_id: str) -> bool:
|
||||
"""Delete chunk file from local filesystem."""
|
||||
try:
|
||||
chunk_path = self.chunks_path / f"{chunk_id}.chunk"
|
||||
|
||||
if chunk_path.exists():
|
||||
await aiofiles.os.remove(chunk_path)
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Failed to delete chunk", chunk_id=chunk_id, error=str(e))
|
||||
return False
|
||||
|
||||
async def assemble_file(self, upload_id: UUID, chunks: List[str]) -> str:
|
||||
"""Assemble chunks into final file."""
|
||||
try:
|
||||
file_id = str(uuid4())
|
||||
file_path = self.files_path / f"{file_id}"
|
||||
|
||||
async with aiofiles.open(file_path, 'wb') as output_file:
|
||||
for chunk_id in chunks:
|
||||
chunk_data = await self.retrieve_chunk(chunk_id)
|
||||
await output_file.write(chunk_data)
|
||||
|
||||
# Clean up chunks after assembly
|
||||
for chunk_id in chunks:
|
||||
await self.delete_chunk(chunk_id)
|
||||
|
||||
await logger.ainfo(
|
||||
"File assembled successfully",
|
||||
upload_id=str(upload_id),
|
||||
file_path=str(file_path),
|
||||
chunks_count=len(chunks)
|
||||
)
|
||||
|
||||
return str(file_path)
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Failed to assemble file",
|
||||
upload_id=str(upload_id),
|
||||
error=str(e)
|
||||
)
|
||||
raise
|
||||
|
||||
async def delete_file(self, file_path: str) -> bool:
|
||||
"""Delete file from local filesystem."""
|
||||
try:
|
||||
path = Path(file_path)
|
||||
|
||||
if path.exists() and path.is_file():
|
||||
await aiofiles.os.remove(path)
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Failed to delete file", file_path=file_path, error=str(e))
|
||||
return False
|
||||
|
||||
async def get_file_stream(self, file_path: str) -> AsyncGenerator[bytes, None]:
|
||||
"""Stream file content for download."""
|
||||
try:
|
||||
path = Path(file_path)
|
||||
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(f"File {file_path} not found")
|
||||
|
||||
async with aiofiles.open(path, 'rb') as f:
|
||||
while True:
|
||||
chunk = await f.read(65536) # 64KB chunks
|
||||
if not chunk:
|
||||
break
|
||||
yield chunk
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Failed to stream file", file_path=file_path, error=str(e))
|
||||
raise
|
||||
|
||||
class StorageManager:
|
||||
"""Main storage manager with upload session management and caching."""
|
||||
|
||||
def __init__(self):
|
||||
self.backend = LocalStorageBackend() # Can be extended to support S3, etc.
|
||||
self.cache_manager = get_cache_manager()
|
||||
|
||||
async def create_upload_session(self, content_id: UUID, total_size: int) -> Dict[str, Any]:
|
||||
"""Create new upload session with chunked upload support."""
|
||||
try:
|
||||
upload_id = uuid4()
|
||||
session_data = {
|
||||
"upload_id": str(upload_id),
|
||||
"content_id": str(content_id),
|
||||
"total_size": total_size,
|
||||
"chunk_size": settings.CHUNK_SIZE,
|
||||
"total_chunks": (total_size + settings.CHUNK_SIZE - 1) // settings.CHUNK_SIZE,
|
||||
"uploaded_chunks": [],
|
||||
"created_at": datetime.utcnow().isoformat(),
|
||||
"expires_at": (datetime.utcnow() + timedelta(hours=24)).isoformat(),
|
||||
"status": "active"
|
||||
}
|
||||
|
||||
# Store session in cache
|
||||
session_key = f"upload_session:{upload_id}"
|
||||
await self.cache_manager.set(session_key, session_data, ttl=86400) # 24 hours
|
||||
|
||||
# Store in database for persistence
|
||||
async with get_async_session() as session:
|
||||
upload_session = ContentUploadSession(
|
||||
id=upload_id,
|
||||
content_id=content_id,
|
||||
total_size=total_size,
|
||||
chunk_size=settings.CHUNK_SIZE,
|
||||
total_chunks=session_data["total_chunks"],
|
||||
expires_at=datetime.fromisoformat(session_data["expires_at"])
|
||||
)
|
||||
session.add(upload_session)
|
||||
await session.commit()
|
||||
|
||||
await logger.ainfo(
|
||||
"Upload session created",
|
||||
upload_id=str(upload_id),
|
||||
content_id=str(content_id),
|
||||
total_size=total_size
|
||||
)
|
||||
|
||||
return {
|
||||
"upload_id": str(upload_id),
|
||||
"chunk_size": settings.CHUNK_SIZE,
|
||||
"total_chunks": session_data["total_chunks"],
|
||||
"upload_url": f"/api/v1/storage/upload/{upload_id}",
|
||||
"expires_at": session_data["expires_at"]
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Failed to create upload session",
|
||||
content_id=str(content_id),
|
||||
error=str(e)
|
||||
)
|
||||
raise
|
||||
|
||||
async def upload_chunk(
|
||||
self,
|
||||
upload_id: UUID,
|
||||
chunk_index: int,
|
||||
chunk_data: bytes,
|
||||
chunk_hash: str
|
||||
) -> Dict[str, Any]:
|
||||
"""Upload and validate a file chunk."""
|
||||
try:
|
||||
# Verify chunk hash
|
||||
calculated_hash = hashlib.sha256(chunk_data).hexdigest()
|
||||
if calculated_hash != chunk_hash:
|
||||
raise ValueError("Chunk hash mismatch")
|
||||
|
||||
# Get upload session
|
||||
session_data = await self._get_upload_session(upload_id)
|
||||
if not session_data:
|
||||
raise ValueError("Upload session not found or expired")
|
||||
|
||||
# Check if chunk already uploaded
|
||||
if chunk_index in session_data.get("uploaded_chunks", []):
|
||||
return {"status": "already_uploaded", "chunk_index": chunk_index}
|
||||
|
||||
# Store chunk
|
||||
chunk_id = await self.backend.store_chunk(upload_id, chunk_index, chunk_data)
|
||||
|
||||
# Update session data
|
||||
session_data["uploaded_chunks"].append(chunk_index)
|
||||
session_data["uploaded_chunks"].sort()
|
||||
|
||||
session_key = f"upload_session:{upload_id}"
|
||||
await self.cache_manager.set(session_key, session_data, ttl=86400)
|
||||
|
||||
# Store chunk info in database
|
||||
async with get_async_session() as session:
|
||||
chunk_record = ContentChunk(
|
||||
upload_id=upload_id,
|
||||
chunk_index=chunk_index,
|
||||
chunk_id=chunk_id,
|
||||
chunk_hash=chunk_hash,
|
||||
chunk_size=len(chunk_data)
|
||||
)
|
||||
session.add(chunk_record)
|
||||
await session.commit()
|
||||
|
||||
await logger.adebug(
|
||||
"Chunk uploaded successfully",
|
||||
upload_id=str(upload_id),
|
||||
chunk_index=chunk_index,
|
||||
chunk_size=len(chunk_data)
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "uploaded",
|
||||
"chunk_index": chunk_index,
|
||||
"uploaded_chunks": len(session_data["uploaded_chunks"]),
|
||||
"total_chunks": session_data["total_chunks"]
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Failed to upload chunk",
|
||||
upload_id=str(upload_id),
|
||||
chunk_index=chunk_index,
|
||||
error=str(e)
|
||||
)
|
||||
raise
|
||||
|
||||
async def finalize_upload(self, upload_id: UUID) -> Dict[str, Any]:
|
||||
"""Finalize upload by assembling chunks into final file."""
|
||||
try:
|
||||
# Get upload session
|
||||
session_data = await self._get_upload_session(upload_id)
|
||||
if not session_data:
|
||||
raise ValueError("Upload session not found")
|
||||
|
||||
# Verify all chunks are uploaded
|
||||
uploaded_chunks = session_data.get("uploaded_chunks", [])
|
||||
total_chunks = session_data["total_chunks"]
|
||||
|
||||
if len(uploaded_chunks) != total_chunks:
|
||||
missing_chunks = set(range(total_chunks)) - set(uploaded_chunks)
|
||||
raise ValueError(f"Missing chunks: {missing_chunks}")
|
||||
|
||||
# Get chunk IDs in order
|
||||
async with get_async_session() as session:
|
||||
stmt = (
|
||||
select(ContentChunk)
|
||||
.where(ContentChunk.upload_id == upload_id)
|
||||
.order_by(ContentChunk.chunk_index)
|
||||
)
|
||||
result = await session.execute(stmt)
|
||||
chunks = result.scalars().all()
|
||||
|
||||
chunk_ids = [chunk.chunk_id for chunk in chunks]
|
||||
|
||||
# Assemble file
|
||||
file_path = await self.backend.assemble_file(upload_id, chunk_ids)
|
||||
|
||||
# Update content record
|
||||
async with get_async_session() as session:
|
||||
stmt = (
|
||||
update(Content)
|
||||
.where(Content.id == UUID(session_data["content_id"]))
|
||||
.values(
|
||||
file_path=file_path,
|
||||
status="completed",
|
||||
updated_at=datetime.utcnow()
|
||||
)
|
||||
)
|
||||
await session.execute(stmt)
|
||||
await session.commit()
|
||||
|
||||
# Clean up session
|
||||
session_key = f"upload_session:{upload_id}"
|
||||
await self.cache_manager.delete(session_key)
|
||||
|
||||
await logger.ainfo(
|
||||
"Upload finalized successfully",
|
||||
upload_id=str(upload_id),
|
||||
file_path=file_path,
|
||||
total_chunks=total_chunks
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "completed",
|
||||
"file_path": file_path,
|
||||
"content_id": session_data["content_id"]
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Failed to finalize upload",
|
||||
upload_id=str(upload_id),
|
||||
error=str(e)
|
||||
)
|
||||
raise
|
||||
|
||||
async def get_file_stream(self, file_path: str) -> AsyncGenerator[bytes, None]:
|
||||
"""Get file stream for download with caching support."""
|
||||
try:
|
||||
# Check if file is cached
|
||||
cache_key = f"file_stream:{hashlib.md5(file_path.encode()).hexdigest()}"
|
||||
|
||||
async for chunk in self.backend.get_file_stream(file_path):
|
||||
yield chunk
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Failed to get file stream", file_path=file_path, error=str(e))
|
||||
raise
|
||||
|
||||
async def delete_content_files(self, content_id: UUID) -> bool:
|
||||
"""Delete all files associated with content."""
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
# Get content
|
||||
stmt = select(Content).where(Content.id == content_id)
|
||||
result = await session.execute(stmt)
|
||||
content = result.scalar_one_or_none()
|
||||
|
||||
if not content or not content.file_path:
|
||||
return True
|
||||
|
||||
# Delete main file
|
||||
await self.backend.delete_file(content.file_path)
|
||||
|
||||
# Delete any remaining chunks
|
||||
chunk_stmt = select(ContentChunk).where(
|
||||
ContentChunk.upload_id == content_id
|
||||
)
|
||||
chunk_result = await session.execute(chunk_stmt)
|
||||
chunks = chunk_result.scalars().all()
|
||||
|
||||
for chunk in chunks:
|
||||
await self.backend.delete_chunk(chunk.chunk_id)
|
||||
|
||||
# Update content record
|
||||
update_stmt = (
|
||||
update(Content)
|
||||
.where(Content.id == content_id)
|
||||
.values(file_path=None, status="deleted")
|
||||
)
|
||||
await session.execute(update_stmt)
|
||||
await session.commit()
|
||||
|
||||
await logger.ainfo(
|
||||
"Content files deleted",
|
||||
content_id=str(content_id)
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Failed to delete content files",
|
||||
content_id=str(content_id),
|
||||
error=str(e)
|
||||
)
|
||||
return False
|
||||
|
||||
async def get_storage_stats(self) -> Dict[str, Any]:
|
||||
"""Get storage usage statistics."""
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
# Get total files and size
|
||||
from sqlalchemy import func
|
||||
stmt = select(
|
||||
func.count(Content.id).label('total_files'),
|
||||
func.sum(Content.file_size).label('total_size')
|
||||
).where(Content.status == 'completed')
|
||||
|
||||
result = await session.execute(stmt)
|
||||
stats = result.first()
|
||||
|
||||
# Get storage by type
|
||||
type_stmt = select(
|
||||
Content.content_type,
|
||||
func.count(Content.id).label('count'),
|
||||
func.sum(Content.file_size).label('size')
|
||||
).where(Content.status == 'completed').group_by(Content.content_type)
|
||||
|
||||
type_result = await session.execute(type_stmt)
|
||||
type_stats = {
|
||||
row.content_type: {
|
||||
'count': row.count,
|
||||
'size': row.size or 0
|
||||
}
|
||||
for row in type_result
|
||||
}
|
||||
|
||||
return {
|
||||
'total_files': stats.total_files or 0,
|
||||
'total_size': stats.total_size or 0,
|
||||
'by_type': type_stats,
|
||||
'updated_at': datetime.utcnow().isoformat()
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror("Failed to get storage stats", error=str(e))
|
||||
return {}
|
||||
|
||||
async def _get_upload_session(self, upload_id: UUID) -> Optional[Dict[str, Any]]:
|
||||
"""Get upload session from cache or database."""
|
||||
# Try cache first
|
||||
session_key = f"upload_session:{upload_id}"
|
||||
session_data = await self.cache_manager.get(session_key)
|
||||
|
||||
if session_data:
|
||||
# Check if session is expired
|
||||
expires_at = datetime.fromisoformat(session_data["expires_at"])
|
||||
if expires_at > datetime.utcnow():
|
||||
return session_data
|
||||
|
||||
# Fallback to database
|
||||
try:
|
||||
async with get_async_session() as session:
|
||||
stmt = (
|
||||
select(ContentUploadSession)
|
||||
.where(ContentUploadSession.id == upload_id)
|
||||
)
|
||||
result = await session.execute(stmt)
|
||||
upload_session = result.scalar_one_or_none()
|
||||
|
||||
if upload_session and upload_session.expires_at > datetime.utcnow():
|
||||
# Rebuild session data
|
||||
chunk_stmt = select(ContentChunk).where(
|
||||
ContentChunk.upload_id == upload_id
|
||||
)
|
||||
chunk_result = await session.execute(chunk_stmt)
|
||||
chunks = chunk_result.scalars().all()
|
||||
|
||||
session_data = {
|
||||
"upload_id": str(upload_session.id),
|
||||
"content_id": str(upload_session.content_id),
|
||||
"total_size": upload_session.total_size,
|
||||
"chunk_size": upload_session.chunk_size,
|
||||
"total_chunks": upload_session.total_chunks,
|
||||
"uploaded_chunks": [chunk.chunk_index for chunk in chunks],
|
||||
"created_at": upload_session.created_at.isoformat(),
|
||||
"expires_at": upload_session.expires_at.isoformat(),
|
||||
"status": "active"
|
||||
}
|
||||
|
||||
# Update cache
|
||||
await self.cache_manager.set(session_key, session_data, ttl=86400)
|
||||
return session_data
|
||||
|
||||
except Exception as e:
|
||||
await logger.aerror(
|
||||
"Failed to get upload session from database",
|
||||
upload_id=str(upload_id),
|
||||
error=str(e)
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
@contextmanager
|
||||
def db_session(auto_commit=False):
|
||||
_session = Session()
|
||||
try:
|
||||
yield _session
|
||||
if auto_commit is True:
|
||||
_session.commit()
|
||||
except BaseException as e:
|
||||
_session.rollback()
|
||||
raise e
|
||||
finally:
|
||||
_session.close()
|
||||
# Additional model for upload sessions
|
||||
from app.core.models.base import Base
|
||||
from sqlalchemy import Column, Integer, DateTime
|
||||
|
||||
class ContentUploadSession(Base):
|
||||
"""Model for tracking upload sessions."""
|
||||
__tablename__ = "content_upload_sessions"
|
||||
|
||||
content_id = Column("content_id", sa.UUID(as_uuid=True), nullable=False)
|
||||
total_size = Column(Integer, nullable=False)
|
||||
chunk_size = Column(Integer, nullable=False, default=1048576) # 1MB
|
||||
total_chunks = Column(Integer, nullable=False)
|
||||
expires_at = Column(DateTime, nullable=False)
|
||||
completed_at = Column(DateTime, nullable=True)
|
||||
@@ -0,0 +1,371 @@
|
||||
"""
|
||||
Comprehensive validation schemas using Pydantic for request/response validation.
|
||||
Provides type safety, data validation, and automatic documentation generation.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Dict, List, Optional, Any, Union
|
||||
from uuid import UUID
|
||||
from enum import Enum
|
||||
|
||||
from pydantic import BaseModel, Field, validator, root_validator
|
||||
from pydantic.networks import EmailStr, HttpUrl
|
||||
|
||||
class ContentTypeEnum(str, Enum):
|
||||
"""Supported content types."""
|
||||
AUDIO = "audio"
|
||||
VIDEO = "video"
|
||||
IMAGE = "image"
|
||||
DOCUMENT = "document"
|
||||
ARCHIVE = "archive"
|
||||
OTHER = "other"
|
||||
|
||||
class VisibilityEnum(str, Enum):
|
||||
"""Content visibility levels."""
|
||||
PUBLIC = "public"
|
||||
PRIVATE = "private"
|
||||
UNLISTED = "unlisted"
|
||||
RESTRICTED = "restricted"
|
||||
|
||||
class StatusEnum(str, Enum):
|
||||
"""Content processing status."""
|
||||
PENDING = "pending"
|
||||
PROCESSING = "processing"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
DELETED = "deleted"
|
||||
|
||||
class PermissionEnum(str, Enum):
|
||||
"""User permissions."""
|
||||
READ = "read"
|
||||
WRITE = "write"
|
||||
DELETE = "delete"
|
||||
ADMIN = "admin"
|
||||
|
||||
class BaseSchema(BaseModel):
|
||||
"""Base schema with common configuration."""
|
||||
|
||||
class Config:
|
||||
use_enum_values = True
|
||||
validate_assignment = True
|
||||
allow_population_by_field_name = True
|
||||
json_encoders = {
|
||||
datetime: lambda v: v.isoformat(),
|
||||
UUID: lambda v: str(v)
|
||||
}
|
||||
|
||||
class ContentSchema(BaseSchema):
|
||||
"""Schema for content creation."""
|
||||
title: str = Field(..., min_length=1, max_length=255, description="Content title")
|
||||
description: Optional[str] = Field(None, max_length=2000, description="Content description")
|
||||
content_type: ContentTypeEnum = Field(..., description="Type of content")
|
||||
file_size: Optional[int] = Field(None, ge=0, le=10737418240, description="File size in bytes (max 10GB)")
|
||||
visibility: VisibilityEnum = Field(VisibilityEnum.PRIVATE, description="Content visibility")
|
||||
tags: List[str] = Field(default_factory=list, max_items=20, description="Content tags")
|
||||
license_id: Optional[UUID] = Field(None, description="License ID if applicable")
|
||||
metadata: Optional[Dict[str, Any]] = Field(None, description="Additional metadata")
|
||||
|
||||
@validator('tags')
|
||||
def validate_tags(cls, v):
|
||||
"""Validate tags format and content."""
|
||||
if not v:
|
||||
return v
|
||||
|
||||
# Check each tag
|
||||
for tag in v:
|
||||
if not isinstance(tag, str):
|
||||
raise ValueError("Tags must be strings")
|
||||
if len(tag) < 1 or len(tag) > 50:
|
||||
raise ValueError("Tag length must be between 1 and 50 characters")
|
||||
if not tag.replace('-', '').replace('_', '').isalnum():
|
||||
raise ValueError("Tags can only contain alphanumeric characters, hyphens, and underscores")
|
||||
|
||||
# Remove duplicates while preserving order
|
||||
seen = set()
|
||||
unique_tags = []
|
||||
for tag in v:
|
||||
tag_lower = tag.lower()
|
||||
if tag_lower not in seen:
|
||||
seen.add(tag_lower)
|
||||
unique_tags.append(tag)
|
||||
|
||||
return unique_tags
|
||||
|
||||
@validator('metadata')
|
||||
def validate_metadata(cls, v):
|
||||
"""Validate metadata structure."""
|
||||
if not v:
|
||||
return v
|
||||
|
||||
# Check metadata size (JSON serialized)
|
||||
import json
|
||||
try:
|
||||
serialized = json.dumps(v)
|
||||
if len(serialized) > 10000: # Max 10KB of metadata
|
||||
raise ValueError("Metadata too large (max 10KB)")
|
||||
except (TypeError, ValueError) as e:
|
||||
raise ValueError(f"Invalid metadata format: {e}")
|
||||
|
||||
return v
|
||||
|
||||
class ContentUpdateSchema(BaseSchema):
|
||||
"""Schema for content updates."""
|
||||
title: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
description: Optional[str] = Field(None, max_length=2000)
|
||||
visibility: Optional[VisibilityEnum] = None
|
||||
tags: Optional[List[str]] = Field(None, max_items=20)
|
||||
license_id: Optional[UUID] = None
|
||||
status: Optional[StatusEnum] = None
|
||||
|
||||
@validator('tags')
|
||||
def validate_tags(cls, v):
|
||||
"""Validate tags if provided."""
|
||||
if v is None:
|
||||
return v
|
||||
return ContentSchema.validate_tags(v)
|
||||
|
||||
class ContentSearchSchema(BaseSchema):
|
||||
"""Schema for content search requests."""
|
||||
query: Optional[str] = Field(None, min_length=1, max_length=200, description="Search query")
|
||||
content_type: Optional[ContentTypeEnum] = None
|
||||
status: Optional[StatusEnum] = None
|
||||
tags: Optional[List[str]] = Field(None, max_items=10)
|
||||
visibility: Optional[VisibilityEnum] = None
|
||||
date_from: Optional[datetime] = None
|
||||
date_to: Optional[datetime] = None
|
||||
sort_by: Optional[str] = Field("updated_at", regex="^(created_at|updated_at|title|file_size)$")
|
||||
sort_order: Optional[str] = Field("desc", regex="^(asc|desc)$")
|
||||
page: int = Field(1, ge=1, le=1000)
|
||||
per_page: int = Field(20, ge=1, le=100)
|
||||
|
||||
@root_validator
|
||||
def validate_date_range(cls, values):
|
||||
"""Validate date range."""
|
||||
date_from = values.get('date_from')
|
||||
date_to = values.get('date_to')
|
||||
|
||||
if date_from and date_to and date_from >= date_to:
|
||||
raise ValueError("date_from must be before date_to")
|
||||
|
||||
return values
|
||||
|
||||
class UserRegistrationSchema(BaseSchema):
|
||||
"""Schema for user registration."""
|
||||
username: str = Field(..., min_length=3, max_length=50, regex="^[a-zA-Z0-9_.-]+$")
|
||||
email: EmailStr = Field(..., description="Valid email address")
|
||||
password: str = Field(..., min_length=8, max_length=128, description="Password (min 8 characters)")
|
||||
full_name: Optional[str] = Field(None, max_length=100)
|
||||
|
||||
@validator('password')
|
||||
def validate_password(cls, v):
|
||||
"""Validate password strength."""
|
||||
if len(v) < 8:
|
||||
raise ValueError("Password must be at least 8 characters long")
|
||||
|
||||
# Check for required character types
|
||||
has_upper = any(c.isupper() for c in v)
|
||||
has_lower = any(c.islower() for c in v)
|
||||
has_digit = any(c.isdigit() for c in v)
|
||||
has_special = any(c in "!@#$%^&*()_+-=[]{}|;:,.<>?" for c in v)
|
||||
|
||||
if not (has_upper and has_lower and has_digit and has_special):
|
||||
raise ValueError(
|
||||
"Password must contain at least one uppercase letter, "
|
||||
"one lowercase letter, one digit, and one special character"
|
||||
)
|
||||
|
||||
return v
|
||||
|
||||
class UserLoginSchema(BaseSchema):
|
||||
"""Schema for user login."""
|
||||
username: str = Field(..., min_length=1, max_length=50)
|
||||
password: str = Field(..., min_length=1, max_length=128)
|
||||
remember_me: bool = Field(False, description="Keep session longer")
|
||||
|
||||
class UserUpdateSchema(BaseSchema):
|
||||
"""Schema for user profile updates."""
|
||||
full_name: Optional[str] = Field(None, max_length=100)
|
||||
email: Optional[EmailStr] = None
|
||||
bio: Optional[str] = Field(None, max_length=500)
|
||||
avatar_url: Optional[HttpUrl] = None
|
||||
settings: Optional[Dict[str, Any]] = None
|
||||
|
||||
@validator('settings')
|
||||
def validate_settings(cls, v):
|
||||
"""Validate user settings."""
|
||||
if not v:
|
||||
return v
|
||||
|
||||
# Allowed settings keys
|
||||
allowed_keys = {
|
||||
'notifications', 'privacy', 'theme', 'language',
|
||||
'timezone', 'auto_save', 'quality_preference'
|
||||
}
|
||||
|
||||
for key in v.keys():
|
||||
if key not in allowed_keys:
|
||||
raise ValueError(f"Invalid settings key: {key}")
|
||||
|
||||
return v
|
||||
|
||||
class StorageUploadSchema(BaseSchema):
|
||||
"""Schema for file upload initiation."""
|
||||
filename: str = Field(..., min_length=1, max_length=255)
|
||||
file_size: int = Field(..., ge=1, le=10737418240) # Max 10GB
|
||||
content_type: str = Field(..., min_length=1, max_length=100)
|
||||
chunk_size: Optional[int] = Field(1048576, ge=65536, le=10485760) # 64KB to 10MB
|
||||
|
||||
@validator('filename')
|
||||
def validate_filename(cls, v):
|
||||
"""Validate filename format."""
|
||||
import re
|
||||
|
||||
# Check for dangerous characters
|
||||
if re.search(r'[<>:"/\\|?*\x00-\x1f]', v):
|
||||
raise ValueError("Filename contains invalid characters")
|
||||
|
||||
# Check for reserved names (Windows)
|
||||
reserved_names = {
|
||||
'CON', 'PRN', 'AUX', 'NUL',
|
||||
'COM1', 'COM2', 'COM3', 'COM4', 'COM5', 'COM6', 'COM7', 'COM8', 'COM9',
|
||||
'LPT1', 'LPT2', 'LPT3', 'LPT4', 'LPT5', 'LPT6', 'LPT7', 'LPT8', 'LPT9'
|
||||
}
|
||||
|
||||
name_part = v.split('.')[0].upper()
|
||||
if name_part in reserved_names:
|
||||
raise ValueError("Filename uses reserved name")
|
||||
|
||||
return v
|
||||
|
||||
class ChunkUploadSchema(BaseSchema):
|
||||
"""Schema for chunk upload."""
|
||||
upload_id: UUID = Field(..., description="Upload session ID")
|
||||
chunk_index: int = Field(..., ge=0, description="Chunk sequence number")
|
||||
chunk_hash: str = Field(..., min_length=64, max_length=64, description="SHA256 hash of chunk")
|
||||
is_final: bool = Field(False, description="Is this the final chunk")
|
||||
|
||||
class BlockchainTransactionSchema(BaseSchema):
|
||||
"""Schema for blockchain transactions."""
|
||||
transaction_type: str = Field(..., regex="^(transfer|mint|burn|stake|unstake)$")
|
||||
amount: Optional[int] = Field(None, ge=0, description="Amount in nanotons")
|
||||
recipient_address: Optional[str] = Field(None, min_length=48, max_length=48)
|
||||
message: Optional[str] = Field(None, max_length=500)
|
||||
|
||||
@validator('recipient_address')
|
||||
def validate_ton_address(cls, v):
|
||||
"""Validate TON address format."""
|
||||
if not v:
|
||||
return v
|
||||
|
||||
# Basic TON address validation
|
||||
import re
|
||||
if not re.match(r'^[a-zA-Z0-9_-]{48}$', v):
|
||||
raise ValueError("Invalid TON address format")
|
||||
|
||||
return v
|
||||
|
||||
class LicenseSchema(BaseSchema):
|
||||
"""Schema for license information."""
|
||||
name: str = Field(..., min_length=1, max_length=100)
|
||||
description: Optional[str] = Field(None, max_length=1000)
|
||||
url: Optional[HttpUrl] = None
|
||||
commercial_use: bool = Field(False, description="Allows commercial use")
|
||||
attribution_required: bool = Field(True, description="Requires attribution")
|
||||
share_alike: bool = Field(False, description="Requires share-alike")
|
||||
|
||||
class AccessControlSchema(BaseSchema):
|
||||
"""Schema for content access control."""
|
||||
user_id: UUID = Field(..., description="User to grant access to")
|
||||
permission: str = Field(..., regex="^(read|write|delete|admin)$")
|
||||
expires_at: Optional[datetime] = Field(None, description="Access expiration time")
|
||||
|
||||
@root_validator
|
||||
def validate_expiration(cls, values):
|
||||
"""Validate access expiration."""
|
||||
expires_at = values.get('expires_at')
|
||||
|
||||
if expires_at and expires_at <= datetime.utcnow():
|
||||
raise ValueError("Expiration time must be in the future")
|
||||
|
||||
return values
|
||||
|
||||
class ApiKeySchema(BaseSchema):
|
||||
"""Schema for API key creation."""
|
||||
name: str = Field(..., min_length=1, max_length=100, description="API key name")
|
||||
permissions: List[str] = Field(..., min_items=1, description="List of permissions")
|
||||
expires_at: Optional[datetime] = Field(None, description="Key expiration time")
|
||||
|
||||
@validator('permissions')
|
||||
def validate_permissions(cls, v):
|
||||
"""Validate permission format."""
|
||||
valid_permissions = {
|
||||
'content.read', 'content.create', 'content.update', 'content.delete',
|
||||
'storage.upload', 'storage.download', 'storage.delete',
|
||||
'user.read', 'user.update', 'admin.read', 'admin.write'
|
||||
}
|
||||
|
||||
for perm in v:
|
||||
if perm not in valid_permissions:
|
||||
raise ValueError(f"Invalid permission: {perm}")
|
||||
|
||||
return list(set(v)) # Remove duplicates
|
||||
|
||||
class WebhookSchema(BaseSchema):
|
||||
"""Schema for webhook configuration."""
|
||||
url: HttpUrl = Field(..., description="Webhook endpoint URL")
|
||||
events: List[str] = Field(..., min_items=1, description="Events to subscribe to")
|
||||
secret: Optional[str] = Field(None, min_length=16, max_length=64, description="Webhook secret")
|
||||
active: bool = Field(True, description="Whether webhook is active")
|
||||
|
||||
@validator('events')
|
||||
def validate_events(cls, v):
|
||||
"""Validate webhook events."""
|
||||
valid_events = {
|
||||
'content.created', 'content.updated', 'content.deleted',
|
||||
'user.registered', 'user.updated', 'upload.completed',
|
||||
'blockchain.transaction', 'system.error'
|
||||
}
|
||||
|
||||
for event in v:
|
||||
if event not in valid_events:
|
||||
raise ValueError(f"Invalid event: {event}")
|
||||
|
||||
return list(set(v))
|
||||
|
||||
# Response schemas
|
||||
class ContentResponseSchema(BaseSchema):
|
||||
"""Schema for content response."""
|
||||
id: UUID
|
||||
title: str
|
||||
description: Optional[str]
|
||||
content_type: ContentTypeEnum
|
||||
file_size: int
|
||||
status: StatusEnum
|
||||
visibility: VisibilityEnum
|
||||
tags: List[str]
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
user_id: UUID
|
||||
|
||||
class UserResponseSchema(BaseSchema):
|
||||
"""Schema for user response."""
|
||||
id: UUID
|
||||
username: str
|
||||
email: EmailStr
|
||||
full_name: Optional[str]
|
||||
created_at: datetime
|
||||
is_active: bool
|
||||
permissions: List[str]
|
||||
|
||||
class ErrorResponseSchema(BaseSchema):
|
||||
"""Schema for error responses."""
|
||||
error: str = Field(..., description="Error message")
|
||||
code: str = Field(..., description="Error code")
|
||||
details: Optional[Dict[str, Any]] = Field(None, description="Additional error details")
|
||||
timestamp: datetime = Field(default_factory=datetime.utcnow)
|
||||
|
||||
class SuccessResponseSchema(BaseSchema):
|
||||
"""Schema for success responses."""
|
||||
message: str = Field(..., description="Success message")
|
||||
data: Optional[Dict[str, Any]] = Field(None, description="Response data")
|
||||
timestamp: datetime = Field(default_factory=datetime.utcnow)
|
||||
Reference in new issue
Block a user