From 3da93fadf5a201da2afa0b63c6717adecc675555 Mon Sep 17 00:00:00 2001 From: GuoQing Zhang Date: Thu, 12 Dec 2024 15:48:35 +0800 Subject: [PATCH] fix: add celery test code --- docker/bisheng/config/config.yaml | 2 + src/backend/bisheng/settings.py | 20 +++++ src/backend/bisheng/worker.py | 76 ------------------- src/backend/bisheng/worker/__init__.py | 3 + src/backend/bisheng/worker/config.py | 9 +++ src/backend/bisheng/worker/main.py | 4 + src/backend/bisheng/worker/test/__init__.py | 0 src/backend/bisheng/worker/test/test.py | 8 ++ .../bisheng/workflow/nodes/input/input.py | 1 + 9 files changed, 47 insertions(+), 76 deletions(-) delete mode 100644 src/backend/bisheng/worker.py create mode 100644 src/backend/bisheng/worker/config.py create mode 100644 src/backend/bisheng/worker/main.py create mode 100644 src/backend/bisheng/worker/test/__init__.py create mode 100644 src/backend/bisheng/worker/test/test.py diff --git a/docker/bisheng/config/config.yaml b/docker/bisheng/config/config.yaml index d160a3d3c..3e1bbbe55 100644 --- a/docker/bisheng/config/config.yaml +++ b/docker/bisheng/config/config.yaml @@ -20,6 +20,8 @@ redis_url: "redis://redis:6379/1" # sentinel_password: encrypt(gAAAAABlp4b4c59FeVGF_OQRVf6NOUIGdxq8246EBD-b0hdK_jVKRs1x4PoAn0A6C5S6IiFKmWn0Nm5eBUWu-7jxcqw6TiVjQA==) # db: 1 +# celery的broken地址 +celery_redis_url: "redis://redis:6379/2" # 知识库的milvus和es配置 支持使用 !env ${PATH} 填写环境变量的值, 若环境变量不存在则会报错 vector_stores: diff --git a/src/backend/bisheng/settings.py b/src/backend/bisheng/settings.py index d690cf78f..08aade354 100644 --- a/src/backend/bisheng/settings.py +++ b/src/backend/bisheng/settings.py @@ -115,6 +115,7 @@ class Settings(BaseModel): environment: Union[dict, str] = 'dev' database_url: Optional[str] = None redis_url: Optional[Union[str, Dict]] = None + celery_redis_url: Optional[Union[str, Dict]] = None redis: Optional[dict] = None admin: dict = {} cache: str = 'InMemoryCache' @@ -173,6 +174,25 @@ class Settings(BaseModel): values['redis_url'] = new_redis_url return values + @root_validator() + def set_celery_redis_url(cls, values): + if 'celery_redis_url' in values: + if isinstance(values['celery_redis_url'], dict): + for k, v in values['celery_redis_url'].items(): + if isinstance(v, str) and v.startswith('encrypt(') and v.endswith(')'): + v = v[8:-1] + values['celery_redis_url'][k] = decrypt_token(v) + else: + import re + pattern = r'(?<=:)[^:]+(?=@)' # 匹配冒号后面到@符号前面的任意字符 + match = re.search(pattern, values['celery_redis_url']) + if match: + password = match.group(0) + new_password = decrypt_token(password) + new_redis_url = re.sub(pattern, f'{new_password}', values['celery_redis_url']) + values['celery_redis_url'] = new_redis_url + return values + @root_validator() def validate_lists(cls, values): for key, value in values.items(): diff --git a/src/backend/bisheng/worker.py b/src/backend/bisheng/worker.py deleted file mode 100644 index d7ec6a1cd..000000000 --- a/src/backend/bisheng/worker.py +++ /dev/null @@ -1,76 +0,0 @@ -from typing import TYPE_CHECKING, Any, Dict, Optional - -from asgiref.sync import async_to_sync -from bisheng.core.celery_app import celery_app -from bisheng.processing.process import Result, generate_result, process_inputs -from bisheng.services.deps import get_session_service -from bisheng.services.manager import initialize_session_service -from celery.exceptions import SoftTimeLimitExceeded # type: ignore -from loguru import logger -from rich import print - -if TYPE_CHECKING: - from bisheng.graph.vertex.base import Vertex - - -@celery_app.task(acks_late=True) -def test_celery(word: str) -> str: - return f'test task return {word}' - - -@celery_app.task(bind=True, soft_time_limit=30, max_retries=3) -def build_vertex(self, vertex: 'Vertex') -> 'Vertex': - """ - Build a vertex - """ - try: - vertex.task_id = self.request.id - async_to_sync(vertex.build)() - return vertex - except SoftTimeLimitExceeded as e: - raise self.retry(exc=SoftTimeLimitExceeded('Task took too long'), countdown=2) from e - - -@celery_app.task(acks_late=True) -def process_graph_cached_task( - data_graph: Dict[str, Any], - inputs: Optional[dict] = None, - clear_cache=False, - session_id=None, -) -> Dict[str, Any]: - try: - initialize_session_service() - session_service = get_session_service() - - if clear_cache: - session_service.clear_session(session_id) - - if session_id is None: - session_id = session_service.generate_key(session_id=session_id, data_graph=data_graph) - - # Use async_to_sync to handle the asynchronous part of the session service - session_data = async_to_sync(session_service.load_session, force_new_loop=True)(session_id, - data_graph) - logger.warning(f'session_data: {session_data}') - graph, artifacts = session_data if session_data else (None, None) - - if not graph: - raise ValueError('Graph not found in the session') - - # Use async_to_sync for the asynchronous build method - built_object = async_to_sync(graph.build, force_new_loop=True)() - - logger.debug(f'Built object: {built_object}') - - processed_inputs = process_inputs(inputs, artifacts or {}) - result = generate_result(built_object, processed_inputs) - - # Update the session with the new data - session_service.update_session(session_id, (graph, artifacts)) - result_object = Result(result=result, session_id=session_id).model_dump() - print(f'Result object: {result_object}') - return result_object - except Exception as e: - logger.error(f'Error in process_graph_cached_task: {e}') - # Handle the exception as needed, maybe re-raise or return an error message - raise diff --git a/src/backend/bisheng/worker/__init__.py b/src/backend/bisheng/worker/__init__.py index e69de29bb..4da0afdf7 100644 --- a/src/backend/bisheng/worker/__init__.py +++ b/src/backend/bisheng/worker/__init__.py @@ -0,0 +1,3 @@ +# register tasks +from bisheng.worker.test.test import * +from bisheng.worker.knowledge.file_worker import * diff --git a/src/backend/bisheng/worker/config.py b/src/backend/bisheng/worker/config.py new file mode 100644 index 000000000..4c2dfd4fe --- /dev/null +++ b/src/backend/bisheng/worker/config.py @@ -0,0 +1,9 @@ +from bisheng.settings import settings + +broker_url = settings.celery_redis_url + +task_serializer = 'json' +result_serializer = 'json' +accept_content = ['json'] +timezone = 'Asia/Shanghai' +enable_utc = False diff --git a/src/backend/bisheng/worker/main.py b/src/backend/bisheng/worker/main.py new file mode 100644 index 000000000..634f6970f --- /dev/null +++ b/src/backend/bisheng/worker/main.py @@ -0,0 +1,4 @@ +from celery import Celery + +bisheng_celery = Celery('bisheng', include=['bisheng.worker']) +bisheng_celery.config_from_object('bisheng.worker.config') diff --git a/src/backend/bisheng/worker/test/__init__.py b/src/backend/bisheng/worker/test/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/src/backend/bisheng/worker/test/test.py b/src/backend/bisheng/worker/test/test.py new file mode 100644 index 000000000..8c4119873 --- /dev/null +++ b/src/backend/bisheng/worker/test/test.py @@ -0,0 +1,8 @@ +from loguru import logger + +from bisheng.worker.main import bisheng_celery + +@bisheng_celery.task +def add(x,y): + logger.info(f"add {x} + {y}") + return x+y diff --git a/src/backend/bisheng/workflow/nodes/input/input.py b/src/backend/bisheng/workflow/nodes/input/input.py index d39fd18c8..a7a1bd0c4 100644 --- a/src/backend/bisheng/workflow/nodes/input/input.py +++ b/src/backend/bisheng/workflow/nodes/input/input.py @@ -49,6 +49,7 @@ class InputNode(BaseNode): 记录文件的metadata数据 """ if not value: + logger.warning(f"{self.id}.{key} value is None") return None # 1、获取默认的embedding模型