fix: add celery test code
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,3 @@
|
||||
# register tasks
|
||||
from bisheng.worker.test.test import *
|
||||
from bisheng.worker.knowledge.file_worker import *
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,4 @@
|
||||
from celery import Celery
|
||||
|
||||
bisheng_celery = Celery('bisheng', include=['bisheng.worker'])
|
||||
bisheng_celery.config_from_object('bisheng.worker.config')
|
||||
@@ -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模型
|
||||
|
||||
Reference in New Issue
Block a user