11import logging
22from pprint import pformat
3+ from typing import Generator
34
45from hayhooks import BasePipelineWrapper
56from haystack import Pipeline
67from haystack .components .builders import ChatPromptBuilder
7- from openinference .instrumentation import _tracers , using_attributes , using_metadata
8- from opentelemetry .trace .status import Status , StatusCode
98from pydantic import BaseModel
109
1110from src .app_config import config
12- from src .common import haystack_utils , phoenix_utils
11+ from src .common import haystack_utils
1312from src .common .components import (
1413 LlmOutputValidator ,
1514 OpenAIWebSearchGenerator ,
1918from src .pipelines .generate_referrals .pipeline_wrapper import Resource
2019
2120logger = logging .getLogger (__name__ )
22- tracer = phoenix_utils .tracer_provider .get_tracer (__name__ )
2321
2422
2523class ActionPlan (BaseModel ):
@@ -48,11 +46,9 @@ def setup(self) -> None:
4846 pipeline = Pipeline ()
4947 pipeline .add_component ("llm" , create_websearch ())
5048
51- prompt_template = haystack_utils .get_phoenix_prompt ("generate_action_plan" )
5249 pipeline .add_component (
5350 instance = ChatPromptBuilder (
54- template = prompt_template ,
55- required_variables = ["resources" , "action_plan_json" , "user_query" ],
51+ variables = ["resources" , "action_plan_json" , "user_query" ],
5652 ),
5753 name = "prompt_builder" ,
5854 )
@@ -67,51 +63,99 @@ def setup(self) -> None:
6763 pipeline .connect ("llm" , "logger" )
6864
6965 self .pipeline = pipeline
66+ self .runner = haystack_utils .TracedPipelineRunner (self .name , self .pipeline )
7067
7168 # Called for the `generate-action-plan/run` endpoint
7269 def run_api (
73- self , resources : list [Resource ] | list [dict ], user_email : str , user_query : str
70+ self ,
71+ resources : list [Resource ] | list [dict ],
72+ user_email : str ,
73+ user_query : str ,
7474 ) -> dict :
7575 resource_objects = get_resources (resources )
76-
77- with using_attributes (user_id = user_email ), using_metadata ({"user_id" : user_email }):
78- # Must set using_metadata context before calling tracer.start_as_current_span()
79- assert isinstance (tracer , _tracers .OITracer ), f"Got unexpected { type (tracer )} "
80- with tracer .start_as_current_span ( # pylint: disable=not-context-manager,unexpected-keyword-arg
81- self .name , openinference_span_kind = "chain"
82- ) as span :
83- result = self ._run (resource_objects , user_email , user_query )
84- span .set_input ([r .name for r in resource_objects ])
85- span .set_output (result ["response" ])
86- span .set_status (Status (StatusCode .OK ))
87- return result
88-
89- def _run (self , resource_objects : list [Resource ], user_email : str , user_query : str ) -> dict :
90- response = self .pipeline .run (
91- {
92- "logger" : {
93- "messages_list" : [
94- {"resource_count" : len (resource_objects ), "user_email" : user_email }
95- ],
96- },
97- "prompt_builder" : {
98- "resources" : format_resources (resource_objects ),
99- "action_plan_json" : action_plan_as_json ,
100- "user_query" : user_query ,
101- },
102- "llm" : {
103- "model" : config .generate_action_plan_model_version ,
104- "reasoning_effort" : config .generate_action_plan_reasoning_level ,
105- },
106- },
76+ pipeline_run_args = self .create_pipeline_args (
77+ user_email ,
78+ resource_objects ,
79+ user_query ,
80+ )
81+ response = self .runner .return_response (
82+ pipeline_run_args ,
83+ user_id = user_email ,
84+ metadata = {"user_id" : user_email },
10785 include_outputs_from = {"llm" , "save_result" },
86+ input_ = [r .name for r in resource_objects ],
87+ extract_output = lambda response : response ["llm" ]["replies" ][0 ]._content [0 ].text ,
10888 )
10989 logger .debug ("Results: %s" , pformat (response , width = 160 ))
90+
11091 return {
11192 "response" : response ["llm" ]["replies" ][0 ]._content [0 ].text ,
11293 "save_result" : response ["save_result" ],
11394 }
11495
96+ def create_pipeline_args (
97+ self ,
98+ user_email : str ,
99+ resource_objects : list [Resource ],
100+ user_query : str ,
101+ * ,
102+ llm_model : str | None = None ,
103+ reasoning_effort : str | None = None ,
104+ streaming : bool = False ,
105+ ) -> dict :
106+ prompt_template = haystack_utils .get_phoenix_prompt ("generate_action_plan" )
107+ return {
108+ "logger" : {
109+ "messages_list" : [
110+ {"resource_count" : len (resource_objects ), "user_email" : user_email }
111+ ],
112+ },
113+ "prompt_builder" : {
114+ "template" : prompt_template ,
115+ "resources" : format_resources (resource_objects ),
116+ "action_plan_json" : action_plan_as_json ,
117+ "user_query" : user_query ,
118+ },
119+ "llm" : {
120+ "model" : llm_model or config .generate_action_plan_model_version ,
121+ "reasoning_effort" : reasoning_effort or config .generate_action_plan_reasoning_level ,
122+ "streaming" : streaming ,
123+ },
124+ }
125+
126+ # https://docs.haystack.deepset.ai/docs/hayhooks#openai-compatibility
127+ # Called for the `{pipeline_name}/chat`, `/chat/completions`, or `/v1/chat/completions` streaming endpoint using Server-Sent Events (SSE)
128+ def run_chat_completion (self , model : str , messages : list , body : dict ) -> Generator :
129+ # Note: 'model' parameter is the pipeline name, not the LLM model
130+ assert model == self .name , f"Unexpected model/pipeline name: { model } "
131+
132+ # Extract custom parameters from the body
133+ resources = body .get ("resources" , [])
134+ user_email = body .get ("user_email" , "" )
135+ user_query = body .get ("user_query" , "" )
136+
137+ if not resources :
138+ raise ValueError ("resources parameter is required" )
139+ if not user_email :
140+ raise ValueError ("user_email parameter is required" )
141+
142+ resource_objects = get_resources (resources )
143+ pipeline_run_args = self .create_pipeline_args (
144+ user_email ,
145+ resource_objects ,
146+ user_query ,
147+ llm_model = body .get ("llm_model" , None ),
148+ reasoning_effort = body .get ("reasoning_effort" , None ),
149+ streaming = True ,
150+ )
151+ logger .info ("Streaming action plan: %s" , pipeline_run_args )
152+ return self .runner .stream_response (
153+ pipeline_run_args ,
154+ user_id = user_email ,
155+ metadata = {"user_id" : user_email },
156+ input_ = [r .name for r in resource_objects ],
157+ )
158+
115159
116160def get_resources (resources : list [Resource ] | list [dict ]) -> list [Resource ]:
117161 """Ensure we have a list of Resource objects."""
0 commit comments