diff --git a/docs-astro/src/content/docs/api/taskworker.md b/docs-astro/src/content/docs/api/taskworker.md index 609509e..73519a8 100644 --- a/docs-astro/src/content/docs/api/taskworker.md +++ b/docs-astro/src/content/docs/api/taskworker.md @@ -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]]: diff --git a/docs-astro/src/content/docs/features/llm-integration.md b/docs-astro/src/content/docs/features/llm-integration.md index 9e33da0..5380768 100644 --- a/docs-astro/src/content/docs/features/llm-integration.md +++ b/docs-astro/src/content/docs/features/llm-integration.md @@ -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: diff --git a/pyproject.toml b/pyproject.toml index 5e31ba6..b8e2a4d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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 "] license = "Apache-2.0" diff --git a/src/planai/_version.py b/src/planai/_version.py index f4e9355..a33c871 100644 --- a/src/planai/_version.py +++ b/src/planai/_version.py @@ -14,4 +14,4 @@ """Version information for PlanAI.""" -__version__ = "0.7.0" +__version__ = "0.7.1" diff --git a/src/planai/llm_task.py b/src/planai/llm_task.py index 2387407..c224909 100644 --- a/src/planai/llm_task.py +++ b/src/planai/llm_task.py @@ -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, @@ -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=( @@ -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( @@ -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() diff --git a/tests/planai/test_llm_task.py b/tests/planai/test_llm_task.py index 70e296b..178440d 100644 --- a/tests/planai/test_llm_task.py +++ b/tests/planai/test_llm_task.py @@ -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)