diff --git a/ollama/_types.py b/ollama/_types.py index 96529d6..1278b54 100644 --- a/ollama/_types.py +++ b/ollama/_types.py @@ -98,7 +98,7 @@ def get(self, key: str, default: Any = None) -> Any: >>> msg.get('tool_calls')[0]['function']['name'] 'foo' """ - return getattr(self, key) if hasattr(self, key) else default + return getattr(self, key) if key in self else default class Options(SubscriptableBaseModel): diff --git a/tests/test_type_serialization.py b/tests/test_type_serialization.py index f458cd2..fd41622 100644 --- a/tests/test_type_serialization.py +++ b/tests/test_type_serialization.py @@ -4,7 +4,19 @@ import pytest -from ollama._types import CreateRequest, Image +from ollama._types import CreateRequest, Image, Message + + +def test_subscriptable_model_get_returns_default_for_unset_field(): + message = Message(role='user') + + assert 'content' not in message + assert message.get('content', 'fallback') == 'fallback' + + message.content = None + + assert 'content' in message + assert message.get('content', 'fallback') is None def test_image_serialization_bytes():