99from hayhooks import BasePipelineWrapper
1010from haystack import Pipeline
1111from haystack .components .builders import ChatPromptBuilder
12- from haystack .core .errors import PipelineRuntimeError
1312from haystack .dataclasses .chat_message import ChatMessage
14- from openinference .instrumentation import _tracers , using_attributes , using_metadata
15- from opentelemetry .trace .status import Status , StatusCode
1613from pydantic import BaseModel
1714
1815from src .app_config import config
@@ -63,6 +60,10 @@ class PipelineWrapper(BasePipelineWrapper):
6360 name = "generate_referrals"
6461
6562 def setup (self ) -> None :
63+ self .pipeline = self ._create_pipeline ()
64+ self .runner = haystack_utils .TracedPipelineRunner (self .name , self .pipeline )
65+
66+ def _create_pipeline (self ) -> Pipeline :
6667 # Do not rely on max_runs_per_component strictly, i.e., a component may run max_runs_per_component+1 times.
6768 # The component_visits counter for max_runs_per_component is reset with each call to pipeline.run()
6869 pipeline = Pipeline (max_runs_per_component = 3 )
@@ -96,31 +97,11 @@ def setup(self) -> None:
9697
9798 pipeline .add_component ("logger" , components .ReadableLogger ())
9899 pipeline .connect ("output_validator.valid_replies" , "logger" )
99-
100- self .pipeline = pipeline
100+ return pipeline
101101
102102 # Called for the `generate-referrals/run` endpoint
103103 def run_api (
104104 self , query : str , user_email : str , prompt_version_id : str = "" , suffix : str = ""
105- ) -> dict :
106- with using_attributes (user_id = user_email ), using_metadata ({"user_id" : user_email }):
107- # Must set using_metadata context before calling tracer.start_as_current_span()
108- assert isinstance (tracer , _tracers .OITracer ), f"Got unexpected { type (tracer )} "
109- with tracer .start_as_current_span ( # pylint: disable=not-context-manager,unexpected-keyword-arg
110- self .name , openinference_span_kind = "chain"
111- ) as span :
112- result = self ._run (query , user_email , prompt_version_id , suffix )
113- span .set_input (query )
114- try :
115- resp_obj = json .loads (result ["llm" ]["replies" ][- 1 ].text )
116- span .set_output ([r ["name" ] for r in resp_obj ["resources" ]])
117- except (KeyError , IndexError ):
118- span .set_output (result ["llm" ]["replies" ][- 1 ].text )
119- span .set_status (Status (StatusCode .OK ))
120- return result
121-
122- def _run (
123- self , query : str , user_email : str , prompt_version_id : str = "" , suffix : str = ""
124105 ) -> dict :
125106 # Retrieve the requested prompt (with optional prompt_version_id and/or suffix)
126107 try :
@@ -132,20 +113,26 @@ def _run(
132113 status_code = 422 ,
133114 detail = f"The requested prompt version '{ prompt_version_id } ' with suffix '{ suffix } ' could not be retrieved due to HTTP status { he .response .status_code } " ,
134115 ) from he
135-
136- try :
137- response = self .pipeline .run (
138- self ._run_arg_data (query , user_email , prompt_template ),
139- include_outputs_from = {"llm" , "save_result" },
140- )
141- logger .debug ("Results: %s" , pformat (response , width = 160 ))
142- return response
143- except PipelineRuntimeError as re :
144- logger .error ("PipelineRuntimeError: %s" , re , exc_info = True )
145- raise HTTPException (status_code = 500 , detail = str (re )) from re
146- except Exception as e :
147- logger .error ("Error %s: %s" , type (e ), e , exc_info = True )
148- raise HTTPException (status_code = 500 , detail = f"Internal error: { str (e )} " ) from e
116+ pipeline_run_args = self ._run_arg_data (query , user_email , prompt_template )
117+
118+ def extract_output (result : dict ) -> list | str :
119+ try :
120+ resp_obj = json .loads (result ["llm" ]["replies" ][- 1 ].text )
121+ return [r ["name" ] for r in resp_obj ["resources" ]]
122+ except (KeyError , IndexError ):
123+ return result ["llm" ]["replies" ][- 1 ].text
124+
125+ response = self .runner .return_response (
126+ pipeline_run_args ,
127+ user_id = user_email ,
128+ metadata = {"user_id" : user_email },
129+ include_outputs_from = {"llm" , "save_result" },
130+ input_ = query ,
131+ extract_output = extract_output ,
132+ parent_span_name_suffix = suffix ,
133+ )
134+ logger .debug ("Results: %s" , pformat (response , width = 160 ))
135+ return response
149136
150137 def _run_arg_data (
151138 self , query : str , user_email : str , prompt_template : list [ChatMessage ]
0 commit comments