Merge leading system messages for OpenAI-compatible endpoints
This commit is contained in:
parent
629cf665a6
commit
55351a8829
2 changed files with 52 additions and 0 deletions
|
|
@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Any, cast
|
||||||
from packaging.version import Version, parse
|
from packaging.version import Version, parse
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from pydantic_ai.profiles.openai import OpenAIModelProfile
|
||||||
from rich.console import RenderableType
|
from rich.console import RenderableType
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
|
|
@ -75,6 +76,16 @@ def apply_common_settings(
|
||||||
return settings_dict
|
return settings_dict
|
||||||
|
|
||||||
|
|
||||||
|
# Strict OpenAI-compatible backends (some vLLM chat templates, e.g. Qwen's)
|
||||||
|
# reject more than one leading system message. Instructions from multiple
|
||||||
|
# sources (agent preamble, capability instructions, dynamic notices) map to
|
||||||
|
# one system message each, so have pydantic-ai merge them for any endpoint
|
||||||
|
# that is not api.openai.com. Harmless on backends that allow multiples.
|
||||||
|
_OPENAI_COMPAT_PROFILE: "OpenAIModelProfile" = {
|
||||||
|
"openai_chat_supports_multiple_system_messages": False
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def get_model(
|
def get_model(
|
||||||
model_config: "ModelConfig",
|
model_config: "ModelConfig",
|
||||||
app_config: "AppConfig | None" = None,
|
app_config: "AppConfig | None" = None,
|
||||||
|
|
@ -125,6 +136,7 @@ def get_model(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OllamaProvider(base_url=base_url),
|
provider=OllamaProvider(base_url=base_url),
|
||||||
settings=model_settings,
|
settings=model_settings,
|
||||||
|
profile=_OPENAI_COMPAT_PROFILE,
|
||||||
)
|
)
|
||||||
|
|
||||||
elif provider == "openai":
|
elif provider == "openai":
|
||||||
|
|
@ -154,6 +166,7 @@ def get_model(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OpenAIProvider(base_url=model_config.base_url),
|
provider=OpenAIProvider(base_url=model_config.base_url),
|
||||||
settings=openai_settings,
|
settings=openai_settings,
|
||||||
|
profile=_OPENAI_COMPAT_PROFILE,
|
||||||
)
|
)
|
||||||
|
|
||||||
return OpenAIChatModel(model_name=model, settings=openai_settings)
|
return OpenAIChatModel(model_name=model, settings=openai_settings)
|
||||||
|
|
|
||||||
|
|
@ -229,6 +229,45 @@ def test_get_model_openai_extra_body_forwarded():
|
||||||
assert result._settings.get("extra_body") == extra
|
assert result._settings.get("extra_body") == extra
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_model_merges_system_messages_for_openai_compatible():
|
||||||
|
"""Instruction parts from multiple sources (agent preamble, capability
|
||||||
|
instructions, dynamic notices) each map to their own system message.
|
||||||
|
Strict OpenAI-compatible templates (e.g. Qwen on vLLM) reject more than
|
||||||
|
one leading system message, so OpenAI-compatible endpoints merge them;
|
||||||
|
plain OpenAI keeps them separate."""
|
||||||
|
from pydantic_ai.messages import InstructionPart, ModelRequest, UserPromptPart
|
||||||
|
from pydantic_ai.models import ModelRequestParameters
|
||||||
|
|
||||||
|
parts = [
|
||||||
|
InstructionPart(content="Base instructions."),
|
||||||
|
InstructionPart(content="Limit notice.", dynamic=True),
|
||||||
|
]
|
||||||
|
messages = [ModelRequest(parts=[UserPromptPart("hi")])]
|
||||||
|
|
||||||
|
async def mapped_for(model):
|
||||||
|
return await model._map_messages(
|
||||||
|
messages, ModelRequestParameters(instruction_parts=parts)
|
||||||
|
)
|
||||||
|
|
||||||
|
vllm = get_model(
|
||||||
|
ModelConfig(provider="openai", name="qwen3.6", base_url="http://vllm:1/v1")
|
||||||
|
)
|
||||||
|
mapped = await mapped_for(vllm)
|
||||||
|
assert [m["role"] for m in mapped] == ["system", "user"]
|
||||||
|
assert mapped[0]["content"] == "Base instructions.\n\nLimit notice."
|
||||||
|
|
||||||
|
ollama = get_model(ModelConfig(provider="ollama", name="llama3"))
|
||||||
|
assert [m["role"] for m in await mapped_for(ollama)] == ["system", "user"]
|
||||||
|
|
||||||
|
openai_native = get_model(ModelConfig(provider="openai", name="gpt-4o"))
|
||||||
|
assert [m["role"] for m in await mapped_for(openai_native)] == [
|
||||||
|
"system",
|
||||||
|
"system",
|
||||||
|
"user",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_ollama_extra_body_forwarded():
|
def test_get_model_ollama_extra_body_forwarded():
|
||||||
"""`extra_body` is forwarded through the Ollama (openai-compatible) branch too."""
|
"""`extra_body` is forwarded through the Ollama (openai-compatible) branch too."""
|
||||||
extra = {"chat_template_kwargs": {"enable_thinking": False}}
|
extra = {"chat_template_kwargs": {"enable_thinking": False}}
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue