@@ -29,6 +29,11 @@ async def __call(self, payload: Dict[str, Any]) -> Dict[str, Any]:
2929 return {"error" : str (e )}
3030 async def __generate_async (self , payload : Dict [str , Any ]) -> Dict [str , Any ]:
3131 return await self .__call (payload = payload )
32+ def __extract_first_text (self , result : dict ) -> str :
33+ for o in result .get ("output" , []):
34+ if o .get ("output" , {}).get ("type" ) == "TEXT" :
35+ return o ["output" ].get ("text" , "" )
36+ return ""
3237 def generate (self , input : Union [str , Dict [str , Any ]], model : Optional [str ] = "auto" , track : Optional [str ] = None ) -> Dict [str , Any ]:
3338 if not track :
3439 raise ValueError ("Track is required." )
@@ -44,4 +49,8 @@ def generate(self, input : Union[str, Dict[str, Any]], model : Optional[str] = "
4449 "model" : model ,
4550 "track" : track
4651 }
47- return asyncio .run (self .__generate_async (payload = payload ))
52+ return asyncio .run (self .__generate_async (payload = payload ))
53+ def get_output_text (self , input : Union [str , Dict [str , Any ]], model : Optional [str ] = "auto" , track : Optional [str ] = None ) -> str :
54+ result = self .generate (input = input , model = model , track = track )
55+ first_text : str = self .__extract_first_text (result )
56+ return first_text
0 commit comments