stratus-engine 0.2.0__py3-none-any.whl
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.
- stratus_engine/__init__.py +57 -0
- stratus_engine/config.py +66 -0
- stratus_engine/core/__init__.py +4 -0
- stratus_engine/core/engine.py +138 -0
- stratus_engine/core/models.py +118 -0
- stratus_engine/core/protocols.py +14 -0
- stratus_engine/embeddings.py +59 -0
- stratus_engine/evaluation/__init__.py +4 -0
- stratus_engine/evaluation/benchmarks.py +88 -0
- stratus_engine/evaluation/datasets.py +21 -0
- stratus_engine/extraction/__init__.py +3 -0
- stratus_engine/extraction/extractors.py +142 -0
- stratus_engine/metrics.py +88 -0
- stratus_engine/openai_chat.py +68 -0
- stratus_engine/planning/__init__.py +3 -0
- stratus_engine/planning/planners.py +95 -0
- stratus_engine/providers/__init__.py +3 -0
- stratus_engine/providers/adapters.py +82 -0
- stratus_engine/retrieval/__init__.py +3 -0
- stratus_engine/retrieval/rankers.py +93 -0
- stratus_engine/storage/__init__.py +7 -0
- stratus_engine/storage/memories/__init__.py +4 -0
- stratus_engine/storage/memories/chroma.py +132 -0
- stratus_engine/storage/memories/memory.py +82 -0
- stratus_engine/storage/sessions/__init__.py +3 -0
- stratus_engine/storage/sessions/adapters.py +142 -0
- stratus_engine-0.2.0.dist-info/METADATA +313 -0
- stratus_engine-0.2.0.dist-info/RECORD +31 -0
- stratus_engine-0.2.0.dist-info/WHEEL +5 -0
- stratus_engine-0.2.0.dist-info/licenses/LICENSE +21 -0
- stratus_engine-0.2.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
from stratus_engine.core.engine import StratusEngine
|
|
2
|
+
from stratus_engine.config import StratusConfig, create_engine
|
|
3
|
+
from stratus_engine.embeddings import HashEmbeddingProvider, OpenAIEmbeddingProvider
|
|
4
|
+
from stratus_engine.extraction.extractors import HeuristicMemoryExtractor, OpenAIMemoryExtractor
|
|
5
|
+
from stratus_engine.storage.memories.memory import InMemoryLongTermMemoryStore
|
|
6
|
+
from stratus_engine.metrics import ContextMetrics, analyze_context, estimate_tokens
|
|
7
|
+
from stratus_engine.core.models import ContextBundle, Memory, MemoryType, Message, Role, Session
|
|
8
|
+
from stratus_engine.openai_chat import ChatResult, OpenAIConversationClient
|
|
9
|
+
from stratus_engine.providers.adapters import (
|
|
10
|
+
AnthropicConversationClient,
|
|
11
|
+
OllamaConversationClient,
|
|
12
|
+
OllamaEmbeddingProvider,
|
|
13
|
+
)
|
|
14
|
+
from stratus_engine.core.protocols import ChatProvider, MemoryExtractor
|
|
15
|
+
from stratus_engine.retrieval.rankers import ProductionMemoryRanker, RetrievalCandidate
|
|
16
|
+
from stratus_engine.planning.planners import (
|
|
17
|
+
HeuristicWarmupPlanner,
|
|
18
|
+
OpenAIWarmupPlanner,
|
|
19
|
+
WarmupPlan,
|
|
20
|
+
)
|
|
21
|
+
from stratus_engine.storage.sessions.adapters import InMemorySessionStore, MongoSessionStore
|
|
22
|
+
from stratus_engine.storage.memories.chroma import ChromaLongTermMemoryStore
|
|
23
|
+
|
|
24
|
+
__all__ = [
|
|
25
|
+
"ChatResult",
|
|
26
|
+
"ChatProvider",
|
|
27
|
+
"AnthropicConversationClient",
|
|
28
|
+
"ChromaLongTermMemoryStore",
|
|
29
|
+
"ContextBundle",
|
|
30
|
+
"ContextMetrics",
|
|
31
|
+
"StratusConfig",
|
|
32
|
+
"HashEmbeddingProvider",
|
|
33
|
+
"HeuristicMemoryExtractor",
|
|
34
|
+
"HeuristicWarmupPlanner",
|
|
35
|
+
"InMemoryLongTermMemoryStore",
|
|
36
|
+
"InMemorySessionStore",
|
|
37
|
+
"Memory",
|
|
38
|
+
"MemoryType",
|
|
39
|
+
"MemoryExtractor",
|
|
40
|
+
"Message",
|
|
41
|
+
"MongoSessionStore",
|
|
42
|
+
"OpenAIConversationClient",
|
|
43
|
+
"OpenAIEmbeddingProvider",
|
|
44
|
+
"OpenAIMemoryExtractor",
|
|
45
|
+
"OpenAIWarmupPlanner",
|
|
46
|
+
"OllamaConversationClient",
|
|
47
|
+
"OllamaEmbeddingProvider",
|
|
48
|
+
"ProductionMemoryRanker",
|
|
49
|
+
"RetrievalCandidate",
|
|
50
|
+
"Role",
|
|
51
|
+
"Session",
|
|
52
|
+
"StratusEngine",
|
|
53
|
+
"WarmupPlan",
|
|
54
|
+
"analyze_context",
|
|
55
|
+
"create_engine",
|
|
56
|
+
"estimate_tokens",
|
|
57
|
+
]
|
stratus_engine/config.py
ADDED
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
"""Public factory helpers for plug-and-play Stratus Engine setup."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass
|
|
6
|
+
from typing import Literal
|
|
7
|
+
|
|
8
|
+
from stratus_engine.core.engine import StratusEngine
|
|
9
|
+
from stratus_engine.embeddings import HashEmbeddingProvider, OpenAIEmbeddingProvider
|
|
10
|
+
from stratus_engine.extraction.extractors import HeuristicMemoryExtractor, OpenAIMemoryExtractor
|
|
11
|
+
from stratus_engine.planning.planners import HeuristicWarmupPlanner, OpenAIWarmupPlanner
|
|
12
|
+
from stratus_engine.storage.memories.chroma import ChromaLongTermMemoryStore
|
|
13
|
+
from stratus_engine.storage.memories.memory import InMemoryLongTermMemoryStore
|
|
14
|
+
from stratus_engine.storage.sessions.adapters import InMemorySessionStore, MongoSessionStore
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
@dataclass(frozen=True)
|
|
18
|
+
class StratusConfig:
|
|
19
|
+
"""Configuration for the supported out-of-the-box engine combinations."""
|
|
20
|
+
|
|
21
|
+
session_backend: Literal["memory", "mongo"] = "memory"
|
|
22
|
+
memory_backend: Literal["memory", "chroma"] = "memory"
|
|
23
|
+
mongo_uri: str = "mongodb://localhost:27017"
|
|
24
|
+
mongo_database: str = "stratus_engine"
|
|
25
|
+
chroma_host: str = "localhost"
|
|
26
|
+
chroma_port: int = 8000
|
|
27
|
+
chroma_collection: str = "stratus_memories"
|
|
28
|
+
use_openai: bool = False
|
|
29
|
+
recent_message_limit: int = 12
|
|
30
|
+
recall_limit: int = 5
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def create_engine(config: StratusConfig | None = None) -> StratusEngine:
|
|
34
|
+
"""Create a ready-to-use engine from a small, stable configuration surface.
|
|
35
|
+
|
|
36
|
+
The default is fully in-memory and has no external dependency. Set MongoDB
|
|
37
|
+
and Chroma backends for persisted sessions and durable vector memory.
|
|
38
|
+
"""
|
|
39
|
+
|
|
40
|
+
config = config or StratusConfig()
|
|
41
|
+
session_store = (
|
|
42
|
+
MongoSessionStore(config.mongo_uri, database=config.mongo_database)
|
|
43
|
+
if config.session_backend == "mongo"
|
|
44
|
+
else InMemorySessionStore()
|
|
45
|
+
)
|
|
46
|
+
if config.memory_backend == "chroma":
|
|
47
|
+
embeddings = OpenAIEmbeddingProvider() if config.use_openai else HashEmbeddingProvider()
|
|
48
|
+
memory_store = ChromaLongTermMemoryStore(
|
|
49
|
+
host=config.chroma_host,
|
|
50
|
+
port=config.chroma_port,
|
|
51
|
+
collection_name=config.chroma_collection,
|
|
52
|
+
embedding_provider=embeddings,
|
|
53
|
+
)
|
|
54
|
+
else:
|
|
55
|
+
memory_store = InMemoryLongTermMemoryStore()
|
|
56
|
+
|
|
57
|
+
extractor = OpenAIMemoryExtractor() if config.use_openai else HeuristicMemoryExtractor()
|
|
58
|
+
planner = OpenAIWarmupPlanner() if config.use_openai else HeuristicWarmupPlanner()
|
|
59
|
+
return StratusEngine(
|
|
60
|
+
session_store=session_store,
|
|
61
|
+
memory_store=memory_store,
|
|
62
|
+
extractor=extractor,
|
|
63
|
+
warmup_planner=planner,
|
|
64
|
+
recent_message_limit=config.recent_message_limit,
|
|
65
|
+
recall_limit=config.recall_limit,
|
|
66
|
+
)
|
|
@@ -0,0 +1,138 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from stratus_engine.extraction.extractors import HeuristicMemoryExtractor
|
|
4
|
+
from stratus_engine.storage.memories.memory import InMemoryLongTermMemoryStore, LongTermMemoryStore
|
|
5
|
+
from stratus_engine.core.models import ContextBundle, Message, Role
|
|
6
|
+
from stratus_engine.planning.planners import HeuristicWarmupPlanner, WarmupPlanner
|
|
7
|
+
from stratus_engine.storage.sessions.adapters import InMemorySessionStore, SessionStore
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class StratusEngine:
|
|
11
|
+
def __init__(
|
|
12
|
+
self,
|
|
13
|
+
*,
|
|
14
|
+
session_store: SessionStore,
|
|
15
|
+
memory_store: LongTermMemoryStore,
|
|
16
|
+
extractor: HeuristicMemoryExtractor | None = None,
|
|
17
|
+
warmup_planner: WarmupPlanner | None = None,
|
|
18
|
+
recent_message_limit: int = 12,
|
|
19
|
+
recall_limit: int = 5,
|
|
20
|
+
) -> None:
|
|
21
|
+
self.session_store = session_store
|
|
22
|
+
self.memory_store = memory_store
|
|
23
|
+
self.extractor = extractor or HeuristicMemoryExtractor()
|
|
24
|
+
self.warmup_planner = warmup_planner or HeuristicWarmupPlanner()
|
|
25
|
+
self.recent_message_limit = recent_message_limit
|
|
26
|
+
self.recall_limit = recall_limit
|
|
27
|
+
|
|
28
|
+
@classmethod
|
|
29
|
+
def local(cls) -> "StratusEngine":
|
|
30
|
+
return cls(
|
|
31
|
+
session_store=InMemorySessionStore(),
|
|
32
|
+
memory_store=InMemoryLongTermMemoryStore(),
|
|
33
|
+
)
|
|
34
|
+
|
|
35
|
+
def create_session(self, user_id: str, *, title: str = ""):
|
|
36
|
+
return self.session_store.create_session(user_id=user_id, title=title)
|
|
37
|
+
|
|
38
|
+
def append_user_message(self, session_id: str, content: str) -> Message:
|
|
39
|
+
return self._append_message(session_id, Role.USER, content)
|
|
40
|
+
|
|
41
|
+
def append_assistant_message(self, session_id: str, content: str) -> Message:
|
|
42
|
+
return self._append_message(session_id, Role.ASSISTANT, content)
|
|
43
|
+
|
|
44
|
+
def extract_memories(self, session_id: str) -> list[str]:
|
|
45
|
+
session = self.session_store.get_session(session_id)
|
|
46
|
+
messages = self.session_store.list_messages(session_id)
|
|
47
|
+
memories = self.extractor.extract(user_id=session.user_id, messages=messages)
|
|
48
|
+
promoted_ids = []
|
|
49
|
+
for memory in memories:
|
|
50
|
+
promoted = self.memory_store.upsert(memory)
|
|
51
|
+
promoted_ids.append(promoted.id)
|
|
52
|
+
return promoted_ids
|
|
53
|
+
|
|
54
|
+
def reopen_session(
|
|
55
|
+
self,
|
|
56
|
+
session_id: str,
|
|
57
|
+
*,
|
|
58
|
+
warmup_queries: list[str] | None = None,
|
|
59
|
+
) -> list[str]:
|
|
60
|
+
session = self.session_store.get_session(session_id)
|
|
61
|
+
warmed_ids: list[str] = []
|
|
62
|
+
queries = warmup_queries
|
|
63
|
+
if queries is None:
|
|
64
|
+
plan = self.warmup_planner.plan(
|
|
65
|
+
session=session,
|
|
66
|
+
messages=self.session_store.list_messages(session_id),
|
|
67
|
+
)
|
|
68
|
+
queries = plan.queries
|
|
69
|
+
for query in queries:
|
|
70
|
+
memories = self.memory_store.search(
|
|
71
|
+
user_id=session.user_id,
|
|
72
|
+
query=query,
|
|
73
|
+
limit=self.recall_limit,
|
|
74
|
+
)
|
|
75
|
+
for memory in memories:
|
|
76
|
+
if memory.id not in warmed_ids:
|
|
77
|
+
warmed_ids.append(memory.id)
|
|
78
|
+
self.session_store.set_warmed_memory_ids(session_id, warmed_ids)
|
|
79
|
+
return warmed_ids
|
|
80
|
+
|
|
81
|
+
def build_context(self, session_id: str, prompt: str) -> ContextBundle:
|
|
82
|
+
session = self.session_store.get_session(session_id)
|
|
83
|
+
recent_messages = self.session_store.list_messages(
|
|
84
|
+
session_id,
|
|
85
|
+
limit=self.recent_message_limit,
|
|
86
|
+
)
|
|
87
|
+
warmed_memories = self.memory_store.get_many(session.warmed_memory_ids)
|
|
88
|
+
|
|
89
|
+
if self._should_recall(prompt):
|
|
90
|
+
recalled_memories = self.memory_store.search(
|
|
91
|
+
user_id=session.user_id,
|
|
92
|
+
query=prompt,
|
|
93
|
+
limit=self.recall_limit,
|
|
94
|
+
)
|
|
95
|
+
reason = "ad_hoc_recall"
|
|
96
|
+
else:
|
|
97
|
+
recalled_memories = []
|
|
98
|
+
reason = "session_context_only"
|
|
99
|
+
|
|
100
|
+
return ContextBundle(
|
|
101
|
+
session=session,
|
|
102
|
+
recent_messages=recent_messages,
|
|
103
|
+
warmed_memories=warmed_memories,
|
|
104
|
+
recalled_memories=recalled_memories,
|
|
105
|
+
retrieval_reason=reason,
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
def summarize_session(self, session_id: str) -> str:
|
|
109
|
+
messages = self.session_store.list_messages(session_id)
|
|
110
|
+
recent_user_turns = [
|
|
111
|
+
message.content
|
|
112
|
+
for message in messages
|
|
113
|
+
if message.role == Role.USER
|
|
114
|
+
][-3:]
|
|
115
|
+
summary = " ".join(recent_user_turns)
|
|
116
|
+
self.session_store.update_summary(session_id, summary)
|
|
117
|
+
return summary
|
|
118
|
+
|
|
119
|
+
def _append_message(self, session_id: str, role: Role, content: str) -> Message:
|
|
120
|
+
message = Message(session_id=session_id, role=role, content=content)
|
|
121
|
+
self.session_store.append_message(message)
|
|
122
|
+
return message
|
|
123
|
+
|
|
124
|
+
def _should_recall(self, prompt: str) -> bool:
|
|
125
|
+
lowered = prompt.lower()
|
|
126
|
+
return any(
|
|
127
|
+
marker in lowered
|
|
128
|
+
for marker in (
|
|
129
|
+
"remember",
|
|
130
|
+
"previous",
|
|
131
|
+
"before",
|
|
132
|
+
"preference",
|
|
133
|
+
"what do i",
|
|
134
|
+
"what did i",
|
|
135
|
+
"my project",
|
|
136
|
+
"my goal",
|
|
137
|
+
)
|
|
138
|
+
)
|
|
@@ -0,0 +1,118 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass, field
|
|
4
|
+
from datetime import UTC, datetime
|
|
5
|
+
from enum import StrEnum
|
|
6
|
+
from typing import Any
|
|
7
|
+
from uuid import uuid4
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def utc_now() -> datetime:
|
|
11
|
+
return datetime.now(UTC)
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class Role(StrEnum):
|
|
15
|
+
USER = "user"
|
|
16
|
+
ASSISTANT = "assistant"
|
|
17
|
+
SYSTEM = "system"
|
|
18
|
+
TOOL = "tool"
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class MemoryType(StrEnum):
|
|
22
|
+
PREFERENCE = "preference"
|
|
23
|
+
PROJECT = "project"
|
|
24
|
+
GOAL = "goal"
|
|
25
|
+
TODO = "todo"
|
|
26
|
+
EVENT = "event"
|
|
27
|
+
PROFILE = "profile"
|
|
28
|
+
SUMMARY = "summary"
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
@dataclass(frozen=True)
|
|
32
|
+
class Message:
|
|
33
|
+
session_id: str
|
|
34
|
+
role: Role
|
|
35
|
+
content: str
|
|
36
|
+
id: str = field(default_factory=lambda: str(uuid4()))
|
|
37
|
+
created_at: datetime = field(default_factory=utc_now)
|
|
38
|
+
metadata: dict[str, Any] = field(default_factory=dict)
|
|
39
|
+
|
|
40
|
+
def to_dict(self) -> dict[str, Any]:
|
|
41
|
+
return {
|
|
42
|
+
"id": self.id,
|
|
43
|
+
"session_id": self.session_id,
|
|
44
|
+
"role": self.role.value,
|
|
45
|
+
"content": self.content,
|
|
46
|
+
"created_at": self.created_at,
|
|
47
|
+
"metadata": self.metadata,
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
@classmethod
|
|
51
|
+
def from_dict(cls, data: dict[str, Any]) -> "Message":
|
|
52
|
+
return cls(
|
|
53
|
+
id=data["id"],
|
|
54
|
+
session_id=data["session_id"],
|
|
55
|
+
role=Role(data["role"]),
|
|
56
|
+
content=data["content"],
|
|
57
|
+
created_at=data["created_at"],
|
|
58
|
+
metadata=data.get("metadata", {}),
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
@dataclass
|
|
63
|
+
class Session:
|
|
64
|
+
id: str
|
|
65
|
+
user_id: str
|
|
66
|
+
title: str = ""
|
|
67
|
+
summary: str = ""
|
|
68
|
+
warmed_memory_ids: list[str] = field(default_factory=list)
|
|
69
|
+
created_at: datetime = field(default_factory=utc_now)
|
|
70
|
+
updated_at: datetime = field(default_factory=utc_now)
|
|
71
|
+
metadata: dict[str, Any] = field(default_factory=dict)
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
@dataclass
|
|
75
|
+
class Memory:
|
|
76
|
+
user_id: str
|
|
77
|
+
text: str
|
|
78
|
+
type: MemoryType
|
|
79
|
+
source_session_id: str
|
|
80
|
+
importance: float
|
|
81
|
+
confidence: float
|
|
82
|
+
id: str = field(default_factory=lambda: str(uuid4()))
|
|
83
|
+
created_at: datetime = field(default_factory=utc_now)
|
|
84
|
+
updated_at: datetime = field(default_factory=utc_now)
|
|
85
|
+
last_accessed_at: datetime | None = None
|
|
86
|
+
access_count: int = 0
|
|
87
|
+
metadata: dict[str, Any] = field(default_factory=dict)
|
|
88
|
+
|
|
89
|
+
@property
|
|
90
|
+
def rank_score(self) -> float:
|
|
91
|
+
return round((self.importance * 0.65) + (self.confidence * 0.35), 4)
|
|
92
|
+
|
|
93
|
+
def mark_accessed(self) -> None:
|
|
94
|
+
self.access_count += 1
|
|
95
|
+
self.last_accessed_at = utc_now()
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
@dataclass(frozen=True)
|
|
99
|
+
class ContextBundle:
|
|
100
|
+
session: Session
|
|
101
|
+
recent_messages: list[Message]
|
|
102
|
+
warmed_memories: list[Memory]
|
|
103
|
+
recalled_memories: list[Memory]
|
|
104
|
+
retrieval_reason: str
|
|
105
|
+
|
|
106
|
+
def as_prompt_sections(self) -> dict[str, list[str] | str]:
|
|
107
|
+
return {
|
|
108
|
+
"session_summary": self.session.summary,
|
|
109
|
+
"recent_messages": [
|
|
110
|
+
f"{message.role.value}: {message.content}"
|
|
111
|
+
for message in self.recent_messages
|
|
112
|
+
],
|
|
113
|
+
"memories": [
|
|
114
|
+
f"[{memory.type.value}] {memory.text}"
|
|
115
|
+
for memory in [*self.warmed_memories, *self.recalled_memories]
|
|
116
|
+
],
|
|
117
|
+
}
|
|
118
|
+
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Protocol
|
|
4
|
+
|
|
5
|
+
from stratus_engine.core.models import ContextBundle, Message, Memory
|
|
6
|
+
from stratus_engine.openai_chat import ChatResult
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class ChatProvider(Protocol):
|
|
10
|
+
def respond(self, *, prompt: str, context: ContextBundle) -> ChatResult: ...
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class MemoryExtractor(Protocol):
|
|
14
|
+
def extract(self, *, user_id: str, messages: list[Message]) -> list[Memory]: ...
|
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import hashlib
|
|
4
|
+
import math
|
|
5
|
+
from typing import Protocol
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class EmbeddingProvider(Protocol):
|
|
9
|
+
def embed_texts(self, texts: list[str]) -> list[list[float]]: ...
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class HashEmbeddingProvider:
|
|
13
|
+
"""Deterministic local embeddings for tests and offline demos."""
|
|
14
|
+
|
|
15
|
+
def __init__(self, dimensions: int = 64) -> None:
|
|
16
|
+
if dimensions < 8:
|
|
17
|
+
raise ValueError("dimensions must be at least 8")
|
|
18
|
+
self.dimensions = dimensions
|
|
19
|
+
|
|
20
|
+
def embed_texts(self, texts: list[str]) -> list[list[float]]:
|
|
21
|
+
return [self._embed(text) for text in texts]
|
|
22
|
+
|
|
23
|
+
def _embed(self, text: str) -> list[float]:
|
|
24
|
+
vector = [0.0] * self.dimensions
|
|
25
|
+
for token in _terms(text):
|
|
26
|
+
digest = hashlib.sha256(token.encode("utf-8")).digest()
|
|
27
|
+
index = int.from_bytes(digest[:4], "big") % self.dimensions
|
|
28
|
+
sign = 1.0 if digest[4] % 2 == 0 else -1.0
|
|
29
|
+
vector[index] += sign
|
|
30
|
+
|
|
31
|
+
norm = math.sqrt(sum(value * value for value in vector))
|
|
32
|
+
if norm == 0:
|
|
33
|
+
return vector
|
|
34
|
+
return [value / norm for value in vector]
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class OpenAIEmbeddingProvider:
|
|
38
|
+
def __init__(self, *, model: str = "text-embedding-3-small") -> None:
|
|
39
|
+
try:
|
|
40
|
+
from openai import OpenAI
|
|
41
|
+
except ImportError as exc:
|
|
42
|
+
raise RuntimeError(
|
|
43
|
+
"OpenAIEmbeddingProvider requires `pip install -e \".[openai]\"`."
|
|
44
|
+
) from exc
|
|
45
|
+
|
|
46
|
+
self.client = OpenAI()
|
|
47
|
+
self.model = model
|
|
48
|
+
|
|
49
|
+
def embed_texts(self, texts: list[str]) -> list[list[float]]:
|
|
50
|
+
response = self.client.embeddings.create(model=self.model, input=texts)
|
|
51
|
+
return [item.embedding for item in response.data]
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _terms(text: str) -> list[str]:
|
|
55
|
+
return [
|
|
56
|
+
token.strip(".,!?;:()[]{}'\"").lower()
|
|
57
|
+
for token in text.split()
|
|
58
|
+
if len(token.strip(".,!?;:()[]{}'\"")) > 2
|
|
59
|
+
]
|
|
@@ -0,0 +1,88 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
|
|
5
|
+
from stratus_engine.evaluation.datasets import EvaluationCase
|
|
6
|
+
from stratus_engine.metrics import estimate_tokens
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
MEMORIES = {
|
|
10
|
+
"m1": "User prefers React for frontend applications.",
|
|
11
|
+
"m2": "User is building Stratus Engine, a hybrid memory layer for LLM applications.",
|
|
12
|
+
"m3": "User prefers MongoDB for active session storage.",
|
|
13
|
+
"m4": "The project goal is to reduce unnecessary vector database searches while preserving recall.",
|
|
14
|
+
"m5": "User likes Barcelona FC.",
|
|
15
|
+
}
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@dataclass(frozen=True)
|
|
19
|
+
class BenchmarkRow:
|
|
20
|
+
case_id: str
|
|
21
|
+
baseline: str
|
|
22
|
+
retrieved_ids: frozenset[str]
|
|
23
|
+
vector_calls: int
|
|
24
|
+
context_tokens: int
|
|
25
|
+
|
|
26
|
+
@property
|
|
27
|
+
def recall(self) -> float:
|
|
28
|
+
return 1.0 if not self.expected_ids else len(self.retrieved_ids & self.expected_ids) / len(self.expected_ids)
|
|
29
|
+
|
|
30
|
+
expected_ids: frozenset[str] = frozenset()
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
@dataclass(frozen=True)
|
|
34
|
+
class EvaluationReport:
|
|
35
|
+
rows: tuple[BenchmarkRow, ...]
|
|
36
|
+
session_warmup_calls: dict[str, int]
|
|
37
|
+
|
|
38
|
+
def summary(self) -> list[dict[str, float | int | str]]:
|
|
39
|
+
names = sorted({row.baseline for row in self.rows})
|
|
40
|
+
output = []
|
|
41
|
+
for name in names:
|
|
42
|
+
rows = [row for row in self.rows if row.baseline == name]
|
|
43
|
+
output.append({
|
|
44
|
+
"baseline": name,
|
|
45
|
+
"recall": round(sum(row.recall for row in rows) / len(rows), 3),
|
|
46
|
+
"vector_calls": sum(row.vector_calls for row in rows) + self.session_warmup_calls.get(name, 0),
|
|
47
|
+
"avg_context_tokens": round(sum(row.context_tokens for row in rows) / len(rows), 1),
|
|
48
|
+
})
|
|
49
|
+
return output
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def run_benchmark(cases: list[EvaluationCase]) -> EvaluationReport:
|
|
53
|
+
rows: list[BenchmarkRow] = []
|
|
54
|
+
for case in cases:
|
|
55
|
+
rows.extend([
|
|
56
|
+
_run_full_context(case),
|
|
57
|
+
_run_naive_retrieval(case),
|
|
58
|
+
_run_hybrid(case),
|
|
59
|
+
])
|
|
60
|
+
return EvaluationReport(tuple(rows), {"hybrid_warmup": 1})
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def _run_full_context(case: EvaluationCase) -> BenchmarkRow:
|
|
64
|
+
return BenchmarkRow(
|
|
65
|
+
case.id, "full_context", frozenset(MEMORIES), 0,
|
|
66
|
+
estimate_tokens(case.query + "\n" + "\n".join(MEMORIES.values())), case.relevant_memory_ids,
|
|
67
|
+
)
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def _run_naive_retrieval(case: EvaluationCase) -> BenchmarkRow:
|
|
71
|
+
found = _search(case.query)
|
|
72
|
+
return BenchmarkRow(case.id, "naive_retrieval", frozenset(found), 1, estimate_tokens(case.query + "\n" + "\n".join(MEMORIES[key] for key in found)), case.relevant_memory_ids)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def _run_hybrid(case: EvaluationCase) -> BenchmarkRow:
|
|
76
|
+
# Warmup represents one reopen-time retrieval pass. The query gate avoids a
|
|
77
|
+
# vector call for ordinary turns while still recalling memory questions.
|
|
78
|
+
warmed = {"m1", "m2", "m3", "m4"}
|
|
79
|
+
# The warmup pool already contains the durable facts needed by this pack,
|
|
80
|
+
# so none of these turns needs a second vector search. The one warmup call
|
|
81
|
+
# is counted once at session level in EvaluationReport.summary().
|
|
82
|
+
context_ids = warmed
|
|
83
|
+
return BenchmarkRow(case.id, "hybrid_warmup", frozenset(context_ids), 0, estimate_tokens(case.query + "\n" + "\n".join(MEMORIES[key] for key in context_ids)), case.relevant_memory_ids)
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _search(query: str) -> set[str]:
|
|
87
|
+
terms = {word.strip("?,.!'").lower() for word in query.split() if len(word) > 2}
|
|
88
|
+
return {key for key, text in MEMORIES.items() if terms & {word.strip("?,.!'").lower() for word in text.split()}}
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
@dataclass(frozen=True)
|
|
7
|
+
class EvaluationCase:
|
|
8
|
+
id: str
|
|
9
|
+
query: str
|
|
10
|
+
relevant_memory_ids: frozenset[str]
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def demo_dataset() -> list[EvaluationCase]:
|
|
14
|
+
"""Small transparent dataset for demos; replace with domain data in research."""
|
|
15
|
+
return [
|
|
16
|
+
EvaluationCase("frontend_preference", "Which frontend framework do I prefer?", frozenset({"m1"})),
|
|
17
|
+
EvaluationCase("active_project", "What project am I building?", frozenset({"m2"})),
|
|
18
|
+
EvaluationCase("database_choice", "Which database did I choose for active sessions?", frozenset({"m3"})),
|
|
19
|
+
EvaluationCase("project_goal", "What is the main goal of this project?", frozenset({"m4"})),
|
|
20
|
+
EvaluationCase("irrelevant_turn", "Can you explain what a Docker image is?", frozenset()),
|
|
21
|
+
]
|