Merge pull request #330 from christianlouis/copilot/add-mypy-type-checking-ci
ci: enforce strict type-checking on app/utils/ module
This commit is contained in:
+2
-2
@@ -269,7 +269,7 @@ class Settings(BaseSettings):
|
|||||||
|
|
||||||
@field_validator("notification_urls", mode="before")
|
@field_validator("notification_urls", mode="before")
|
||||||
@classmethod
|
@classmethod
|
||||||
def parse_notification_urls(cls, v):
|
def parse_notification_urls(cls, v: str | list[str]) -> list[str]:
|
||||||
"""Parse notification URLs from string or list"""
|
"""Parse notification URLs from string or list"""
|
||||||
if isinstance(v, str):
|
if isinstance(v, str):
|
||||||
if "," in v:
|
if "," in v:
|
||||||
@@ -281,7 +281,7 @@ class Settings(BaseSettings):
|
|||||||
|
|
||||||
@field_validator("session_secret")
|
@field_validator("session_secret")
|
||||||
@classmethod
|
@classmethod
|
||||||
def validate_session_secret(cls, v, info):
|
def validate_session_secret(cls, v: str | None, info: object) -> str | None:
|
||||||
"""Validate that session_secret is set and has sufficient length when auth is enabled"""
|
"""Validate that session_secret is set and has sufficient length when auth is enabled"""
|
||||||
if info.data.get("auth_enabled") and not v:
|
if info.data.get("auth_enabled") and not v:
|
||||||
raise ValueError("SESSION_SECRET must be set when AUTH_ENABLED=True")
|
raise ValueError("SESSION_SECRET must be set when AUTH_ENABLED=True")
|
||||||
|
|||||||
+6
-4
@@ -2,10 +2,12 @@
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
from collections.abc import Generator
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from sqlalchemy import create_engine, exc
|
from sqlalchemy import create_engine, exc
|
||||||
from sqlalchemy.engine.url import make_url
|
from sqlalchemy.engine.url import make_url
|
||||||
from sqlalchemy.orm import declarative_base, sessionmaker
|
from sqlalchemy.orm import Session, declarative_base, sessionmaker
|
||||||
|
|
||||||
from app.config import settings
|
from app.config import settings
|
||||||
|
|
||||||
@@ -19,7 +21,7 @@ engine = create_engine(DB_URL, connect_args={"check_same_thread": False})
|
|||||||
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
||||||
|
|
||||||
|
|
||||||
def init_db():
|
def init_db() -> None:
|
||||||
"""
|
"""
|
||||||
Ensures the SQLite database file and its parent directory exist (if using sqlite).
|
Ensures the SQLite database file and its parent directory exist (if using sqlite).
|
||||||
Then runs Base.metadata.create_all(bind=engine) to initialize tables.
|
Then runs Base.metadata.create_all(bind=engine) to initialize tables.
|
||||||
@@ -55,7 +57,7 @@ def init_db():
|
|||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
||||||
def _run_schema_migrations(engine):
|
def _run_schema_migrations(engine: Any) -> None:
|
||||||
"""
|
"""
|
||||||
Apply lightweight schema migrations for columns added after the initial release.
|
Apply lightweight schema migrations for columns added after the initial release.
|
||||||
Each migration is idempotent and safe to run multiple times.
|
Each migration is idempotent and safe to run multiple times.
|
||||||
@@ -117,7 +119,7 @@ def _run_schema_migrations(engine):
|
|||||||
logger.warning(f"Skipping filehash unique index drop: {exc}")
|
logger.warning(f"Skipping filehash unique index drop: {exc}")
|
||||||
|
|
||||||
|
|
||||||
def get_db():
|
def get_db() -> Generator[Session, None, None]:
|
||||||
"""
|
"""
|
||||||
Dependency for FastAPI routes or general DB usage.
|
Dependency for FastAPI routes or general DB usage.
|
||||||
Yields a SQLAlchemy session, and closes it upon exit.
|
Yields a SQLAlchemy session, and closes it upon exit.
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ from app.models import ApplicationSettings
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def load_settings_from_db(settings_obj, db_session: Session) -> None:
|
def load_settings_from_db(settings_obj: object, db_session: Session) -> None:
|
||||||
"""
|
"""
|
||||||
Load settings from database and apply them to the settings object.
|
Load settings from database and apply them to the settings object.
|
||||||
|
|
||||||
@@ -111,7 +111,7 @@ def convert_setting_value(value: Optional[str], field_type: Any) -> Any:
|
|||||||
return str(value)
|
return str(value)
|
||||||
|
|
||||||
|
|
||||||
def reload_settings_from_db(settings_obj) -> bool:
|
def reload_settings_from_db(settings_obj: object) -> bool:
|
||||||
"""
|
"""
|
||||||
Reload settings from database.
|
Reload settings from database.
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ Module for masking sensitive information in configuration values
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
def mask_sensitive_value(value):
|
def mask_sensitive_value(value: str | None) -> str | None:
|
||||||
"""
|
"""
|
||||||
Masks sensitive values like API keys in logs and output
|
Masks sensitive values like API keys in logs and output
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ from app.config import settings
|
|||||||
from app.utils.config_validator.masking import mask_sensitive_value
|
from app.utils.config_validator.masking import mask_sensitive_value
|
||||||
|
|
||||||
|
|
||||||
def get_provider_status():
|
def get_provider_status() -> dict[str, dict[str, object]]:
|
||||||
"""
|
"""
|
||||||
Returns status information for all configured providers
|
Returns status information for all configured providers
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ logger = logging.getLogger(__name__)
|
|||||||
_PYDANTIC_INTERNALS = {"model_computed_fields", "model_config", "model_extra", "model_fields", "model_fields_set"}
|
_PYDANTIC_INTERNALS = {"model_computed_fields", "model_config", "model_extra", "model_fields", "model_fields_set"}
|
||||||
|
|
||||||
|
|
||||||
def dump_all_settings():
|
def dump_all_settings() -> None:
|
||||||
"""Log all settings values for diagnostic purposes"""
|
"""Log all settings values for diagnostic purposes"""
|
||||||
logger.info("--- DUMPING ALL SETTINGS FOR DIAGNOSTIC PURPOSES ---")
|
logger.info("--- DUMPING ALL SETTINGS FOR DIAGNOSTIC PURPOSES ---")
|
||||||
for key in dir(settings):
|
for key in dir(settings):
|
||||||
@@ -60,7 +60,7 @@ def dump_all_settings():
|
|||||||
logger.info("--- END OF SETTINGS DUMP ---")
|
logger.info("--- END OF SETTINGS DUMP ---")
|
||||||
|
|
||||||
|
|
||||||
def get_settings_for_display(show_values=False):
|
def get_settings_for_display(show_values: bool = False) -> dict[str, list[dict[str, object]]]:
|
||||||
"""
|
"""
|
||||||
Group settings into logical categories and check if they are configured.
|
Group settings into logical categories and check if they are configured.
|
||||||
Returns a dictionary with categories as keys and lists of setting items as values.
|
Returns a dictionary with categories as keys and lists of setting items as values.
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from app.config import settings
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def validate_email_config():
|
def validate_email_config() -> list[str]:
|
||||||
"""Validates email configuration settings"""
|
"""Validates email configuration settings"""
|
||||||
issues = []
|
issues = []
|
||||||
|
|
||||||
@@ -36,7 +36,7 @@ def validate_email_config():
|
|||||||
return issues
|
return issues
|
||||||
|
|
||||||
|
|
||||||
def validate_auth_config():
|
def validate_auth_config() -> list[str]:
|
||||||
"""Validates authentication configuration settings"""
|
"""Validates authentication configuration settings"""
|
||||||
issues = []
|
issues = []
|
||||||
|
|
||||||
@@ -69,7 +69,7 @@ def validate_auth_config():
|
|||||||
return issues
|
return issues
|
||||||
|
|
||||||
|
|
||||||
def validate_storage_configs():
|
def validate_storage_configs() -> dict[str, list[str]]:
|
||||||
"""Validates configuration for all storage providers"""
|
"""Validates configuration for all storage providers"""
|
||||||
issues = {}
|
issues = {}
|
||||||
|
|
||||||
@@ -178,7 +178,7 @@ def validate_storage_configs():
|
|||||||
return issues
|
return issues
|
||||||
|
|
||||||
|
|
||||||
def validate_notification_config():
|
def validate_notification_config() -> list[str]:
|
||||||
"""Check notification configuration"""
|
"""Check notification configuration"""
|
||||||
issues = []
|
issues = []
|
||||||
|
|
||||||
@@ -210,7 +210,7 @@ def validate_notification_config():
|
|||||||
return issues
|
return issues
|
||||||
|
|
||||||
|
|
||||||
def check_all_configs():
|
def check_all_configs() -> dict[str, list[str] | dict[str, list[str]]]:
|
||||||
"""Run all configuration validations and log results"""
|
"""Run all configuration validations and log results"""
|
||||||
from app.utils.config_validator.settings_display import dump_all_settings
|
from app.utils.config_validator.settings_display import dump_all_settings
|
||||||
|
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ logger = logging.getLogger(__name__)
|
|||||||
_cipher_suite = None
|
_cipher_suite = None
|
||||||
|
|
||||||
|
|
||||||
def _get_cipher_suite():
|
def _get_cipher_suite() -> object | None:
|
||||||
"""
|
"""
|
||||||
Get or create the Fernet cipher suite for encryption/decryption.
|
Get or create the Fernet cipher suite for encryption/decryption.
|
||||||
|
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
import hashlib
|
import hashlib
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
def hash_file(filepath, chunk_size=65536):
|
def hash_file(filepath: str | Path, chunk_size: int = 65536) -> str:
|
||||||
"""
|
"""
|
||||||
Returns the SHA-256 hash of the file at 'filepath'.
|
Returns the SHA-256 hash of the file at 'filepath'.
|
||||||
Reads the file in chunks to handle large files efficiently.
|
Reads the file in chunks to handle large files efficiently.
|
||||||
|
|||||||
@@ -2,13 +2,14 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
import uuid
|
import uuid
|
||||||
|
from collections.abc import Callable
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def get_unique_filename(original_path, check_exists_func=None):
|
def get_unique_filename(original_path: str, check_exists_func: Callable[[str], bool] | None = None) -> str:
|
||||||
"""
|
"""
|
||||||
Generates a unique filename by appending a timestamp or counter when a collision occurs.
|
Generates a unique filename by appending a timestamp or counter when a collision occurs.
|
||||||
|
|
||||||
@@ -69,7 +70,7 @@ def get_unique_filename(original_path, check_exists_func=None):
|
|||||||
return new_path
|
return new_path
|
||||||
|
|
||||||
|
|
||||||
def get_unique_filepath_with_counter(directory, base_filename, extension=".pdf"):
|
def get_unique_filepath_with_counter(directory: str, base_filename: str, extension: str = ".pdf") -> str:
|
||||||
"""
|
"""
|
||||||
Returns a unique filepath in the specified directory using a numeric counter suffix.
|
Returns a unique filepath in the specified directory using a numeric counter suffix.
|
||||||
If 'base_filename.pdf' exists, it will append '-0001', '-0002', etc.
|
If 'base_filename.pdf' exists, it will append '-0001', '-0002', etc.
|
||||||
@@ -120,7 +121,7 @@ def get_unique_filepath_with_counter(directory, base_filename, extension=".pdf")
|
|||||||
return candidate
|
return candidate
|
||||||
|
|
||||||
|
|
||||||
def sanitize_filename(filename):
|
def sanitize_filename(filename: str) -> str:
|
||||||
r"""
|
r"""
|
||||||
Sanitize a filename to ensure it's valid across different file systems
|
Sanitize a filename to ensure it's valid across different file systems
|
||||||
and prevent path traversal attacks.
|
and prevent path traversal attacks.
|
||||||
@@ -162,7 +163,7 @@ def sanitize_filename(filename):
|
|||||||
return sanitized
|
return sanitized
|
||||||
|
|
||||||
|
|
||||||
def extract_remote_path(file_path, base_dir, remote_base=""):
|
def extract_remote_path(file_path: str, base_dir: str, remote_base: str = "") -> str:
|
||||||
"""
|
"""
|
||||||
Extract a remote path for a file by preserving its directory structure
|
Extract a remote path for a file by preserving its directory structure
|
||||||
relative to the base directory, but with a new remote base path.
|
relative to the base directory, but with a new remote base path.
|
||||||
|
|||||||
+10
-3
@@ -16,9 +16,9 @@ class TaskLogCollector(logging.Handler):
|
|||||||
This captures all logger.info/error/warning output automatically.
|
This captures all logger.info/error/warning output automatically.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self._buffers = defaultdict(list)
|
self._buffers: defaultdict[str, list[str]] = defaultdict(list)
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
def emit(self, record: logging.LogRecord) -> None:
|
def emit(self, record: logging.LogRecord) -> None:
|
||||||
@@ -60,7 +60,14 @@ def _ensure_collector_installed() -> None:
|
|||||||
_collector_installed = True
|
_collector_installed = True
|
||||||
|
|
||||||
|
|
||||||
def log_task_progress(task_id, step_name, status, message=None, file_id=None, detail=None):
|
def log_task_progress(
|
||||||
|
task_id: str,
|
||||||
|
step_name: str,
|
||||||
|
status: str,
|
||||||
|
message: str | None = None,
|
||||||
|
file_id: int | None = None,
|
||||||
|
detail: str | None = None,
|
||||||
|
) -> None:
|
||||||
"""
|
"""
|
||||||
Logs the progress of a Celery task to the database.
|
Logs the progress of a Celery task to the database.
|
||||||
|
|
||||||
|
|||||||
@@ -194,6 +194,12 @@ disable_error_code = [
|
|||||||
"call-arg", # Dynamic call signatures in framework code
|
"call-arg", # Dynamic call signatures in framework code
|
||||||
]
|
]
|
||||||
|
|
||||||
|
# Stricter type-checking for utility modules
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "app.utils.*"
|
||||||
|
disallow_untyped_defs = true
|
||||||
|
disallow_incomplete_defs = true
|
||||||
|
|
||||||
# Coverage configuration
|
# Coverage configuration
|
||||||
[tool.coverage.run]
|
[tool.coverage.run]
|
||||||
source = ["app"]
|
source = ["app"]
|
||||||
|
|||||||
Reference in New Issue
Block a user