Repository navigation
Expand file tree
/
Copy pathmodule_registration.py
More file actions
executable file
·142 lines (110 loc) · 4.84 KB
/
Copy pathmodule_registration.py
File metadata and controls
executable file
·142 lines (110 loc) · 4.84 KB
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
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
# Copyright(C) [2026] Advanced Micro Devices, Inc. All rights reserved.
import logging
from enum import Enum
from typing import Callable
from agents import AGENT_REGISTRY
AGENT_ALIASES = {"geak_v4": "geak", "forge_operator2flydsl": "forge"}
class AgentType(Enum):
"""Enumeration of supported agent types."""
CURSOR = "cursor"
CLAUDE_CODE = "claude_code"
CODEX = "codex"
DEEPSEEK_HARNESS = "deepseek_harness"
TASK_VALIDATOR = "task_validator"
GEAK = "geak"
FORGE = "forge"
@classmethod
def from_string(cls, agent_string: str) -> 'AgentType':
"""
Convert string to AgentType enum.
Args:
agent_string: String representation of agent name
Returns:
AgentType enum
Raises:
ValueError: If agent_string is not a valid agent type
"""
normalized = agent_string.lower().replace("-", "_")
normalized = AGENT_ALIASES.get(normalized, normalized)
for agent_type in cls:
if agent_type.value == normalized:
return agent_type
# If no match found, raise error with available options
valid_options = [agent.value for agent in cls]
raise ValueError(f"Invalid agent type: '{agent_string}'. Valid options are: {valid_options}")
def load_agent_launcher(agent_type: AgentType, logger: logging.Logger) -> Callable[..., str]:
"""
Dynamically load agent launcher function from agent registry.
Args:
agent_type: AgentType enum
logger: Logger instance
Returns:
Agent launcher function
"""
agent_name = agent_type.value
# Import all agent modules to trigger registration
try:
if agent_type == AgentType.CURSOR:
from agents.cursor import launch_agent # noqa: F401
elif agent_type == AgentType.CLAUDE_CODE:
from agents.claude_code import launch_agent # noqa: F401
elif agent_type == AgentType.CODEX:
from agents.codex import launch_agent # noqa: F401
elif agent_type == AgentType.DEEPSEEK_HARNESS:
from agents.deepseek_harness import launch_agent # noqa: F401
elif agent_type == AgentType.TASK_VALIDATOR:
from agents.task_validator import launch_agent # noqa: F401
elif agent_type == AgentType.GEAK:
from agents.geak import launch_agent # noqa: F401
elif agent_type == AgentType.FORGE:
from agents.forge import launch_agent # noqa: F401
except ImportError as e:
logger.error(f"Failed to import agent {agent_name}: {e}")
raise
# Get agent from registry
if agent_name not in AGENT_REGISTRY:
raise ValueError(f"Agent '{agent_name}' not found in registry. Available agents: {list(AGENT_REGISTRY.keys())}")
logger.info(f"Loaded agent: {agent_name}")
return AGENT_REGISTRY[agent_name]
def load_post_processing_handler(agent_type: AgentType, logger: logging.Logger) -> Callable[[list, logging.Logger], None]:
"""
Dynamically load post-processing function based on agent type.
Args:
agent_type: AgentType enum
logger: Logger instance
Returns:
Post-processing function for the agent
Raises:
NotImplementedError: If agent doesn't have post-processing support
"""
from src.postprocessing import general_post_processing
agent_name = agent_type.value
# Map agents to their post-processing functions
if agent_type == AgentType.TASK_VALIDATOR:
from agents.task_validator.validation_postprocessing import validation_post_processing
logger.info(f"Using validation_post_processing for agent: {agent_name}")
return validation_post_processing
elif agent_type in [AgentType.CURSOR, AgentType.CLAUDE_CODE, AgentType.CODEX, AgentType.DEEPSEEK_HARNESS, AgentType.GEAK, AgentType.FORGE]:
logger.info(f"Using general_post_processing for agent: {agent_name}")
return general_post_processing
else:
raise NotImplementedError(f"Post-processing not implemented for agent: {agent_name}")
def load_prompt_builder(agent_type: AgentType, logger: logging.Logger) -> Callable:
"""
Dynamically load prompt builder function based on agent type.
Args:
agent_type: AgentType enum
logger: Logger instance
Returns:
Prompt builder function
Raises:
NotImplementedError: If agent doesn't have prompt builder support
"""
from src.prompt_builder import prompt_builder
agent_name = agent_type.value
# Map agents to their prompt builder functions
if agent_type in [AgentType.CURSOR, AgentType.CLAUDE_CODE, AgentType.CODEX, AgentType.DEEPSEEK_HARNESS]:
logger.info(f"Using standard prompt_builder for agent: {agent_name}")
return prompt_builder
else:
raise NotImplementedError(f"Prompt builder not implemented for agent: {agent_name}")