Repository navigation
get_system_prompt(task) hook on LLMTaskWorker #3
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 1 commit
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -14,4 +14,4 @@ | |
|
|
||
| """Version information for PlanAI.""" | ||
|
|
||
| __version__ = "0.7.0" | ||
| __version__ = "0.7.1" | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -146,6 +146,21 @@ def get_tools(self, task: Task) -> Optional[List[Tool]]: | |
| """ | ||
| return self.tools | ||
|
|
||
| def get_system_prompt(self, task: Task) -> Optional[str]: | ||
| """ | ||
| Returns the system prompt to use for this task. Defaults to the static | ||
| ``system_prompt`` field. Override it to build a per-task system prompt, for | ||
| example to place reference material that several tasks share ahead of the | ||
| per-task instructions, where provider prompt caching can reuse it. | ||
|
|
||
| Args: | ||
| task (Task): The input task. | ||
|
|
||
| Returns: | ||
| Optional[str]: The system prompt, or None for the provider default. | ||
| """ | ||
| return self.system_prompt | ||
|
|
||
| def get_cache_salt(self, task: Task) -> Optional[str]: | ||
| """ | ||
| Returns an optional salt to forward as ``cache_salt`` to generate_pydantic. | ||
|
|
@@ -203,7 +218,7 @@ def extra_validation_with_task(response: BaseModel): | |
| ) | ||
| ), | ||
| output_schema=self._output_type(), | ||
| system=self.system_prompt, | ||
| system=self.get_system_prompt(task), | ||
| tools=tools if tools else None, | ||
| task=self._format_task(processed_task), | ||
| temperature=self.temperature, | ||
|
|
@@ -233,7 +248,7 @@ def get_full_prompt(self, task: Task) -> str: | |
| else "" | ||
| ) | ||
| ), | ||
| system=self.system_prompt, | ||
| system=self.get_system_prompt(task), | ||
| task=self._format_task(processed_task), | ||
| instructions=task_prompt, | ||
| format_instructions=LLMInterface.get_format_instructions( | ||
|
|
@@ -338,5 +353,5 @@ def _get_cache_key(self, task: Task) -> str: | |
| """Generate a unique cache key for the input task including the prompt template and model name.""" | ||
| upstream_cache_key = super()._get_cache_key(task) | ||
|
|
||
| upstream_cache_key += f" - {self.system_prompt} - {self.get_full_prompt(task)} - {self.llm.model_name}" | ||
| upstream_cache_key += f" - {self.get_system_prompt(task)} - {self.get_full_prompt(task)} - {self.llm.model_name}" | ||
|
Owner
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Partly: generate_full_prompt() ignores its system argument, so the full prompt does not contain the system prompt and it has to stay in the key explicitly. What was redundant is the double evaluation: the key now evaluates get_system_prompt(task) once and passes it to get_full_prompt(task, system_prompt=...). Test added asserting a single call. |
||
| return hashlib.sha1(upstream_cache_key.encode()).hexdigest() | ||
Uh oh!
There was an error while loading. Please reload this page.