From 67af03ee9491e208e793cf690f913094d851859b Mon Sep 17 00:00:00 2001 From: vitor-aignosi Date: Mon, 27 Oct 2025 13:29:01 -0300 Subject: [PATCH] SIENTIAPDE-1314 Enhance thread safety in MLFlowRepository model caching - Introduced a reentrant lock to synchronize access to the model cache, ensuring thread safety during cache checks and updates. - Updated the cache management logic to acquire the lock when checking for existing models and when updating the cache after downloading a new model. - Reduced the maximum cached workflows in the worker configuration for improved resource management. --- .../utils/repository/model_repository.py | 35 +++++++++++-------- laborious/worker/worker.py | 3 +- 2 files changed, 21 insertions(+), 17 deletions(-) diff --git a/laborious/utils/repository/model_repository.py b/laborious/utils/repository/model_repository.py index 4572f09..c585d7d 100644 --- a/laborious/utils/repository/model_repository.py +++ b/laborious/utils/repository/model_repository.py @@ -16,6 +16,7 @@ Capabilities: import ctypes import gc +import threading import traceback from datetime import datetime, timedelta from os import environ, makedirs, path @@ -74,6 +75,7 @@ class MLFlowRepository: self.client = mlflow.tracking.MlflowClient() self.model_cache: dict[str, Any] = {} + self._cache_lock = threading.RLock() self.logger = logger """ @@ -466,28 +468,31 @@ class MLFlowRepository: model_key = f'{model_name}_{model_type}' - if model_key in self.model_cache: - cache = self.model_cache[model_key] + # Acquire lock to check cache + with self._cache_lock: + if model_key in self.model_cache: + cache = self.model_cache[model_key] - # Check if config has changed or is outdated - if self.check_cache_retention(cache, retention): - return self.handle_valid_model(model_name=model_name, cache=cache) + # Check if config has changed or is outdated + if self.check_cache_retention(cache, retention): + return self.handle_valid_model(model_name=model_name, cache=cache) + else: + # Model is outdated, delete old model files + self.handle_outdated_model(model_name=model_name, model_key=model_key) else: - # Model is outdated, delete old model files - self.handle_outdated_model(model_name=model_name, model_key=model_key) - else: - self.logger.debug( - f'Model {model_name} is not in {model_type} cache, downloading a new one' - ) + self.logger.debug( + f'Model {model_name} is not in {model_type} cache, downloading a new one' + ) - # Donwload new model + # Donwload new model (without lock to avoid blocking other threads) model, _artifact_path = self.download_model( model_name=model_name, model_type=model_type, flavor=flavor, load_wrapper=False ) - cache = {'target': model, 'timestamp': datetime.now()} - - self.model_cache[model_key] = cache + # Update cache with lock + with self._cache_lock: + cache = {'target': model, 'timestamp': datetime.now()} + self.model_cache[model_key] = cache return model diff --git a/laborious/worker/worker.py b/laborious/worker/worker.py index 781ad55..5faf004 100644 --- a/laborious/worker/worker.py +++ b/laborious/worker/worker.py @@ -152,7 +152,7 @@ async def main(): max_concurrent_workflow_tasks=50, max_concurrent_activities=50, max_concurrent_local_activities=50, - max_cached_workflows=200, + max_cached_workflows=2, workflow_task_poller_behavior=PollerBehaviorAutoscaling(), activity_task_poller_behavior=PollerBehaviorAutoscaling(), ), @@ -164,7 +164,6 @@ async def main(): # MLFlow activities.request_predict, activities.request_transform, - activities.query_to_minio, # Gates activities.input_gate, activities.mlflow_response_gate,