Skip to content

Commit df7789b

Browse files
authored
feat: Generate action plan accounts for user's query and its language (#100)
1 parent b52bf94 commit df7789b

4 files changed

Lines changed: 12 additions & 6 deletions

File tree

app/src/app_config.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@ def db_session(self) -> db.Session:
3333
PROMPT_VERSIONS: dict = {
3434
"extract_supports": "UHJvbXB0VmVyc2lvbjo0Ng==",
3535
"generate_referrals": "UHJvbXB0VmVyc2lvbjo0OA==",
36-
"generate_action_plan": "UHJvbXB0VmVyc2lvbjo1MQ==",
36+
"generate_action_plan": "UHJvbXB0VmVyc2lvbjo1Mg==",
3737
"crawl_gcta": "UHJvbXB0VmVyc2lvbjozNg==",
3838
"crawl_indeed": "UHJvbXB0VmVyc2lvbjozNA==",
3939
}

app/src/pipelines/generate_action_plan/pipeline_wrapper.py

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -45,7 +45,8 @@ def setup(self) -> None:
4545
prompt_template = haystack_utils.get_phoenix_prompt("generate_action_plan")
4646
pipeline.add_component(
4747
instance=ChatPromptBuilder(
48-
template=prompt_template, required_variables=["resources", "action_plan_json"]
48+
template=prompt_template,
49+
required_variables=["resources", "action_plan_json", "user_query"],
4950
),
5051
name="prompt_builder",
5152
)
@@ -57,7 +58,9 @@ def setup(self) -> None:
5758
self.pipeline = pipeline
5859

5960
# Called for the `generate-action-plan/run` endpoint
60-
def run_api(self, resources: list[Resource] | list[dict], user_email: str) -> dict:
61+
def run_api(
62+
self, resources: list[Resource] | list[dict], user_email: str, user_query: str
63+
) -> dict:
6164
resource_objects = get_resources(resources)
6265

6366
with using_attributes(user_id=user_email), using_metadata({"user_id": user_email}):
@@ -66,13 +69,13 @@ def run_api(self, resources: list[Resource] | list[dict], user_email: str) -> di
6669
with tracer.start_as_current_span( # pylint: disable=not-context-manager,unexpected-keyword-arg
6770
self.name, openinference_span_kind="chain"
6871
) as span:
69-
result = self._run(resource_objects, user_email)
72+
result = self._run(resource_objects, user_email, user_query)
7073
span.set_input([r.name for r in resource_objects])
7174
span.set_output(result["response"])
7275
span.set_status(Status(StatusCode.OK))
7376
return result
7477

75-
def _run(self, resource_objects: list[Resource], user_email: str) -> dict:
78+
def _run(self, resource_objects: list[Resource], user_email: str, user_query: str) -> dict:
7679
response = self.pipeline.run(
7780
{
7881
"logger": {
@@ -83,6 +86,7 @@ def _run(self, resource_objects: list[Resource], user_email: str) -> dict:
8386
"prompt_builder": {
8487
"resources": format_resources(resource_objects),
8588
"action_plan_json": action_plan_as_json,
89+
"user_query": user_query,
8690
},
8791
"llm": {"model": "gpt-5-mini", "reasoning_effort": "low"},
8892
},

frontend/src/app/[locale]/generate-referrals/page.tsx

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -162,7 +162,7 @@ export default function Page() {
162162

163163
try {
164164
const { actionPlan: plan, errorMessage: planError } =
165-
await fetchActionPlan(selectedResources, userEmail);
165+
await fetchActionPlan(selectedResources, userEmail, clientDescription);
166166
setActionPlan(plan);
167167
if (planError) {
168168
setErrorMessage(planError);

frontend/src/util/fetchActionPlan.ts

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -52,6 +52,7 @@ function fixJsonControlCharacters(jsonString: string): string {
5252
export async function fetchActionPlan(
5353
resources: Resource[],
5454
userEmail: string,
55+
userQuery: string,
5556
): Promise<{ actionPlan: ActionPlan | null; errorMessage?: string }> {
5657
const apiDomain = await getApiDomain();
5758
const url = apiDomain + "generate_action_plan/run";
@@ -69,6 +70,7 @@ export async function fetchActionPlan(
6970
body: JSON.stringify({
7071
resources: resources,
7172
user_email: userEmail,
73+
user_query: userQuery,
7274
}),
7375
cache: "no-store",
7476
signal: ac.signal,

0 commit comments

Comments
 (0)