@@ -20,12 +20,31 @@ class ProfileAgent:
2020 3. summarize() - If round 2 executed, generate final summary
2121 """
2222
23- def __init__ (self , model_name : str ) -> None :
23+ def __init__ (
24+ self ,
25+ model_name : str ,
26+ * ,
27+ prompt_dir : Path | None = None ,
28+ template_dir : Path | None = None ,
29+ profile_prompt_name : str = "profile_agent" ,
30+ profile_template_name : str = "profile_agent.jinja2" ,
31+ decide_prompt_name : str = "profile_agent_decide" ,
32+ decide_template_name : str = "profile_agent_decide.jinja2" ,
33+ summarize_prompt_name : str = "profile_agent_summarize" ,
34+ summarize_template_name : str = "profile_agent_summarize.jinja2" ,
35+ ) -> None :
2436 self .model_name = model_name
2537 self ._llm = None
26- template_dir = Path (__file__ ).parent / "prompts" / "templates"
38+ self .prompt_dir = prompt_dir
39+ tmpl_dir = template_dir or (Path (__file__ ).parent / "prompts" / "templates" )
40+ self .profile_prompt_name = profile_prompt_name
41+ self .profile_template_name = profile_template_name
42+ self .decide_prompt_name = decide_prompt_name
43+ self .decide_template_name = decide_template_name
44+ self .summarize_prompt_name = summarize_prompt_name
45+ self .summarize_template_name = summarize_template_name
2746 self .jinja_env = Environment (
28- loader = FileSystemLoader (template_dir ),
47+ loader = FileSystemLoader (tmpl_dir ),
2948 trim_blocks = True ,
3049 lstrip_blocks = True ,
3150 )
@@ -37,9 +56,13 @@ def _get_llm(self):
3756 raise RuntimeError ("LLM configuration not found." )
3857 return self ._llm
3958
40- def _render_prompt (self , name : str , context : Dict [str , Any ]) -> str :
41- prompt_cfg = load_prompt_yaml (name , required_keys = ("system" , "guidelines" ))
42- template = self .jinja_env .get_template (f"{ name } .jinja2" )
59+ def _render_prompt (self , prompt_name : str , template_name : str , context : Dict [str , Any ]) -> str :
60+ prompt_cfg = load_prompt_yaml (
61+ prompt_name ,
62+ required_keys = ("system" , "guidelines" ),
63+ prompt_dir = self .prompt_dir ,
64+ )
65+ template = self .jinja_env .get_template (template_name )
4366 return template .render (
4467 system_prompt_text = prompt_cfg ["system" ],
4568 guidelines_text = prompt_cfg ["guidelines" ],
@@ -59,7 +82,8 @@ def generate_profile_code(
5982 (code, raw_response, messages)
6083 """
6184 prompt_content = self ._render_prompt (
62- "profile_agent" ,
85+ self .profile_prompt_name ,
86+ self .profile_template_name ,
6387 {
6488 "query_md" : session_state .get ("query" , "" ),
6589 "task_dir" : session_state .get ("task_dir" , "" ),
@@ -102,7 +126,8 @@ def decide_or_summarize(
102126 stderr_section = f"- Stderr:\n ```\n { stderr [:1000 ]} \n ```"
103127
104128 prompt_content = self ._render_prompt (
105- "profile_agent_decide" ,
129+ self .decide_prompt_name ,
130+ self .decide_template_name ,
106131 {
107132 "query_md" : session_state .get ("query" , "" ),
108133 "task_dir" : session_state .get ("task_dir" , "" ),
@@ -206,7 +231,8 @@ def summarize(
206231 r2_stdout = (round2_result .get ("stdout" ) or "" )[:2000 ]
207232
208233 prompt_content = self ._render_prompt (
209- "profile_agent_summarize" ,
234+ self .summarize_prompt_name ,
235+ self .summarize_template_name ,
210236 {
211237 "query_md" : session_state .get ("query" , "" ),
212238 "round1_status" : r1_status ,
0 commit comments