fix: add celery test code

This commit is contained in:
GuoQing Zhang
2024-12-12 15:48:35 +08:00
parent 3e0b7b4737
commit 3da93fadf5
9 changed files with 47 additions and 76 deletions
+2
View File
@@ -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:
+20
View File
@@ -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():
-76
View File
@@ -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
+3
View File
@@ -0,0 +1,3 @@
# register tasks
from bisheng.worker.test.test import *
from bisheng.worker.knowledge.file_worker import *
+9
View File
@@ -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
+4
View File
@@ -0,0 +1,4 @@
from celery import Celery
bisheng_celery = Celery('bisheng', include=['bisheng.worker'])
bisheng_celery.config_from_object('bisheng.worker.config')
+8
View File
@@ -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
@@ -49,6 +49,7 @@ class InputNode(BaseNode):
记录文件的metadata数据
"""
if not value:
logger.warning(f"{self.id}.{key} value is None")
return None
# 1、获取默认的embedding模型