Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions docs-astro/src/content/docs/api/taskworker.md
Original file line number Diff line number Diff line change
Expand Up @@ -264,6 +264,15 @@ def extra_validation(self, response: Task, input_task: Task) -> Optional[str]:
return None
```

##### get_system_prompt
```python
def get_system_prompt(self, task: Task) -> Optional[str]:
"""System prompt for this task; defaults to the ``system_prompt`` field"""
return self.system_prompt + "\n\n# Reference\n" + load_reference(task)
Comment thread
Copilot marked this conversation as resolved.
Outdated
```

Use it to put material that many tasks share, such as a set of notes every section writer works from, into the system prompt. Provider prompt caching matches prefixes at block boundaries, so a shared system block is reused across tasks while the per-task instructions vary.

##### get_tools
```python
def get_tools(self, task: Task) -> Optional[List[Tool]]:
Expand Down
2 changes: 2 additions & 0 deletions docs-astro/src/content/docs/features/llm-integration.md
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,8 @@ class ExpertAnalyzer(LLMTaskWorker):
output_types: List[Type[Task]] = [Analysis]
```

Override `get_system_prompt(task)` when the system prompt depends on the task, for example to include reference material that several tasks share so the provider's prompt cache can reuse it.

### Structured Output

PlanAI automatically handles structured output using Pydantic models:
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[tool.poetry]
name = "planai"
version = "0.7.0"
version = "0.7.1"
description = "A simple framework for coordinating classical compute and LLM-based tasks."
authors = ["Niels Provos <planai@provos.org>"]
license = "Apache-2.0"
Expand Down
2 changes: 1 addition & 1 deletion src/planai/_version.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,4 +14,4 @@

"""Version information for PlanAI."""

__version__ = "0.7.0"
__version__ = "0.7.1"
21 changes: 18 additions & 3 deletions src/planai/llm_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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}"

Copy link
Copy Markdown
Owner Author

Choose a reason for hiding this comment

The 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()
26 changes: 26 additions & 0 deletions tests/planai/test_llm_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,32 @@ def test_invoke_llm_uses_get_tools_hook(self, mock_publish_work):
call_args = mock_generate_pydantic.call_args
self.assertEqual(call_args.kwargs["tools"], [mock_tool])

@patch("planai.llm_task.LLMTaskWorker.publish_work")
def test_invoke_llm_uses_get_system_prompt_hook(self, mock_publish_work):
output_task_payload = DummyOutputTask(result="out")
input_task = DummyTask(content="in")

with patch.object(
self.llm, "generate_pydantic", return_value=output_task_payload
) as mock_generate_pydantic:
with patch(
"planai.llm_task.LLMTaskWorker.get_system_prompt",
return_value="per-task system prompt",
) as mock_hook:
self.worker._invoke_llm(input_task)

mock_hook.assert_called_once_with(input_task)
self.assertEqual(
mock_generate_pydantic.call_args.kwargs["system"],
"per-task system prompt",
)

def test_get_system_prompt_defaults_to_the_field(self):
self.assertEqual(
self.worker.get_system_prompt(DummyTask(content="x")),
self.worker.system_prompt,
)

@patch("planai.llm_task.LLMTaskWorker.publish_work")
def test_max_tool_rounds_forwarded_only_with_tools(self, mock_publish_work):
mock_tool = MagicMock(spec=LLMToolInstance)
Expand Down
Loading