From 2095fc91076bb9b68acc48b5958c6d82b863ea1b Mon Sep 17 00:00:00 2001 From: jmorganca Date: Sat, 23 Nov 2024 19:02:28 -0800 Subject: [PATCH 1/5] make subscription methods more consistent with maps --- ollama/_types.py | 38 ++++++++++++++++++++++++++++++++++++-- 1 file changed, 36 insertions(+), 2 deletions(-) diff --git a/ollama/_types.py b/ollama/_types.py index 5be4850..2c9f3cb 100644 --- a/ollama/_types.py +++ b/ollama/_types.py @@ -17,9 +17,30 @@ from pydantic import ( class SubscriptableBaseModel(BaseModel): def __getitem__(self, key: str) -> Any: - return getattr(self, key) + """ + >>> msg = Message(role='user') + >>> msg['role'] + 'user' + >>> tool = Tool() + >>> tool['type'] + 'function' + >>> msg = Message(role='user') + >>> msg['nonexistent'] + Traceback (most recent call last): + KeyError: 'nonexistent' + """ + if key in self: + return getattr(self, key) + + raise KeyError(key) def __setitem__(self, key: str, value: Any) -> None: + """ + >>> msg = Message(role='user') + >>> msg['role'] = 'assistant' + >>> msg['role'] + 'assistant' + """ setattr(self, key, value) def __contains__(self, key: str) -> bool: @@ -61,7 +82,20 @@ class SubscriptableBaseModel(BaseModel): return False def get(self, key: str, default: Any = None) -> Any: - return getattr(self, key, default) + """ + >>> msg = Message(role='user') + >>> msg.get('role') + 'user' + >>> tool = Tool() + >>> tool.get('type') + 'function' + >>> msg = Message(role='user') + >>> msg.get('nonexistent') + >>> msg = Message(role='user') + >>> msg.get('nonexistent', 'default') + 'default' + """ + return self[key] if key in self else default class Options(SubscriptableBaseModel): From ea0e0dc692c7017cb0ebbd1e9275732d68b7d38a Mon Sep 17 00:00:00 2001 From: Jeffrey Morgan Date: Tue, 26 Nov 2024 10:35:57 -0800 Subject: [PATCH 2/5] Update ollama/_types.py Co-authored-by: Parth Sareen --- ollama/_types.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/ollama/_types.py b/ollama/_types.py index 2c9f3cb..acf44de 100644 --- a/ollama/_types.py +++ b/ollama/_types.py @@ -40,6 +40,11 @@ class SubscriptableBaseModel(BaseModel): >>> msg['role'] = 'assistant' >>> msg['role'] 'assistant' + >>> tool_call = Message.ToolCall(function=Message.ToolCall.Function(name='foo', arguments={})) + >>> msg = Message(role='user', content='hello') + >>> msg['tool_calls'] = [tool_call] + >>> msg['tool_calls'][0]['function']['name'] + 'foo' """ setattr(self, key, value) From ec2c8fdd8d98d80d9928a0b40197fedabde8a0cd Mon Sep 17 00:00:00 2001 From: Jeffrey Morgan Date: Tue, 26 Nov 2024 10:41:45 -0800 Subject: [PATCH 3/5] Update ollama/_types.py Co-authored-by: Parth Sareen --- ollama/_types.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/ollama/_types.py b/ollama/_types.py index acf44de..40dac38 100644 --- a/ollama/_types.py +++ b/ollama/_types.py @@ -21,9 +21,6 @@ class SubscriptableBaseModel(BaseModel): >>> msg = Message(role='user') >>> msg['role'] 'user' - >>> tool = Tool() - >>> tool['type'] - 'function' >>> msg = Message(role='user') >>> msg['nonexistent'] Traceback (most recent call last): From d8d98e17b28c1bb747a04f6c0fe0dc2dfda5ac8a Mon Sep 17 00:00:00 2001 From: Jeffrey Morgan Date: Tue, 26 Nov 2024 10:41:50 -0800 Subject: [PATCH 4/5] Update ollama/_types.py Co-authored-by: Parth Sareen --- ollama/_types.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/ollama/_types.py b/ollama/_types.py index 40dac38..93d898f 100644 --- a/ollama/_types.py +++ b/ollama/_types.py @@ -88,9 +88,6 @@ class SubscriptableBaseModel(BaseModel): >>> msg = Message(role='user') >>> msg.get('role') 'user' - >>> tool = Tool() - >>> tool.get('type') - 'function' >>> msg = Message(role='user') >>> msg.get('nonexistent') >>> msg = Message(role='user') From d4c38978d17082d81a90adc7e7b6043bd5aae7ce Mon Sep 17 00:00:00 2001 From: Jeffrey Morgan Date: Tue, 26 Nov 2024 10:41:53 -0800 Subject: [PATCH 5/5] Update ollama/_types.py Co-authored-by: Parth Sareen --- ollama/_types.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/ollama/_types.py b/ollama/_types.py index 93d898f..293bfa8 100644 --- a/ollama/_types.py +++ b/ollama/_types.py @@ -93,6 +93,9 @@ class SubscriptableBaseModel(BaseModel): >>> msg = Message(role='user') >>> msg.get('nonexistent', 'default') 'default' + >>> msg = Message(role='user', tool_calls=[ Message.ToolCall(function=Message.ToolCall.Function(name='foo', arguments={}))]) + >>> msg.get('tool_calls')[0]['function']['name'] + 'foo' """ return self[key] if key in self else default