Skip to content

Commit bd38d76

Browse files
FEAT: add prompt suffix handling (#134)
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
1 parent 7a841e1 commit bd38d76

8 files changed

Lines changed: 67 additions & 31 deletions

File tree

app/src/app_config.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -39,8 +39,9 @@ def db_session(self) -> db.Session:
3939
# so they are not unique across different Phoenix instances.
4040
PROMPT_VERSIONS: dict = {
4141
"extract_supports": "UHJvbXB0VmVyc2lvbjo0Ng==",
42-
"generate_referrals": "UHJvbXB0VmVyc2lvbjo0OA==",
43-
"generate_action_plan": "UHJvbXB0VmVyc2lvbjo1Mg==",
42+
"generate_referrals": "UHJvbXB0VmVyc2lvbjo2MA==", # if no suffix, the default Austin area prompt will be used
43+
"generate_referrals_keystone": "UHJvbXB0VmVyc2lvbjo2Nw==",
44+
"generate_action_plan": "UHJvbXB0VmVyc2lvbjo2OQ==",
4445
"crawl_gcta": "UHJvbXB0VmVyc2lvbjozNg==",
4546
"crawl_indeed": "UHJvbXB0VmVyc2lvbjozNA==",
4647
}

app/src/common/haystack_utils.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,11 @@
66
from src.common import phoenix_utils
77

88

9-
def get_phoenix_prompt(prompt_name: str, prompt_version_id: str = "") -> list[ChatMessage]:
10-
prompt_ver = phoenix_utils.get_prompt_template(prompt_name, prompt_version_id)
9+
def get_phoenix_prompt(
10+
prompt_name: str, prompt_version_id: str = "", suffix: str = ""
11+
) -> list[ChatMessage]:
12+
full_prompt_name = f"{prompt_name}_{suffix}" if suffix else prompt_name
13+
prompt_ver = phoenix_utils.get_prompt_template(full_prompt_name, prompt_version_id)
1114
return to_chat_messages(prompt_ver._template["messages"])
1215

1316

app/src/pipelines/generate_action_plan/pipeline_wrapper.py

Lines changed: 12 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -48,11 +48,9 @@ def setup(self) -> None:
4848
pipeline = Pipeline()
4949
pipeline.add_component("llm", create_websearch())
5050

51-
prompt_template = haystack_utils.get_phoenix_prompt("generate_action_plan")
5251
pipeline.add_component(
5352
instance=ChatPromptBuilder(
54-
template=prompt_template,
55-
required_variables=["resources", "action_plan_json", "user_query"],
53+
variables=["resources", "action_plan_json", "user_query"],
5654
),
5755
name="prompt_builder",
5856
)
@@ -70,7 +68,10 @@ def setup(self) -> None:
7068

7169
# Called for the `generate-action-plan/run` endpoint
7270
def run_api(
73-
self, resources: list[Resource] | list[dict], user_email: str, user_query: str
71+
self,
72+
resources: list[Resource] | list[dict],
73+
user_email: str,
74+
user_query: str,
7475
) -> dict:
7576
resource_objects = get_resources(resources)
7677

@@ -87,14 +88,20 @@ def run_api(
8788
return result
8889

8990
def _run(self, resource_objects: list[Resource], user_email: str, user_query: str) -> dict:
91+
prompt_template = haystack_utils.get_phoenix_prompt("generate_action_plan")
92+
9093
response = self.pipeline.run(
9194
{
9295
"logger": {
9396
"messages_list": [
94-
{"resource_count": len(resource_objects), "user_email": user_email}
97+
{
98+
"resource_count": len(resource_objects),
99+
"user_email": user_email,
100+
}
95101
],
96102
},
97103
"prompt_builder": {
104+
"template": prompt_template,
98105
"resources": format_resources(resource_objects),
99106
"action_plan_json": action_plan_as_json,
100107
"user_query": user_query,

app/src/pipelines/generate_referrals/pipeline_wrapper.py

Lines changed: 10 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -100,14 +100,16 @@ def setup(self) -> None:
100100
self.pipeline = pipeline
101101

102102
# Called for the `generate-referrals/run` endpoint
103-
def run_api(self, query: str, user_email: str, prompt_version_id: str = "") -> dict:
103+
def run_api(
104+
self, query: str, user_email: str, prompt_version_id: str = "", suffix: str = ""
105+
) -> dict:
104106
with using_attributes(user_id=user_email), using_metadata({"user_id": user_email}):
105107
# Must set using_metadata context before calling tracer.start_as_current_span()
106108
assert isinstance(tracer, _tracers.OITracer), f"Got unexpected {type(tracer)}"
107109
with tracer.start_as_current_span( # pylint: disable=not-context-manager,unexpected-keyword-arg
108110
self.name, openinference_span_kind="chain"
109111
) as span:
110-
result = self._run(query, user_email, prompt_version_id)
112+
result = self._run(query, user_email, prompt_version_id, suffix)
111113
span.set_input(query)
112114
try:
113115
resp_obj = json.loads(result["llm"]["replies"][-1].text)
@@ -117,16 +119,18 @@ def run_api(self, query: str, user_email: str, prompt_version_id: str = "") -> d
117119
span.set_status(Status(StatusCode.OK))
118120
return result
119121

120-
def _run(self, query: str, user_email: str, prompt_version_id: str = "") -> dict:
121-
# Retrieve the requested prompt_version_id and error if requested prompt version is not found
122+
def _run(
123+
self, query: str, user_email: str, prompt_version_id: str = "", suffix: str = ""
124+
) -> dict:
125+
# Retrieve the requested prompt (with optional prompt_version_id and/or suffix)
122126
try:
123127
prompt_template = haystack_utils.get_phoenix_prompt(
124-
"generate_referrals", prompt_version_id
128+
"generate_referrals", prompt_version_id=prompt_version_id, suffix=suffix
125129
)
126130
except httpx.HTTPStatusError as he:
127131
raise HTTPException(
128132
status_code=422,
129-
detail=f"The requested prompt version '{prompt_version_id}' could not be retrieved due to HTTP status {he.response.status_code}",
133+
detail=f"The requested prompt version '{prompt_version_id}' with suffix '{suffix}' could not be retrieved due to HTTP status {he.response.status_code}",
130134
) from he
131135

132136
try:

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

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -107,6 +107,7 @@ export default function Page() {
107107

108108
async function findResources() {
109109
const prompt_version_id = searchParams?.get("prompt_version_id") ?? null;
110+
const suffix = searchParams?.get("suffix") ?? undefined;
110111

111112
setLoading(true);
112113
setResult(null);
@@ -118,6 +119,7 @@ export default function Page() {
118119
request,
119120
userEmail,
120121
prompt_version_id,
122+
suffix,
121123
);
122124
setResourcesResultId(resultId);
123125
setErrorMessage(errorMessage);

frontend/src/components/ResourcesList.tsx

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -117,7 +117,9 @@ const ResourcesList = ({
117117
</button>
118118
<CardHeader className="p-3 ml-3">
119119
<CardTitle className="text-xl font-semibold text-gray-900 flex items-center gap-2">
120-
<span className={`flex-shrink-0 w-7 h-7 ${getNumberBadgeClass(r.referral_type)} text-white rounded-full flex items-center justify-center text-m font-medium`}>
120+
<span
121+
className={`flex-shrink-0 w-7 h-7 ${getNumberBadgeClass(r.referral_type)} text-white rounded-full flex items-center justify-center text-m font-medium`}
122+
>
121123
{i + 1}
122124
</span>
123125
<div>{r.name}</div>

frontend/src/util/fetchActionPlan.ts

Lines changed: 11 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -68,14 +68,20 @@ export async function fetchActionPlan(
6868
const timer = setTimeout(() => ac.abort(), 120_000);
6969

7070
try {
71+
const requestBody: {
72+
resources: Resource[];
73+
user_email: string;
74+
user_query: string;
75+
} = {
76+
resources: resources,
77+
user_email: userEmail,
78+
user_query: userQuery,
79+
};
80+
7181
const upstream = await fetch(url, {
7282
method: "POST",
7383
headers,
74-
body: JSON.stringify({
75-
resources: resources,
76-
user_email: userEmail,
77-
user_query: userQuery,
78-
}),
84+
body: JSON.stringify(requestBody),
7985
cache: "no-store",
8086
signal: ac.signal,
8187
});

frontend/src/util/fetchResources.ts

Lines changed: 21 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ export async function fetchResources(
55
clientDescription: string,
66
userEmail: string,
77
prompt_version_id: string | null,
8+
suffix?: string,
89
) {
910
const apiDomain = await getApiDomain();
1011

@@ -20,16 +21,26 @@ export async function fetchResources(
2021

2122
const ac = new AbortController();
2223
const timer = setTimeout(() => ac.abort(), 600_000); // TODO make configurable
23-
const requestBody = prompt_version_id
24-
? JSON.stringify({
25-
query: clientDescription,
26-
user_email: userEmail,
27-
prompt_version_id: prompt_version_id,
28-
})
29-
: JSON.stringify({
30-
query: clientDescription,
31-
user_email: userEmail,
32-
});
24+
25+
const requestData: {
26+
query: string;
27+
user_email: string;
28+
prompt_version_id?: string;
29+
suffix?: string;
30+
} = {
31+
query: clientDescription,
32+
user_email: userEmail,
33+
};
34+
35+
if (prompt_version_id) {
36+
requestData.prompt_version_id = prompt_version_id;
37+
}
38+
39+
if (suffix) {
40+
requestData.suffix = suffix;
41+
}
42+
43+
const requestBody = JSON.stringify(requestData);
3344

3445
try {
3546
const upstream = await fetch(url, {

0 commit comments

Comments
 (0)