nora-lib 0.0.1.dev0__tar.gz → 0.0.3__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {nora_lib-0.0.1.dev0 → nora_lib-0.0.3}/PKG-INFO +4 -1
- nora_lib-0.0.3/README.md +37 -0
- nora_lib-0.0.3/nora_lib/context/context_service.py +88 -0
- nora_lib-0.0.3/nora_lib/context/models.py +21 -0
- nora_lib-0.0.3/nora_lib/interactions/__init__.py +0 -0
- nora_lib-0.0.3/nora_lib/interactions/interactions_service.py +117 -0
- nora_lib-0.0.3/nora_lib/interactions/models.py +114 -0
- nora_lib-0.0.3/nora_lib/py.typed +0 -0
- nora_lib-0.0.3/nora_lib/tasks/__init__.py +0 -0
- nora_lib-0.0.3/nora_lib/tasks/models.py +25 -0
- nora_lib-0.0.3/nora_lib/tasks/state.py +53 -0
- {nora_lib-0.0.1.dev0 → nora_lib-0.0.3}/nora_lib.egg-info/PKG-INFO +4 -1
- nora_lib-0.0.3/nora_lib.egg-info/SOURCES.txt +22 -0
- nora_lib-0.0.3/nora_lib.egg-info/requires.txt +8 -0
- {nora_lib-0.0.1.dev0 → nora_lib-0.0.3}/nora_lib.egg-info/top_level.txt +1 -0
- {nora_lib-0.0.1.dev0 → nora_lib-0.0.3}/setup.py +3 -3
- nora_lib-0.0.3/tests/tasks/__init__.py +0 -0
- nora_lib-0.0.3/tests/tasks/test_state.py +77 -0
- nora_lib-0.0.3/tests/test_placeholder.py +6 -0
- nora_lib-0.0.1.dev0/README.md +0 -3
- nora_lib-0.0.1.dev0/nora_lib.egg-info/SOURCES.txt +0 -10
- nora_lib-0.0.1.dev0/nora_lib.egg-info/requires.txt +0 -5
- {nora_lib-0.0.1.dev0 → nora_lib-0.0.3}/nora_lib/__init__.py +0 -0
- /nora_lib-0.0.1.dev0/nora_lib/py.typed → /nora_lib-0.0.3/nora_lib/context/__init__.py +0 -0
- {nora_lib-0.0.1.dev0 → nora_lib-0.0.3}/nora_lib.egg-info/dependency_links.txt +0 -0
- {nora_lib-0.0.1.dev0 → nora_lib-0.0.3}/pyproject.toml +0 -0
- {nora_lib-0.0.1.dev0 → nora_lib-0.0.3}/setup.cfg +0 -0
|
@@ -1,10 +1,13 @@
|
|
|
1
1
|
Metadata-Version: 2.1
|
|
2
2
|
Name: nora_lib
|
|
3
|
-
Version: 0.0.
|
|
3
|
+
Version: 0.0.3
|
|
4
4
|
Summary: For making and coordinating agents and tools
|
|
5
5
|
Home-page: https://github.com/allenai/nora_lib
|
|
6
6
|
Requires-Python: >=3.9
|
|
7
|
+
Requires-Dist: pydantic<3,>=2
|
|
8
|
+
Requires-Dist: requests
|
|
7
9
|
Provides-Extra: dev
|
|
8
10
|
Requires-Dist: mypy; extra == "dev"
|
|
9
11
|
Requires-Dist: pytest; extra == "dev"
|
|
10
12
|
Requires-Dist: black; extra == "dev"
|
|
13
|
+
Requires-Dist: types-requests; extra == "dev"
|
nora_lib-0.0.3/README.md
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
# nora_lib
|
|
2
|
+
|
|
3
|
+
For making and coordinating agents and tools.
|
|
4
|
+
|
|
5
|
+
# Development
|
|
6
|
+
|
|
7
|
+
When preparing a PR, make sure you first run the verify script
|
|
8
|
+
to confirm the code is buildable, passes type- and formatting-checks,
|
|
9
|
+
as well as unit tests.
|
|
10
|
+
|
|
11
|
+
```bash
|
|
12
|
+
cd <project_root>
|
|
13
|
+
./verify.sh
|
|
14
|
+
```
|
|
15
|
+
|
|
16
|
+
The script will tell you what's wrong if anything fails.
|
|
17
|
+
|
|
18
|
+
Update the `version=` field in `setup.py` when you make changes
|
|
19
|
+
as part of your changeset.
|
|
20
|
+
|
|
21
|
+
# Publication
|
|
22
|
+
|
|
23
|
+
After your PR merges to `main` you will need to publish
|
|
24
|
+
the library to public pypi for it to be useable by client applications.
|
|
25
|
+
|
|
26
|
+
This requires one environment variable to be set, which can be found in
|
|
27
|
+
the 1Pass NORA Vault under the secret named "NORA pypi token".
|
|
28
|
+
|
|
29
|
+
```bash
|
|
30
|
+
export AI2_NORA_PYPI_TOKEN=<SECRET IN NORA VAULT>
|
|
31
|
+
cd <project_root>
|
|
32
|
+
git checkout main
|
|
33
|
+
git pull origin main
|
|
34
|
+
git tag v<YOUR_NEW_VERSION>
|
|
35
|
+
git push origin --tags
|
|
36
|
+
./publish.sh
|
|
37
|
+
```
|
|
@@ -0,0 +1,88 @@
|
|
|
1
|
+
from datetime import datetime, timezone
|
|
2
|
+
from typing import List, Optional
|
|
3
|
+
from uuid import UUID
|
|
4
|
+
|
|
5
|
+
from nora_lib.interactions.interactions_service import InteractionsService
|
|
6
|
+
from nora_lib.interactions.models import (
|
|
7
|
+
ReturnedMessage,
|
|
8
|
+
ReturnedAgentContextMessage,
|
|
9
|
+
ReturnedAgentContextEvent,
|
|
10
|
+
EventType,
|
|
11
|
+
AgentMessageData,
|
|
12
|
+
Event,
|
|
13
|
+
)
|
|
14
|
+
from nora_lib.context.models import WrappedTaskObject
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class ContextService:
|
|
18
|
+
"""
|
|
19
|
+
Save and retrieve task agent context from interaction store
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
def __init__(
|
|
23
|
+
self,
|
|
24
|
+
agent_actor_id: str, # uuid representing this agent in interaction store
|
|
25
|
+
interactions_base_url: str,
|
|
26
|
+
interactions_bearer_token: Optional[str],
|
|
27
|
+
timeout: int = 30,
|
|
28
|
+
):
|
|
29
|
+
self.interactions_service = self._get_interactions_service(
|
|
30
|
+
interactions_base_url, interactions_bearer_token, timeout
|
|
31
|
+
)
|
|
32
|
+
self.agent_actor_id = agent_actor_id
|
|
33
|
+
|
|
34
|
+
def _get_interactions_service(self, url, token, timeout) -> InteractionsService:
|
|
35
|
+
return InteractionsService(url, timeout, token)
|
|
36
|
+
|
|
37
|
+
def fetch_context(
|
|
38
|
+
self, request: WrappedTaskObject
|
|
39
|
+
) -> List[ReturnedAgentContextMessage]:
|
|
40
|
+
message_id = request.message_id
|
|
41
|
+
|
|
42
|
+
returned_messages: List[ReturnedMessage] = (
|
|
43
|
+
self.interactions_service.fetch_messages_and_events_for_forked_thread(
|
|
44
|
+
message_id, EventType.AGENT_CONTEXT
|
|
45
|
+
)
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
messages_with_filtered_events: List[ReturnedAgentContextMessage] = []
|
|
49
|
+
for message in returned_messages:
|
|
50
|
+
events_saved_by_this_agent: List[ReturnedAgentContextEvent] = []
|
|
51
|
+
if message.events:
|
|
52
|
+
for event in message.events:
|
|
53
|
+
context_event = ReturnedAgentContextEvent.model_validate(event)
|
|
54
|
+
if context_event.actor_id == self.agent_actor_id:
|
|
55
|
+
events_saved_by_this_agent.append(context_event)
|
|
56
|
+
|
|
57
|
+
events_saved_by_this_agent.sort(
|
|
58
|
+
key=lambda event: datetime.fromisoformat(event.timestamp)
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
updated_message = ReturnedAgentContextMessage(
|
|
62
|
+
message_id=message.message_id,
|
|
63
|
+
actor_id=message.actor_id,
|
|
64
|
+
text=message.text,
|
|
65
|
+
ts=message.ts,
|
|
66
|
+
annotated_text=message.annotated_text,
|
|
67
|
+
events=events_saved_by_this_agent,
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
messages_with_filtered_events.append(updated_message)
|
|
71
|
+
|
|
72
|
+
return messages_with_filtered_events
|
|
73
|
+
|
|
74
|
+
def save_context(self, event_data: WrappedTaskObject):
|
|
75
|
+
agent_data = AgentMessageData(
|
|
76
|
+
message_data=event_data.model_dump(),
|
|
77
|
+
data_sender_actor_id=event_data.sender_actor_id,
|
|
78
|
+
virtual_thread_id=event_data.virtual_thread_id,
|
|
79
|
+
)
|
|
80
|
+
event = Event(
|
|
81
|
+
type=EventType.AGENT_CONTEXT,
|
|
82
|
+
actor_id=UUID(self.agent_actor_id),
|
|
83
|
+
timestamp=datetime.now(timezone.utc),
|
|
84
|
+
data=agent_data.model_dump(),
|
|
85
|
+
message_id=event_data.message_id,
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
self.interactions_service.save_event(event)
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
from typing import Generic, Optional, TypeVar
|
|
2
|
+
from pydantic import BaseModel, Field
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
R = TypeVar("R", bound=BaseModel)
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class WrappedTaskObject(BaseModel, Generic[R]):
|
|
9
|
+
"""Encloses request or response object with additional metadata"""
|
|
10
|
+
|
|
11
|
+
message_id: str = Field(
|
|
12
|
+
description="id of originating message; key for istore retrieval"
|
|
13
|
+
)
|
|
14
|
+
sender_actor_id: str = Field(
|
|
15
|
+
description="string representation of uuid identifying agent sending data"
|
|
16
|
+
)
|
|
17
|
+
virtual_thread_id: Optional[str] = Field(
|
|
18
|
+
description="Tool-defined local thread to associate follow up requests"
|
|
19
|
+
)
|
|
20
|
+
task_id: Optional[str] = Field(description="Reference to a long-running task")
|
|
21
|
+
data: R = Field(description="Tool-defined request or response")
|
|
File without changes
|
|
@@ -0,0 +1,117 @@
|
|
|
1
|
+
from datetime import datetime
|
|
2
|
+
import requests
|
|
3
|
+
from typing import List, Optional
|
|
4
|
+
|
|
5
|
+
from nora_lib.interactions.models import (
|
|
6
|
+
Event,
|
|
7
|
+
EventType,
|
|
8
|
+
ReturnedMessage,
|
|
9
|
+
ThreadRelationsResponse,
|
|
10
|
+
ThreadForkEventData,
|
|
11
|
+
thread_message_lookup_request,
|
|
12
|
+
)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class InteractionsService:
|
|
16
|
+
"""
|
|
17
|
+
Service which saves interactions to the Interactions API
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
def __init__(self, base_url, timeout, token):
|
|
21
|
+
self.base_url = base_url
|
|
22
|
+
self.timeout = timeout
|
|
23
|
+
self.headers = {"Authorization": f"Bearer {token}"}
|
|
24
|
+
|
|
25
|
+
def save_event(self, event: Event) -> None:
|
|
26
|
+
"""Save an event to the Interactions API"""
|
|
27
|
+
event_url = f"{self.base_url}/interaction/v1/event"
|
|
28
|
+
response = requests.post(
|
|
29
|
+
event_url,
|
|
30
|
+
json=event.model_dump(),
|
|
31
|
+
headers=self.headers,
|
|
32
|
+
timeout=int(self.timeout),
|
|
33
|
+
)
|
|
34
|
+
response.raise_for_status()
|
|
35
|
+
|
|
36
|
+
def fetch_thread_messages_and_events_for_message(
|
|
37
|
+
self, message_id: str, event_type: str
|
|
38
|
+
) -> ThreadRelationsResponse:
|
|
39
|
+
"""Fetch messages and associated events from the same thread as provided messagev id"""
|
|
40
|
+
message_url = f"{self.base_url}/interaction/v1/search/message"
|
|
41
|
+
request_body = thread_message_lookup_request(message_id, event_type=event_type)
|
|
42
|
+
response = requests.post(
|
|
43
|
+
message_url,
|
|
44
|
+
json=request_body,
|
|
45
|
+
headers=self.headers,
|
|
46
|
+
timeout=int(self.timeout),
|
|
47
|
+
)
|
|
48
|
+
response.raise_for_status()
|
|
49
|
+
json_response = response.json()
|
|
50
|
+
|
|
51
|
+
return ThreadRelationsResponse.model_validate(
|
|
52
|
+
json_response.get("message", {}).get("thread", {})
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
def fetch_messages_and_events_for_thread(
|
|
56
|
+
self,
|
|
57
|
+
thread_id: str,
|
|
58
|
+
event_type: Optional[str] = None,
|
|
59
|
+
min_timestamp: Optional[str] = None,
|
|
60
|
+
) -> dict:
|
|
61
|
+
"""Fetch messages and events for the thread containing a given message from the Interactions API"""
|
|
62
|
+
THREAD_SEARCH_URL = f"{self.base_url}/interaction/v1/search/thread"
|
|
63
|
+
request_body = {
|
|
64
|
+
"id": thread_id,
|
|
65
|
+
"relations": {
|
|
66
|
+
"messages": (
|
|
67
|
+
{"filter": {"min_timestamp": min_timestamp}}
|
|
68
|
+
if min_timestamp
|
|
69
|
+
else {}
|
|
70
|
+
),
|
|
71
|
+
"events": {"filter": {"type": event_type}} if event_type else {},
|
|
72
|
+
},
|
|
73
|
+
}
|
|
74
|
+
|
|
75
|
+
response = requests.post(
|
|
76
|
+
THREAD_SEARCH_URL,
|
|
77
|
+
json=request_body,
|
|
78
|
+
headers=self.headers,
|
|
79
|
+
timeout=int(self.timeout),
|
|
80
|
+
)
|
|
81
|
+
response.raise_for_status()
|
|
82
|
+
return response.json()
|
|
83
|
+
|
|
84
|
+
def fetch_messages_and_events_for_forked_thread(
|
|
85
|
+
self, message_id: str, event_type: str
|
|
86
|
+
) -> List[ReturnedMessage]:
|
|
87
|
+
"""Build a history of messages for a given message including associated events.
|
|
88
|
+
This includes messages from pre-forked threads."""
|
|
89
|
+
returned_messages: List[ReturnedMessage] = []
|
|
90
|
+
|
|
91
|
+
messages_for_thread: ThreadRelationsResponse = (
|
|
92
|
+
self.fetch_thread_messages_and_events_for_message(message_id, event_type)
|
|
93
|
+
)
|
|
94
|
+
if messages_for_thread.messages:
|
|
95
|
+
returned_messages.extend(messages_for_thread.messages)
|
|
96
|
+
|
|
97
|
+
# Lookup any thread_fork events (conversation across surfaces)
|
|
98
|
+
thread_fork_events = self.fetch_messages_and_events_for_thread(
|
|
99
|
+
messages_for_thread.thread_id, EventType.THREAD_FORK.value
|
|
100
|
+
)
|
|
101
|
+
for forked_thread_event in thread_fork_events.get("thread", {}).get(
|
|
102
|
+
"events", []
|
|
103
|
+
):
|
|
104
|
+
event_data = ThreadForkEventData.model_validate(
|
|
105
|
+
forked_thread_event.get("data", {})
|
|
106
|
+
)
|
|
107
|
+
forked_thread: ThreadRelationsResponse = (
|
|
108
|
+
self.fetch_thread_messages_and_events_for_message(
|
|
109
|
+
event_data.previous_message_id, event_type
|
|
110
|
+
)
|
|
111
|
+
)
|
|
112
|
+
if forked_thread.messages:
|
|
113
|
+
returned_messages.extend(forked_thread.messages)
|
|
114
|
+
|
|
115
|
+
returned_messages.sort(key=lambda x: datetime.fromisoformat(x.ts))
|
|
116
|
+
|
|
117
|
+
return returned_messages
|
|
@@ -0,0 +1,114 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Model for interactions to be sent to the interactions service.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from datetime import datetime
|
|
6
|
+
from enum import Enum
|
|
7
|
+
from typing import Optional, List
|
|
8
|
+
from uuid import UUID
|
|
9
|
+
|
|
10
|
+
from pydantic import BaseModel, Field, field_serializer
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class EventType(str, Enum):
|
|
14
|
+
"""Event types for the interactions service"""
|
|
15
|
+
|
|
16
|
+
AGENT_CONTEXT = "agent:message_context"
|
|
17
|
+
THREAD_FORK = "thread_fork"
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class Event(BaseModel):
|
|
21
|
+
"""event object to be sent to the interactions service; requires association with a message, thread or channel id"""
|
|
22
|
+
|
|
23
|
+
type: str
|
|
24
|
+
actor_id: UUID = Field(
|
|
25
|
+
description="identifies actor writing the event to the interaction service"
|
|
26
|
+
)
|
|
27
|
+
timestamp: datetime
|
|
28
|
+
text: Optional[str] = None
|
|
29
|
+
data: Optional[dict] = Field(default_factory=dict)
|
|
30
|
+
message_id: Optional[str] = None
|
|
31
|
+
thread_id: Optional[str] = None
|
|
32
|
+
channel_id: Optional[str] = None
|
|
33
|
+
|
|
34
|
+
@field_serializer("actor_id")
|
|
35
|
+
def serialize_actor_id(self, actor_id: UUID):
|
|
36
|
+
return str(actor_id)
|
|
37
|
+
|
|
38
|
+
@field_serializer("timestamp")
|
|
39
|
+
def serialize_timestamp(self, timestamp: datetime):
|
|
40
|
+
return timestamp.isoformat()
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
class AgentMessageData(BaseModel):
|
|
44
|
+
"""capture requests to and responses from tools within Events"""
|
|
45
|
+
|
|
46
|
+
message_data: dict # dict of agent/tool request/response format
|
|
47
|
+
data_sender_actor_id: Optional[str] = None # agent sending the data
|
|
48
|
+
virtual_thread_id: Optional[str] = None # tool-provided thread
|
|
49
|
+
tool_call_id: Optional[str] = None # llm-provided thread
|
|
50
|
+
tool_name: Optional[str] = None # llm identifier for tool
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class ReturnedAgentContextEvent(BaseModel):
|
|
54
|
+
"""Event format returned by interaction service for agent context events"""
|
|
55
|
+
|
|
56
|
+
actor_id: str # agent that saved this context
|
|
57
|
+
timestamp: str
|
|
58
|
+
data: AgentMessageData
|
|
59
|
+
type: str
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
class ReturnedAgentContextMessage(BaseModel):
|
|
63
|
+
"""Message format returned by interaction service for search by thread"""
|
|
64
|
+
|
|
65
|
+
message_id: str
|
|
66
|
+
actor_id: str
|
|
67
|
+
text: str
|
|
68
|
+
ts: str
|
|
69
|
+
annotated_text: Optional[str] = None
|
|
70
|
+
events: Optional[List[ReturnedAgentContextEvent]] = None
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
class ThreadForkEventData(BaseModel):
|
|
74
|
+
"""Event data for a thread fork event"""
|
|
75
|
+
|
|
76
|
+
previous_message_id: str
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
class ReturnedMessage(BaseModel):
|
|
80
|
+
"""Message format returned by interaction service for search by thread"""
|
|
81
|
+
|
|
82
|
+
message_id: str
|
|
83
|
+
actor_id: str
|
|
84
|
+
text: str
|
|
85
|
+
ts: str
|
|
86
|
+
annotated_text: Optional[str] = None
|
|
87
|
+
events: Optional[List[dict]] = None
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
class ThreadRelationsResponse(BaseModel):
|
|
91
|
+
"""Thread format returned by interaction service for thread relations in a search response"""
|
|
92
|
+
|
|
93
|
+
thread_id: str
|
|
94
|
+
events: Optional[List[Event]] = None # events associated only with the thread
|
|
95
|
+
messages: Optional[List[ReturnedMessage]] = (
|
|
96
|
+
None # includes events associated with each message
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def thread_message_lookup_request(message_id: str, event_type: str) -> dict:
|
|
101
|
+
"""retrieve messages and events for the thread associated with a message"""
|
|
102
|
+
return {
|
|
103
|
+
"id": message_id,
|
|
104
|
+
"relations": {
|
|
105
|
+
"thread": {
|
|
106
|
+
"relations": {
|
|
107
|
+
"messages": {
|
|
108
|
+
"relations": {"events": {"filter": {"type": event_type}}},
|
|
109
|
+
"apply_annotations_from_actors": ["*"],
|
|
110
|
+
},
|
|
111
|
+
}
|
|
112
|
+
}
|
|
113
|
+
},
|
|
114
|
+
}
|
|
File without changes
|
|
File without changes
|
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
from typing import Any, Dict, Generic, Optional, TypeVar
|
|
2
|
+
from pydantic import BaseModel, Field
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
R = TypeVar("R", bound=BaseModel)
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
TASK_STATUSES = dict(STARTED="STARTED", FAILED="FAILED", COMPLETED="COMPLETED")
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class AsyncTaskState(BaseModel, Generic[R]):
|
|
12
|
+
"""Models the current state of an asynchronous request."""
|
|
13
|
+
|
|
14
|
+
task_id: str = Field(
|
|
15
|
+
"Identifies the long-running task so that its status and eventual result"
|
|
16
|
+
"can be checked in follow-up calls."
|
|
17
|
+
)
|
|
18
|
+
estimated_time: str = Field(
|
|
19
|
+
description="How long we expect this task to take from start to finish."
|
|
20
|
+
)
|
|
21
|
+
task_status: str = Field(description="Current human-readable status of the task.")
|
|
22
|
+
task_result: Optional[R] = Field(description="Final result of the task.")
|
|
23
|
+
extra_state: Dict[str, Any] = Field(
|
|
24
|
+
description="Any extra task-specific state can go in here as free-form JSON-serializable dictionary."
|
|
25
|
+
)
|
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Handles IO for asynchronous task-related state.
|
|
3
|
+
Currently just reads/writes from local disk,
|
|
4
|
+
the best and most robust mechanism.
|
|
5
|
+
|
|
6
|
+
A StateManager instance must be initialized with
|
|
7
|
+
a concrete subclass of `AsyncTaskState`, as implemented
|
|
8
|
+
by dependent projects.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
import json
|
|
12
|
+
import os
|
|
13
|
+
from typing import Generic, Optional, Type
|
|
14
|
+
|
|
15
|
+
from nora_lib.tasks.models import AsyncTaskState, R, TASK_STATUSES
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class NoSuchTaskException(Exception):
|
|
19
|
+
def __init__(self, task_id: str):
|
|
20
|
+
self._task_id = task_id
|
|
21
|
+
|
|
22
|
+
def __str__(self):
|
|
23
|
+
return f"No record found for task {self._task_id}"
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class StateManager(Generic[R]):
|
|
27
|
+
def __init__(self, task_state_class: Type[AsyncTaskState[R]], state_dir) -> None:
|
|
28
|
+
self._task_state_class = task_state_class
|
|
29
|
+
self._state_dir = state_dir
|
|
30
|
+
|
|
31
|
+
def read_state(self, task_id: str) -> AsyncTaskState[R]:
|
|
32
|
+
task_state_path = os.path.join(self._state_dir, f"{task_id}.json")
|
|
33
|
+
if not os.path.isfile(task_state_path):
|
|
34
|
+
raise NoSuchTaskException(task_id)
|
|
35
|
+
|
|
36
|
+
with open(task_state_path, "r") as f:
|
|
37
|
+
return self._task_state_class(**json.loads(f.read()))
|
|
38
|
+
|
|
39
|
+
def write_state(self, state: AsyncTaskState[R]) -> None:
|
|
40
|
+
task_state_path = os.path.join(self._state_dir, f"{state.task_id}.json")
|
|
41
|
+
with open(task_state_path, "w") as f:
|
|
42
|
+
json.dump(state.model_dump(), f)
|
|
43
|
+
|
|
44
|
+
def update_status(self, task_id: str, new_status: str) -> None:
|
|
45
|
+
state = self.read_state(task_id)
|
|
46
|
+
state.task_status = new_status
|
|
47
|
+
self.write_state(state)
|
|
48
|
+
|
|
49
|
+
def save_result(self, task_id: str, task_result: R) -> None:
|
|
50
|
+
state = self.read_state(task_id)
|
|
51
|
+
state.task_status = TASK_STATUSES["COMPLETED"]
|
|
52
|
+
state.task_result = task_result
|
|
53
|
+
self.write_state(state)
|
|
@@ -1,10 +1,13 @@
|
|
|
1
1
|
Metadata-Version: 2.1
|
|
2
2
|
Name: nora_lib
|
|
3
|
-
Version: 0.0.
|
|
3
|
+
Version: 0.0.3
|
|
4
4
|
Summary: For making and coordinating agents and tools
|
|
5
5
|
Home-page: https://github.com/allenai/nora_lib
|
|
6
6
|
Requires-Python: >=3.9
|
|
7
|
+
Requires-Dist: pydantic<3,>=2
|
|
8
|
+
Requires-Dist: requests
|
|
7
9
|
Provides-Extra: dev
|
|
8
10
|
Requires-Dist: mypy; extra == "dev"
|
|
9
11
|
Requires-Dist: pytest; extra == "dev"
|
|
10
12
|
Requires-Dist: black; extra == "dev"
|
|
13
|
+
Requires-Dist: types-requests; extra == "dev"
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
README.md
|
|
2
|
+
pyproject.toml
|
|
3
|
+
setup.py
|
|
4
|
+
nora_lib/__init__.py
|
|
5
|
+
nora_lib/py.typed
|
|
6
|
+
nora_lib.egg-info/PKG-INFO
|
|
7
|
+
nora_lib.egg-info/SOURCES.txt
|
|
8
|
+
nora_lib.egg-info/dependency_links.txt
|
|
9
|
+
nora_lib.egg-info/requires.txt
|
|
10
|
+
nora_lib.egg-info/top_level.txt
|
|
11
|
+
nora_lib/context/__init__.py
|
|
12
|
+
nora_lib/context/context_service.py
|
|
13
|
+
nora_lib/context/models.py
|
|
14
|
+
nora_lib/interactions/__init__.py
|
|
15
|
+
nora_lib/interactions/interactions_service.py
|
|
16
|
+
nora_lib/interactions/models.py
|
|
17
|
+
nora_lib/tasks/__init__.py
|
|
18
|
+
nora_lib/tasks/models.py
|
|
19
|
+
nora_lib/tasks/state.py
|
|
20
|
+
tests/test_placeholder.py
|
|
21
|
+
tests/tasks/__init__.py
|
|
22
|
+
tests/tasks/test_state.py
|
|
@@ -1,13 +1,13 @@
|
|
|
1
1
|
import setuptools
|
|
2
2
|
|
|
3
|
-
runtime_requirements = []
|
|
3
|
+
runtime_requirements = ["pydantic>=2,<3", "requests"]
|
|
4
4
|
|
|
5
5
|
# For running tests, linting, etc
|
|
6
|
-
dev_requirements = ["mypy", "pytest", "black"]
|
|
6
|
+
dev_requirements = ["mypy", "pytest", "black", "types-requests"]
|
|
7
7
|
|
|
8
8
|
setuptools.setup(
|
|
9
9
|
name="nora_lib",
|
|
10
|
-
version="0.0.
|
|
10
|
+
version="0.0.3",
|
|
11
11
|
description="For making and coordinating agents and tools",
|
|
12
12
|
url="https://github.com/allenai/nora_lib",
|
|
13
13
|
packages=setuptools.find_packages(exclude=(["tests"])),
|
|
File without changes
|
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
import tempfile
|
|
2
|
+
import unittest
|
|
3
|
+
|
|
4
|
+
from pydantic import BaseModel, Field
|
|
5
|
+
|
|
6
|
+
from nora_lib.tasks.models import AsyncTaskState, TASK_STATUSES
|
|
7
|
+
from nora_lib.tasks.state import NoSuchTaskException, StateManager
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class MyTaskResult(BaseModel):
|
|
11
|
+
a: str
|
|
12
|
+
b: int
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class MyAsyncTaskState(AsyncTaskState[MyTaskResult]):
|
|
16
|
+
pass
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class TestState(unittest.TestCase):
|
|
20
|
+
def test__can_read_and_write_state(self):
|
|
21
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
22
|
+
manager = StateManager(MyAsyncTaskState, tmpdir)
|
|
23
|
+
state = MyAsyncTaskState(
|
|
24
|
+
task_id="asdf",
|
|
25
|
+
estimated_time="40 days and 40 nights",
|
|
26
|
+
task_status="STARTED",
|
|
27
|
+
task_result=None,
|
|
28
|
+
extra_state={"foo": "bar"},
|
|
29
|
+
)
|
|
30
|
+
manager.write_state(state)
|
|
31
|
+
fetched_state = manager.read_state("asdf")
|
|
32
|
+
|
|
33
|
+
self.assertEqual(state, fetched_state)
|
|
34
|
+
|
|
35
|
+
def test__raises_an_error_if_referencing_nonexistent_task(self):
|
|
36
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
37
|
+
manager = StateManager(MyAsyncTaskState, tmpdir)
|
|
38
|
+
with self.assertRaises(NoSuchTaskException):
|
|
39
|
+
manager.read_state("asdf")
|
|
40
|
+
|
|
41
|
+
def test__allows_specific_update_of_status_field(self):
|
|
42
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
43
|
+
manager = StateManager(MyAsyncTaskState, tmpdir)
|
|
44
|
+
state = MyAsyncTaskState(
|
|
45
|
+
task_id="asdf",
|
|
46
|
+
estimated_time="40 days and 40 nights",
|
|
47
|
+
task_status="STARTED",
|
|
48
|
+
task_result=None,
|
|
49
|
+
extra_state={"foo": "bar"},
|
|
50
|
+
)
|
|
51
|
+
manager.write_state(state)
|
|
52
|
+
manager.update_status(state.task_id, "fail")
|
|
53
|
+
fetched_state = manager.read_state("asdf")
|
|
54
|
+
state.task_status = "fail"
|
|
55
|
+
|
|
56
|
+
self.assertEqual(state, fetched_state)
|
|
57
|
+
|
|
58
|
+
def test__allows_specific_update_of_result_field(self):
|
|
59
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
60
|
+
manager = StateManager(MyAsyncTaskState, tmpdir)
|
|
61
|
+
state = MyAsyncTaskState(
|
|
62
|
+
task_id="asdf",
|
|
63
|
+
estimated_time="40 days and 40 nights",
|
|
64
|
+
task_status="STARTED",
|
|
65
|
+
task_result=None,
|
|
66
|
+
extra_state={"foo": "bar"},
|
|
67
|
+
)
|
|
68
|
+
manager.write_state(state)
|
|
69
|
+
|
|
70
|
+
result = MyTaskResult(a="asdf", b=123)
|
|
71
|
+
manager.save_result(state.task_id, result)
|
|
72
|
+
fetched_state = manager.read_state("asdf")
|
|
73
|
+
|
|
74
|
+
state.task_status = TASK_STATUSES["COMPLETED"]
|
|
75
|
+
state.task_result = result
|
|
76
|
+
|
|
77
|
+
self.assertEqual(state, fetched_state)
|
nora_lib-0.0.1.dev0/README.md
DELETED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|