File size: 2,737 Bytes
d82bbe4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
from __future__ import annotations

import asyncio
from typing import Any, Dict, List, Optional
from langchain_openai import ChatOpenAI
from langchain_core.messages import SystemMessage, HumanMessage, BaseMessage
from langchain_core.tools import Tool

from dataflow_agent.promptstemplates.prompt_template import PromptsTemplateGenerator
from dataflow_agent.state import DFState
from dataflow_agent.utils import robust_parse_json
from dataflow_agent.toolkits.tool_manager import ToolManager
from dataflow_agent.logger import get_logger

log = get_logger(__name__)

from dataflow_agent.agentroles.cores.base_agent import BaseAgent

class DataContentClassifier(BaseAgent):
    """数据内容分类器 - 继承自BaseAgent"""
    # 后续都改成类方法
    @classmethod
    def create(cls, tool_manager: Optional[ToolManager] = None, **kwargs):
        return cls(tool_manager=tool_manager, **kwargs)

    @property
    def role_name(self) -> str:
        return "classifier"
    
    @property
    def system_prompt_template_name(self) -> str:
        return "system_prompt_for_data_content_classification"
    
    @property
    def task_prompt_template_name(self) -> str:
        return "task_prompt_for_data_content_classification"
    
    def get_task_prompt_params(self, pre_tool_results: Dict[str, Any]) -> Dict[str, Any]:
        """数据分类器特有的提示词参数"""
        return {
            'local_tool_for_sample': pre_tool_results.get('sample', ''),
            'local_tool_for_get_categories': pre_tool_results.get('categories', '[]'),
        }
    
    def get_default_pre_tool_results(self) -> Dict[str, Any]:
        """数据分类器的默认前置工具结果"""
        return {
            'sample': '',
            'categories': '[]'
        }
    
    def update_state_result(self, state: DFState, result: Dict[str, Any], pre_tool_results: Dict[str, Any]):
        """自定义状态更新 - 保持向后兼容"""
        state.category = result 
        super().update_state_result(state, result, pre_tool_results)

async def data_content_classification(
    state: DFState, 
    model_name: Optional[str] = None,
    tool_manager: Optional[ToolManager] = None,
    temperature: float = 0.0,
    max_tokens: int = 512,
    use_agent: bool = False,
    **kwargs,
) -> DFState:
    classifier = DataContentClassifier(
        tool_manager=tool_manager,
        model_name=model_name,
        temperature=temperature,
        max_tokens=max_tokens,
    )
    return await classifier.execute(state, use_agent=use_agent, **kwargs)

def create_classifier(tool_manager: Optional[ToolManager] = None, **kwargs) -> DataContentClassifier:
    return DataContentClassifier(tool_manager=tool_manager, **kwargs)