diff --git a/pyproject.toml b/pyproject.toml index 613538c..b360aa9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -13,7 +13,7 @@ dependencies = [ 'defusedxml>=0.7', 'httpx>=0.28', 'json-repair>=0.55', - 'jsonschema>=4', + 'jsonschema>=4.18', 'odfdo>=3.22.8', 'orjson>=3', 'pandas[excel]>=2.1', diff --git a/src/graphon/dsl/slim/package_loader.py b/src/graphon/dsl/slim/package_loader.py index fc9359c..a7aee1f 100644 --- a/src/graphon/dsl/slim/package_loader.py +++ b/src/graphon/dsl/slim/package_loader.py @@ -284,9 +284,10 @@ def _convert_model_entity(self, raw_model: dict[str, Any]) -> AIModelEntity | No if model_type is None or fetch_from is None: return None + raw_features = raw_model.get("features") features = [ feature - for item in raw_model.get("features", []) or [] + for item in raw_features or [] if (feature := self._convert_model_feature(item)) is not None ] model_properties = { @@ -299,7 +300,7 @@ def _convert_model_entity(self, raw_model: dict[str, Any]) -> AIModelEntity | No model=str(raw_model["model"]), label=self._convert_i18n(raw_model.get("label")), model_type=model_type, - features=features or None, + features=features if raw_features is not None else None, fetch_from=fetch_from, model_properties=model_properties, deprecated=bool(raw_model.get("deprecated")), diff --git a/src/graphon/model_runtime/entities/model_entities.py b/src/graphon/model_runtime/entities/model_entities.py index ccda57c..add0c93 100644 --- a/src/graphon/model_runtime/entities/model_entities.py +++ b/src/graphon/model_runtime/entities/model_entities.py @@ -200,12 +200,11 @@ def validate_model(self) -> Self: ), None, ) - if not schema_key: + # Explicit feature lists are authoritative; infer support only for legacy + # model schemas that omit the feature declaration. + if not schema_key or self.features is not None: return self - if self.features is None: - self.features = [ModelFeature.STRUCTURED_OUTPUT] - elif ModelFeature.STRUCTURED_OUTPUT not in self.features: - self.features.append(ModelFeature.STRUCTURED_OUTPUT) + self.features = [ModelFeature.STRUCTURED_OUTPUT] return self diff --git a/src/graphon/nodes/llm/entities.py b/src/graphon/nodes/llm/entities.py index 18e233b..133e18d 100644 --- a/src/graphon/nodes/llm/entities.py +++ b/src/graphon/nodes/llm/entities.py @@ -1,7 +1,7 @@ from collections.abc import Mapping, Sequence from typing import Any, Literal -from pydantic import BaseModel, Field, field_validator +from pydantic import BaseModel, Field, field_validator, model_validator from graphon.entities.base_node_data import BaseNodeData from graphon.enums import BuiltinNodeTypes, NodeType @@ -75,8 +75,7 @@ class LLMNodeData(BaseNodeData): context: ContextConfig vision: VisionConfig = Field(default_factory=VisionConfig) structured_output: Mapping[str, Any] | None = None - # We used 'structured_output_enabled' in the past, but it's not a good name. - structured_output_switch_on: bool = Field(False, alias="structured_output_enabled") + structured_output_switch_on: bool = False reasoning_format: Literal["separated", "tagged"] = Field( # Keep tagged as default for backward compatibility default="tagged", @@ -96,6 +95,16 @@ class LLMNodeData(BaseNodeData): ), ) + @model_validator(mode="before") + @classmethod + def migrate_legacy_structured_output_switch(cls, data: Any) -> Any: + if not isinstance(data, Mapping) or "structured_output_enabled" not in data: + return data + data = dict(data) + legacy_value = data.pop("structured_output_enabled") + data.setdefault("structured_output_switch_on", legacy_value) + return data + @field_validator("prompt_config", mode="before") @classmethod def convert_none_prompt_config(cls, v: Any) -> Any: diff --git a/src/graphon/nodes/llm/node.py b/src/graphon/nodes/llm/node.py index 4754f05..e69f157 100644 --- a/src/graphon/nodes/llm/node.py +++ b/src/graphon/nodes/llm/node.py @@ -11,6 +11,10 @@ from datetime import UTC, datetime, timedelta from typing import Any, Literal, assert_never, override +from jsonschema import Draft7Validator, SchemaError, ValidationError +from referencing import Registry +from referencing.exceptions import Unresolvable + from graphon.entities.graph_init_params import GraphInitParams from graphon.enums import ( BuiltinNodeTypes, @@ -39,6 +43,7 @@ PromptMessageContentUnionTypes, TextPromptMessageContent, ) +from graphon.model_runtime.entities.model_entities import ModelFeature from graphon.model_runtime.memory.prompt_message_memory import PromptMessageMemory from graphon.model_runtime.utils.encoders import jsonable_encoder from graphon.node_events.base import ( @@ -101,6 +106,7 @@ PromptMessageContentType.AUDIO: FileType.AUDIO, PromptMessageContentType.DOCUMENT: FileType.DOCUMENT, } +_NO_REMOTE_SCHEMA_REGISTRY = Registry() @dataclass(frozen=True) @@ -262,6 +268,18 @@ def _prepare_run_prompt( node_inputs=node_inputs, ) model_instance = self._prepare_model_instance() + if self.node_data.structured_output_enabled: + model_schema = llm_utils.fetch_model_schema(model_instance=model_instance) + if ( + model_schema.features is not None + and ModelFeature.STRUCTURED_OUTPUT not in model_schema.features + ): + msg = ( + "Structured output is not supported by model " + f"{model_instance.provider}/{model_instance.model_name} " + "(stage=capability)" + ) + raise LLMNodeError(msg) node_inputs.update( llm_utils.build_model_identity_inputs(model_instance=model_instance), ) @@ -370,7 +388,7 @@ def _yield_run_completion( raise LLMNodeError(msg) completed_event = event - if completed_event.structured_output: + if completed_event.structured_output is not None: structured_output = LLMStructuredOutput( structured_output=completed_event.structured_output, ) @@ -518,6 +536,7 @@ def _invoke_llm_with_polling( model_instance=self._model_instance, reasoning_format=self.node_data.reasoning_format, request_start_time=request_start_time, + json_schema=json_schema, ) return case LLMPollingStatus.FAILED: @@ -726,6 +745,7 @@ def invoke_llm( model_parameters = model_instance.parameters invoke_model_parameters = dict(model_parameters) invoke_result: LLMResult | Generator[LLMResultChunk, None, None] + output_schema: dict[str, Any] | None = None if structured_output_enabled: output_schema = LLMNode.fetch_structured_output_schema( structured_output=structured_output or {}, @@ -758,6 +778,7 @@ def invoke_llm( model_instance=model_instance, reasoning_format=reasoning_format, request_start_time=request_start_time, + json_schema=output_schema, ) @staticmethod @@ -771,6 +792,7 @@ def handle_invoke_result( model_instance: LLMProtocol, reasoning_format: Literal["separated", "tagged"] = "tagged", request_start_time: float | None = None, + json_schema: Mapping[str, Any] | None = None, ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: if isinstance(invoke_result, LLMResult): yield from LLMNode._yield_blocking_invoke_result( @@ -779,6 +801,7 @@ def handle_invoke_result( file_outputs=file_outputs, reasoning_format=reasoning_format, request_start_time=request_start_time, + json_schema=json_schema, ) return @@ -790,6 +813,7 @@ def handle_invoke_result( model_instance=model_instance, reasoning_format=reasoning_format, request_start_time=request_start_time, + json_schema=json_schema, ) @staticmethod @@ -800,19 +824,25 @@ def _yield_blocking_invoke_result( file_outputs: list[File], reasoning_format: Literal["separated", "tagged"] = "tagged", request_start_time: float | None = None, + json_schema: Mapping[str, Any] | None = None, ) -> Generator[ModelInvokeCompletedEvent, None, None]: duration = None if request_start_time is not None: duration = time.perf_counter() - request_start_time invoke_result.usage.latency = round(duration, 3) - yield LLMNode.handle_blocking_result( + event = LLMNode.handle_blocking_result( invoke_result=invoke_result, saver=file_saver, file_outputs=file_outputs, reasoning_format=reasoning_format, request_latency=duration, ) + LLMNode._validate_structured_output_result( + structured_output=event.structured_output, + json_schema=json_schema, + ) + yield event @staticmethod def _yield_streaming_invoke_result( @@ -824,6 +854,7 @@ def _yield_streaming_invoke_result( model_instance: LLMProtocol, reasoning_format: Literal["separated", "tagged"] = "tagged", request_start_time: float | None = None, + json_schema: Mapping[str, Any] | None = None, ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: start_time = ( request_start_time @@ -847,10 +878,10 @@ def _yield_streaming_invoke_result( model_instance=model_instance, error=e, ): - msg = f"Failed to parse structured output: {e}" + msg = f"Failed to parse structured output (stage=result, path=$): {e}" raise LLMNodeError(msg) from e if type(e).__name__ == "OutputParserError": - msg = f"Failed to parse structured output: {e}" + msg = f"Failed to parse structured output (stage=result, path=$): {e}" raise LLMNodeError(msg) from e raise @@ -888,6 +919,10 @@ def _yield_streaming_invoke_result( first_token_time=state.first_token_time, start_time=state.start_time, ) + LLMNode._validate_structured_output_result( + structured_output=state.structured_output, + json_schema=json_schema, + ) yield ModelInvokeCompletedEvent( # Use clean_text for separated mode, full_text for tagged mode @@ -1051,6 +1086,38 @@ def _is_structured_output_parse_error( and is_structured_output_parse_error(error) ) + @staticmethod + def _validate_structured_output_result( + *, + structured_output: Mapping[str, Any] | None, + json_schema: Mapping[str, Any] | None, + ) -> None: + if json_schema is None: + return + if structured_output is None: + msg = ( + "Structured output validation failed " + "(stage=result, path=$): structured output is missing" + ) + raise LLMNodeError(msg) + try: + Draft7Validator( + json_schema, + registry=_NO_REMOTE_SCHEMA_REGISTRY, + ).validate(structured_output) + except ValidationError as error: + msg = ( + "Structured output validation failed " + f"(stage=result, path={error.json_path}): {error.message}" + ) + raise LLMNodeError(msg) from error + except Unresolvable as error: + msg = ( + "Structured output validation failed " + "(stage=result, path=$): schema reference could not be resolved" + ) + raise LLMNodeError(msg) from error + @staticmethod def _finalize_streaming_usage( *, @@ -1602,27 +1669,23 @@ def fetch_structured_output_schema( or not a JSON object. """ - if not structured_output: - msg = "Please provide a valid structured output schema" - raise LLMNodeError(msg) - structured_output_schema = json.dumps( - structured_output.get("schema", {}), - ensure_ascii=False, - ) - if not structured_output_schema: - msg = "Please provide a valid structured output schema" + raw_schema = structured_output.get("schema") + if not isinstance(raw_schema, Mapping): + msg = ( + "Invalid structured output schema " + "(stage=schema, path=$.schema): expected a JSON object" + ) raise LLMNodeError(msg) - + schema = dict(raw_schema) try: - schema = json.loads(structured_output_schema) - if not isinstance(schema, dict): - msg = "structured_output_schema must be a JSON object" - raise LLMNodeError(msg) - except json.JSONDecodeError as error: - msg = "structured_output_schema is not valid JSON format" + Draft7Validator.check_schema(schema) + except SchemaError as error: + msg = ( + "Invalid structured output schema " + f"(stage=schema, path={error.json_path}): {error.message}" + ) raise LLMNodeError(msg) from error - else: - return schema + return schema @staticmethod def _save_multimodal_output_and_convert_result_to_markdown( diff --git a/tests/dsl/test_slim_llm.py b/tests/dsl/test_slim_llm.py index e6b4dec..3d7ae1b 100644 --- a/tests/dsl/test_slim_llm.py +++ b/tests/dsl/test_slim_llm.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json from collections.abc import Iterable, Mapping from pathlib import Path from typing import Any @@ -26,14 +27,16 @@ def invoke_chunks( self.calls.append((plugin_id, action, data)) if action == "get_llm_num_tokens": return [{"num_tokens": 7}] - return [ - { - "delta": { - "index": 0, - "message": {"content": "hello"}, - } - } - ] + chunk = { + "delta": { + "index": 0, + "message": {"content": "hello"}, + }, + } + model_parameters = data.get("model_parameters") + if isinstance(model_parameters, Mapping) and "json_schema" in model_parameters: + chunk["structured_output"] = {"ok": True} + return [chunk] class _FailingSlimClient: @@ -179,6 +182,29 @@ def test_slim_llm_counts_tokens_and_collects_blocking_result( } +def test_slim_llm_passes_merged_parameters_and_json_schema( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + client = _patch_recording_slim_client(monkeypatch) + llm = _build_llm(tmp_path) + schema = {"type": "object", "required": ["ok"]} + + result = llm.invoke_llm_with_structured_output( + prompt_messages=[], + json_schema=schema, + model_parameters={"max_tokens": 8}, + stop=None, + stream=False, + ) + + assert result.structured_output == {"ok": True} + model_parameters = client.calls[-1][2]["model_parameters"] + assert model_parameters["temperature"] == pytest.approx(0.2) + assert model_parameters["max_tokens"] == 8 + assert json.loads(model_parameters["json_schema"]) == schema + + def test_slim_llm_preserves_slim_client_errors( monkeypatch: pytest.MonkeyPatch, tmp_path: Path, diff --git a/tests/dsl/test_slim_package_loader.py b/tests/dsl/test_slim_package_loader.py index 89d0b6c..260c036 100644 --- a/tests/dsl/test_slim_package_loader.py +++ b/tests/dsl/test_slim_package_loader.py @@ -60,6 +60,51 @@ def test_slim_config_auto_discovers_uv_and_python( assert config.local.uv_path == "/usr/local/bin/uv" +@pytest.mark.parametrize( + ("raw_features", "infer_from_rule", "expected_features"), + [ + (None, False, None), + ([], True, []), + (None, True, [ModelFeature.STRUCTURED_OUTPUT]), + ( + ["structured-output"], + False, + [ModelFeature.STRUCTURED_OUTPUT], + ), + ], + ids=["unknown", "unsupported", "inferred", "supported"], +) +def test_slim_package_loader_preserves_model_feature_tri_state( + tmp_path: Path, + raw_features: list[str] | None, + infer_from_rule: bool, + expected_features: list[ModelFeature] | None, +) -> None: + loader = SlimPackageLoader( + SlimConfig( + bindings=[SlimProviderBinding(plugin_id="author/fake:0.0.1@test")], + local=SlimLocalSettings(folder=tmp_path), + ), + ) + raw_model = { + "model": "chat-model", + "label": {"en_US": "Chat Model"}, + "model_type": "llm", + "fetch_from": "predefined-model", + "model_properties": {}, + "parameter_rules": ( + [{"name": "json_schema", "type": "string"}] if infer_from_rule else [] + ), + } + if raw_features is not None: + raw_model["features"] = raw_features + + model = loader.convert_model_entity(raw_model) + + assert model is not None + assert model.features == expected_features + + def _write_multi_provider_plugin(plugin_root: Path) -> None: (plugin_root / "_assets").mkdir(parents=True, exist_ok=True) (plugin_root / "provider").mkdir(parents=True, exist_ok=True) diff --git a/tests/nodes/llm/test_node.py b/tests/nodes/llm/test_node.py index 40289b0..11aca21 100644 --- a/tests/nodes/llm/test_node.py +++ b/tests/nodes/llm/test_node.py @@ -8,6 +8,7 @@ import pytest +from graphon.entities.base_node_data import BaseNodeData from graphon.enums import WorkflowNodeExecutionStatus from graphon.file import helpers as file_helpers from graphon.file.enums import FileTransferMethod, FileType @@ -24,6 +25,7 @@ LLMResultChunk, LLMResultChunkDelta, LLMResultChunkWithStructuredOutput, + LLMResultWithStructuredOutput, LLMStructuredOutput, LLMUsage, ) @@ -188,6 +190,229 @@ def _stub_simple_prompt(monkeypatch: pytest.MonkeyPatch, node: LLMNode) -> None: ) +@pytest.mark.parametrize( + "switch_values", + [ + {"structured_output_enabled": True}, + {"structured_output_switch_on": True}, + { + "structured_output_enabled": False, + "structured_output_switch_on": True, + }, + ], + ids=["legacy", "current", "current-wins-conflict"], +) +def test_structured_output_switch_survives_node_data_round_trip( + switch_values: dict[str, bool], +) -> None: + payload = { + "type": "llm", + "model": { + "provider": "openai", + "name": "gpt-4o", + "mode": "chat", + }, + "prompt_template": [{"role": "user", "text": "Hello"}], + "context": {"enabled": False}, + "structured_output": {"schema": {"type": "object"}}, + **switch_values, + } + + node_data = LLMNode.validate_node_data(BaseNodeData.model_validate(payload)) + restored = LLMNode.validate_node_data(node_data) + dumped = restored.model_dump(mode="python", by_alias=True) + restored_from_dump = LLMNode.validate_node_data(dumped) + + assert restored.structured_output_switch_on is True + assert restored.structured_output_enabled is True + assert dumped["structured_output_switch_on"] is True + assert "structured_output_enabled" not in dumped + assert restored_from_dump.structured_output_switch_on is True + + +def test_fetch_structured_output_schema_checks_draft7_schema() -> None: + with pytest.raises(LLMNodeError) as exc_info: + LLMNode.fetch_structured_output_schema( + structured_output={"schema": {"type": "not-a-json-schema-type"}}, + ) + + assert "stage=schema" in str(exc_info.value) + assert "path=$.type" in str(exc_info.value) + + +def test_final_structured_output_validation_disables_remote_ref_retrieval( + monkeypatch: pytest.MonkeyPatch, +) -> None: + def fail_urlopen(*_args: Any, **_kwargs: Any) -> None: + pytest.fail("remote schema retrieval was attempted") + + monkeypatch.setattr("urllib.request.urlopen", fail_urlopen) + + with pytest.raises(LLMNodeError) as exc_info: + LLMNode._validate_structured_output_result( + structured_output={}, + json_schema={"$ref": "http://127.0.0.1/schema"}, + ) + + assert "stage=result" in str(exc_info.value) + assert "path=$" in str(exc_info.value) + assert "schema reference could not be resolved" in str(exc_info.value) + + +def test_final_structured_output_validation_supports_local_refs() -> None: + schema = { + "definitions": { + "result": { + "type": "object", + "properties": {"ok": {"type": "boolean"}}, + "required": ["ok"], + }, + }, + "$ref": "#/definitions/result", + } + + LLMNode._validate_structured_output_result( + structured_output={"ok": True}, + json_schema=schema, + ) + + +_RESULT_SCHEMA = { + "type": "object", + "properties": { + "profile": { + "type": "object", + "properties": { + "role": {"type": "string", "enum": ["admin"]}, + "scores": {"type": "array", "items": {"type": "integer"}}, + }, + "required": ["role", "scores"], + }, + }, + "required": ["profile"], +} + + +@pytest.mark.parametrize( + ("structured_output", "path"), + [ + (None, "$"), + ({}, "$"), + ({"profile": {}}, "$.profile"), + ({"profile": {"role": "viewer", "scores": [1]}}, "$.profile.role"), + ( + {"profile": {"role": "admin", "scores": [1, "invalid"]}}, + "$.profile.scores[1]", + ), + ], +) +def test_final_structured_output_is_validated_against_full_schema( + structured_output: dict[str, Any] | None, + path: str, +) -> None: + chunk = LLMResultChunkWithStructuredOutput( + model="gpt-4o", + delta=LLMResultChunkDelta( + index=0, + message=AssistantPromptMessage(content=""), + usage=LLMUsage.empty_usage(), + ), + structured_output=structured_output, + ) + model = MagicMock(is_structured_output_parse_error=lambda _error: False) + + with pytest.raises(LLMNodeError) as exc_info: + list( + LLMNode.handle_invoke_result( + invoke_result=_stream_results(chunk), + file_saver=MagicMock(), + file_outputs=[], + node_id="llm", + model_instance=cast(LLMProtocol, model), + json_schema=_RESULT_SCHEMA, + ), + ) + + assert "stage=result" in str(exc_info.value) + assert f"path={path}" in str(exc_info.value) + if structured_output is None: + assert "structured output is missing" in str(exc_info.value) + + +def test_blocking_structured_output_is_validated_against_schema() -> None: + result = LLMResultWithStructuredOutput( + model="gpt-4o", + message=AssistantPromptMessage(content=""), + usage=LLMUsage.empty_usage(), + structured_output={"profile": {"role": "admin", "scores": ["invalid"]}}, + ) + + with pytest.raises(LLMNodeError, match=r"stage=result, path=\$\.profile\.scores"): + list( + LLMNode.handle_invoke_result( + invoke_result=result, + file_saver=MagicMock(), + file_outputs=[], + node_id="llm", + model_instance=cast(LLMProtocol, MagicMock()), + json_schema=_RESULT_SCHEMA, + ), + ) + + +@pytest.mark.parametrize( + ("features", "should_invoke"), + [ + ([ModelFeature.STRUCTURED_OUTPUT], True), + ([], False), + (None, True), + ], + ids=["supported", "unsupported", "unknown"], +) +def test_structured_output_capability_is_tri_state( + features: list[ModelFeature] | None, + should_invoke: bool, +) -> None: + model = MagicMock( + provider="openai", + model_name="gpt-4o", + parameters={}, + stop=(), + is_structured_output_parse_error=lambda _error: False, + ) + model.get_model_schema.return_value = SimpleNamespace( + features=features, + supports_prompt_content_type=lambda _content_type: True, + ) + model.invoke_llm_with_structured_output.return_value = _stream_results( + LLMResultChunkWithStructuredOutput( + model="gpt-4o", + delta=LLMResultChunkDelta( + index=0, + message=AssistantPromptMessage(content=""), + usage=LLMUsage.empty_usage(), + ), + structured_output={}, + ), + ) + node = _build_llm_node(model_instance=model) + node.node_data.structured_output_switch_on = True + node.node_data.structured_output = {"schema": {"type": "object"}} + + completed = next( + event for event in node._run() if isinstance(event, StreamCompletedEvent) + ) + + if should_invoke: + model.invoke_llm_with_structured_output.assert_called_once() + assert completed.node_run_result.status == WorkflowNodeExecutionStatus.SUCCEEDED + assert completed.node_run_result.outputs["structured_output"] == {} + else: + model.invoke_llm_with_structured_output.assert_not_called() + assert completed.node_run_result.status == WorkflowNodeExecutionStatus.FAILED + assert "stage=capability" in completed.node_run_result.error + + def test_run_emits_model_identity_in_node_result_inputs( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/uv.lock b/uv.lock index 4896f3c..a40a8d6 100644 --- a/uv.lock +++ b/uv.lock @@ -366,7 +366,7 @@ requires-dist = [ { name = "defusedxml", specifier = ">=0.7" }, { name = "httpx", specifier = ">=0.28" }, { name = "json-repair", specifier = ">=0.55" }, - { name = "jsonschema", specifier = ">=4" }, + { name = "jsonschema", specifier = ">=4.18" }, { name = "odfdo", specifier = ">=3.22.8" }, { name = "orjson", specifier = ">=3" }, { name = "pandas", extras = ["excel"], specifier = ">=2.1" },