Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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 or "") + "\n\n# Reference\n" + load_reference(task)
```

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"
34 changes: 30 additions & 4 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 All @@ -219,10 +234,15 @@ def extra_validation_with_task(response: BaseModel):
assert isinstance(response, Task) or response is None
self.post_process(response=response, input_task=task)

def get_full_prompt(self, task: Task) -> str:
def get_full_prompt(self, task: Task, system_prompt: Optional[str] = None) -> str:
"""The formatted prompt for ``task``. ``system_prompt`` lets a caller that
has already evaluated ``get_system_prompt(task)`` pass it in, so the hook
runs once per task."""
task_prompt = self.format_prompt(task)

processed_task = self.pre_process(task)
if system_prompt is None:
system_prompt = self.get_system_prompt(task)

return self.llm.generate_full_prompt(
prompt_template=(
Expand All @@ -233,7 +253,7 @@ def get_full_prompt(self, task: Task) -> str:
else ""
)
),
system=self.system_prompt,
system=system_prompt,
task=self._format_task(processed_task),
instructions=task_prompt,
format_instructions=LLMInterface.get_format_instructions(
Expand Down Expand Up @@ -338,5 +358,11 @@ 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}"
# generate_full_prompt() does not fold the system prompt into its result,
# so the key carries it explicitly; the hook is evaluated once here
system_prompt = self.get_system_prompt(task)
full_prompt = self.get_full_prompt(task, system_prompt=system_prompt)
upstream_cache_key += (
f" - {system_prompt} - {full_prompt} - {self.llm.model_name}"
)
return hashlib.sha1(upstream_cache_key.encode()).hexdigest()
45 changes: 45 additions & 0 deletions tests/planai/test_llm_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,51 @@ 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_cache_key_evaluates_the_system_prompt_hook_once(self):
from planai.llm_task import CachedLLMTaskWorker

class Worker(CachedLLMTaskWorker):
output_types: list = [DummyOutputTask]
calls: int = 0

def get_system_prompt(self, task):
self.calls += 1
return "per task"

import tempfile

with tempfile.TemporaryDirectory() as cache_dir:
worker = Worker(llm=self.llm, prompt="p", cache_dir=cache_dir)
key = worker._get_cache_key(DummyTask(content="x"))
self.assertEqual(worker.calls, 1)
self.assertTrue(key)

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