protoprompt 0.1.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.
- protoprompt/__init__.py +58 -0
- protoprompt/context.py +33 -0
- protoprompt/exceptions.py +20 -0
- protoprompt/injector.py +69 -0
- protoprompt/injector_budgeted.py +219 -0
- protoprompt/llm.py +12 -0
- protoprompt/pipeline.py +71 -0
- protoprompt/profile/__init__.py +4 -0
- protoprompt/profile/builder.py +62 -0
- protoprompt/profile/schema.py +24 -0
- protoprompt/profile/types.py +12 -0
- protoprompt/session/__init__.py +16 -0
- protoprompt/session/compressor.py +17 -0
- protoprompt/session/strategy.py +200 -0
- protoprompt/session/types.py +17 -0
- protoprompt/store/__init__.py +4 -0
- protoprompt/store/chroma.py +71 -0
- protoprompt/store/memory.py +79 -0
- protoprompt/store/protocol.py +39 -0
- protoprompt/tokens/__init__.py +7 -0
- protoprompt/tokens/protocol.py +25 -0
- protoprompt/tokens/regex_counter.py +52 -0
- protoprompt/tokens/tiktoken_adapter.py +57 -0
- protoprompt-0.1.0.dist-info/METADATA +201 -0
- protoprompt-0.1.0.dist-info/RECORD +28 -0
- protoprompt-0.1.0.dist-info/WHEEL +5 -0
- protoprompt-0.1.0.dist-info/licenses/LICENSE +21 -0
- protoprompt-0.1.0.dist-info/top_level.txt +1 -0
protoprompt/__init__.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
"""protoprompt: layered context builder for LLM prompts.
|
|
2
|
+
|
|
3
|
+
Public top-level exports live here; per-layer subpackages re-export
|
|
4
|
+
their own public API as well, so both styles work:
|
|
5
|
+
|
|
6
|
+
from protoprompt import ContextBuilder
|
|
7
|
+
from protoprompt.context import ContextInput
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from protoprompt.context import ContextInput, ContextOutput
|
|
11
|
+
from protoprompt.exceptions import TokenBudgetExceededError
|
|
12
|
+
from protoprompt.injector import ContextBuilder
|
|
13
|
+
from protoprompt.injector_budgeted import (
|
|
14
|
+
DEFAULT_PRIORITIES,
|
|
15
|
+
BudgetReport,
|
|
16
|
+
TokenBudgetedContextBuilder,
|
|
17
|
+
)
|
|
18
|
+
from protoprompt.llm import LLMClientProtocol
|
|
19
|
+
from protoprompt.pipeline import Pipeline
|
|
20
|
+
from protoprompt.profile.builder import ProfileBuilder
|
|
21
|
+
from protoprompt.profile.types import UserProfile
|
|
22
|
+
from protoprompt.session.compressor import Compressor
|
|
23
|
+
from protoprompt.session.strategy import (
|
|
24
|
+
HeuristicStrategy,
|
|
25
|
+
LLMSummaryStrategy,
|
|
26
|
+
StrategyProtocol,
|
|
27
|
+
)
|
|
28
|
+
from protoprompt.session.types import CompressedBlock, Session
|
|
29
|
+
from protoprompt.store.memory import InMemStore
|
|
30
|
+
from protoprompt.store.protocol import StoreProtocol
|
|
31
|
+
from protoprompt.tokens.protocol import TokenCounter
|
|
32
|
+
from protoprompt.tokens.regex_counter import RegexTokenCounter
|
|
33
|
+
|
|
34
|
+
__version__ = "0.1.0"
|
|
35
|
+
|
|
36
|
+
__all__ = [
|
|
37
|
+
"ContextBuilder",
|
|
38
|
+
"TokenBudgetedContextBuilder",
|
|
39
|
+
"ContextInput",
|
|
40
|
+
"ContextOutput",
|
|
41
|
+
"BudgetReport",
|
|
42
|
+
"DEFAULT_PRIORITIES",
|
|
43
|
+
"TokenBudgetExceededError",
|
|
44
|
+
"Pipeline",
|
|
45
|
+
"Compressor",
|
|
46
|
+
"StrategyProtocol",
|
|
47
|
+
"HeuristicStrategy",
|
|
48
|
+
"LLMSummaryStrategy",
|
|
49
|
+
"Session",
|
|
50
|
+
"CompressedBlock",
|
|
51
|
+
"UserProfile",
|
|
52
|
+
"ProfileBuilder",
|
|
53
|
+
"StoreProtocol",
|
|
54
|
+
"InMemStore",
|
|
55
|
+
"TokenCounter",
|
|
56
|
+
"RegexTokenCounter",
|
|
57
|
+
"LLMClientProtocol",
|
|
58
|
+
]
|
protoprompt/context.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
"""Public dataclasses for context assembly."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass, field
|
|
6
|
+
from typing import TYPE_CHECKING
|
|
7
|
+
|
|
8
|
+
if TYPE_CHECKING:
|
|
9
|
+
from protoprompt.injector_budgeted import BudgetReport
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@dataclass
|
|
13
|
+
class ContextInput:
|
|
14
|
+
query: str
|
|
15
|
+
chat_id: str = ""
|
|
16
|
+
system_prompt: str = ""
|
|
17
|
+
doc_ids: list[int] = field(default_factory=list)
|
|
18
|
+
embedding_model: str = "nomic-embed-text"
|
|
19
|
+
top_k_rag: int = 5
|
|
20
|
+
top_k_session: int = 3
|
|
21
|
+
include_rag: bool = True
|
|
22
|
+
include_session: bool = True
|
|
23
|
+
include_profile: bool = False
|
|
24
|
+
profile_text: str = ""
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass
|
|
28
|
+
class ContextOutput:
|
|
29
|
+
system_prompt: str
|
|
30
|
+
rag_blocks: list[str] = field(default_factory=list)
|
|
31
|
+
session_blocks: list[str] = field(default_factory=list)
|
|
32
|
+
profile_used: bool = False
|
|
33
|
+
budget_report: "BudgetReport | None" = None
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""Exceptions raised by the context assembly pipeline."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class TokenBudgetExceededError(RuntimeError):
|
|
7
|
+
"""Raised when a hard-required section (system prompt) does not fit
|
|
8
|
+
into the configured token budget.
|
|
9
|
+
|
|
10
|
+
Soft sections (RAG, session) are dropped silently by the budget
|
|
11
|
+
allocator; only mandatory sections ever trigger this.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
def __init__(self, used: int, budget: int, section: str) -> None:
|
|
15
|
+
super().__init__(
|
|
16
|
+
f"Section '{section}' needs {used} tokens but budget is {budget}"
|
|
17
|
+
)
|
|
18
|
+
self.used = used
|
|
19
|
+
self.budget = budget
|
|
20
|
+
self.section = section
|
protoprompt/injector.py
ADDED
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import logging
|
|
4
|
+
|
|
5
|
+
from protoprompt.context import ContextInput, ContextOutput
|
|
6
|
+
from protoprompt.llm import LLMClientProtocol
|
|
7
|
+
from protoprompt.store.protocol import StoreProtocol
|
|
8
|
+
|
|
9
|
+
logger = logging.getLogger(__name__)
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class ContextBuilder:
|
|
13
|
+
"""Default context assembler.
|
|
14
|
+
|
|
15
|
+
RAG blocks (if any) are queried, session memory is queried when a
|
|
16
|
+
``chat_id`` is supplied, and the system prompt is the anchor. The
|
|
17
|
+
query is embedded once and reused for both retrievals.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
def __init__(
|
|
21
|
+
self,
|
|
22
|
+
store: StoreProtocol,
|
|
23
|
+
llm: LLMClientProtocol,
|
|
24
|
+
) -> None:
|
|
25
|
+
self._store = store
|
|
26
|
+
self._llm = llm
|
|
27
|
+
|
|
28
|
+
async def build(self, inp: ContextInput) -> ContextOutput:
|
|
29
|
+
parts: list[str] = [inp.system_prompt] if inp.system_prompt else []
|
|
30
|
+
rag_blocks: list[str] = []
|
|
31
|
+
session_blocks: list[str] = []
|
|
32
|
+
profile_used = False
|
|
33
|
+
|
|
34
|
+
query_emb: list[float] | None = None
|
|
35
|
+
if (inp.include_rag and inp.doc_ids) or (inp.include_session and inp.chat_id):
|
|
36
|
+
query_emb = (await self._llm.embed([inp.query], model=inp.embedding_model))[0]
|
|
37
|
+
|
|
38
|
+
if inp.include_rag and inp.doc_ids and query_emb is not None:
|
|
39
|
+
str_doc_ids = [str(d) for d in inp.doc_ids]
|
|
40
|
+
where = {"doc_id": {"$in": str_doc_ids}} if len(str_doc_ids) > 1 else {"doc_id": str_doc_ids[0]}
|
|
41
|
+
results = self._store.query(query_emb, top_k=inp.top_k_rag, where=where)
|
|
42
|
+
if results:
|
|
43
|
+
rag_texts = [r["document"] for r in results]
|
|
44
|
+
rag_blocks = rag_texts
|
|
45
|
+
parts.append("\n\n---\n\n".join(rag_texts))
|
|
46
|
+
|
|
47
|
+
if inp.include_session and inp.chat_id and query_emb is not None:
|
|
48
|
+
session_results = self._store.query(
|
|
49
|
+
query_emb,
|
|
50
|
+
top_k=inp.top_k_session,
|
|
51
|
+
where={"doc_id": f"session_{inp.chat_id}"},
|
|
52
|
+
)
|
|
53
|
+
if session_results:
|
|
54
|
+
session_texts = [r["document"] for r in session_results]
|
|
55
|
+
session_blocks = session_texts
|
|
56
|
+
parts.append(
|
|
57
|
+
"История диалога (сжатая):\n" + "\n---\n".join(session_texts)
|
|
58
|
+
)
|
|
59
|
+
|
|
60
|
+
if inp.include_profile and inp.profile_text:
|
|
61
|
+
parts.append(f"Профиль пользователя:\n{inp.profile_text}")
|
|
62
|
+
profile_used = True
|
|
63
|
+
|
|
64
|
+
return ContextOutput(
|
|
65
|
+
system_prompt="\n\n".join(parts),
|
|
66
|
+
rag_blocks=rag_blocks,
|
|
67
|
+
session_blocks=session_blocks,
|
|
68
|
+
profile_used=profile_used,
|
|
69
|
+
)
|
|
@@ -0,0 +1,219 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import logging
|
|
4
|
+
from dataclasses import dataclass, field
|
|
5
|
+
|
|
6
|
+
from protoprompt.context import ContextInput, ContextOutput
|
|
7
|
+
from protoprompt.exceptions import TokenBudgetExceededError
|
|
8
|
+
from protoprompt.injector import ContextBuilder
|
|
9
|
+
from protoprompt.llm import LLMClientProtocol
|
|
10
|
+
from protoprompt.store.protocol import StoreProtocol
|
|
11
|
+
from protoprompt.tokens.protocol import TokenCounter
|
|
12
|
+
from protoprompt.tokens.regex_counter import RegexTokenCounter
|
|
13
|
+
|
|
14
|
+
logger = logging.getLogger(__name__)
|
|
15
|
+
|
|
16
|
+
Priority = str
|
|
17
|
+
SEGMENT_RAG = "rag"
|
|
18
|
+
SEGMENT_SESSION = "session"
|
|
19
|
+
SEGMENT_PROFILE = "profile"
|
|
20
|
+
SEGMENT_SYSTEM = "system"
|
|
21
|
+
|
|
22
|
+
DEFAULT_PRIORITIES: tuple[Priority, ...] = (
|
|
23
|
+
SEGMENT_SYSTEM,
|
|
24
|
+
SEGMENT_PROFILE,
|
|
25
|
+
SEGMENT_SESSION,
|
|
26
|
+
SEGMENT_RAG,
|
|
27
|
+
)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@dataclass
|
|
31
|
+
class BudgetReport:
|
|
32
|
+
"""Observability for a single context build.
|
|
33
|
+
|
|
34
|
+
``used_tokens`` is the final size of the assembled ``system_prompt``
|
|
35
|
+
counted via the supplied ``TokenCounter``. ``dropped_blocks`` lists
|
|
36
|
+
block identifiers that did not fit; UI may surface this to the user.
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
used_tokens: int = 0
|
|
40
|
+
budget: int = 0
|
|
41
|
+
remaining_tokens: int = 0
|
|
42
|
+
dropped_blocks: list[str] = field(default_factory=list)
|
|
43
|
+
section_tokens: dict[str, int] = field(default_factory=dict)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
@dataclass
|
|
47
|
+
class _Candidate:
|
|
48
|
+
section: Priority
|
|
49
|
+
text: str
|
|
50
|
+
label: str
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class TokenBudgetedContextBuilder(ContextBuilder):
|
|
54
|
+
"""ContextBuilder that enforces a hard token ceiling on the final
|
|
55
|
+
``system_prompt``.
|
|
56
|
+
|
|
57
|
+
Behaviour:
|
|
58
|
+
1. ``system_prompt`` is mandatory and never truncated. If it does not
|
|
59
|
+
fit into ``max_tokens`` alone, ``TokenBudgetExceededError`` is
|
|
60
|
+
raised.
|
|
61
|
+
2. ``profile_text`` is appended in full when ``include_profile`` is
|
|
62
|
+
true; if it would push us over budget, the profile is dropped and
|
|
63
|
+
``dropped_blocks`` is updated (not raised — profile is a hint).
|
|
64
|
+
3. RAG and session blocks are pooled (top_k * 2) and allocated
|
|
65
|
+
greedily in priority order. The last accepted block is truncated
|
|
66
|
+
at a word boundary if it does not fit whole.
|
|
67
|
+
4. ``BudgetReport`` is attached to the returned ``ContextOutput``
|
|
68
|
+
via ``budget_report`` so the caller can surface a usage indicator.
|
|
69
|
+
"""
|
|
70
|
+
|
|
71
|
+
def __init__(
|
|
72
|
+
self,
|
|
73
|
+
store: StoreProtocol,
|
|
74
|
+
llm: LLMClientProtocol,
|
|
75
|
+
counter: TokenCounter | None = None,
|
|
76
|
+
max_tokens: int = 4096,
|
|
77
|
+
priorities: tuple[Priority, ...] = DEFAULT_PRIORITIES,
|
|
78
|
+
) -> None:
|
|
79
|
+
super().__init__(store, llm)
|
|
80
|
+
self._counter: TokenCounter = counter or RegexTokenCounter()
|
|
81
|
+
self._max_tokens = max_tokens
|
|
82
|
+
self._priorities = priorities
|
|
83
|
+
|
|
84
|
+
async def build(self, inp: ContextInput) -> ContextOutput:
|
|
85
|
+
report = BudgetReport(budget=self._max_tokens)
|
|
86
|
+
|
|
87
|
+
system_cost = self._counter.count(inp.system_prompt) if inp.system_prompt else 0
|
|
88
|
+
if system_cost > self._max_tokens:
|
|
89
|
+
raise TokenBudgetExceededError(system_cost, self._max_tokens, SEGMENT_SYSTEM)
|
|
90
|
+
report.section_tokens[SEGMENT_SYSTEM] = system_cost
|
|
91
|
+
remaining = self._max_tokens - system_cost
|
|
92
|
+
|
|
93
|
+
if inp.include_profile and inp.profile_text:
|
|
94
|
+
profile_block = f"Профиль пользователя:\n{inp.profile_text}"
|
|
95
|
+
profile_cost = self._counter.count(profile_block)
|
|
96
|
+
if profile_cost > remaining:
|
|
97
|
+
logger.warning(
|
|
98
|
+
"Profile block (%d tokens) exceeds remaining budget (%d); dropping",
|
|
99
|
+
profile_cost, remaining,
|
|
100
|
+
)
|
|
101
|
+
report.dropped_blocks.append(SEGMENT_PROFILE)
|
|
102
|
+
profile_block = ""
|
|
103
|
+
else:
|
|
104
|
+
report.section_tokens[SEGMENT_PROFILE] = profile_cost
|
|
105
|
+
remaining -= profile_cost
|
|
106
|
+
else:
|
|
107
|
+
profile_block = ""
|
|
108
|
+
|
|
109
|
+
pool: dict[Priority, list[_Candidate]] = {p: [] for p in self._priorities}
|
|
110
|
+
|
|
111
|
+
if (inp.include_rag and inp.doc_ids) or (inp.include_session and inp.chat_id):
|
|
112
|
+
query_emb = (await self._llm.embed([inp.query], model=inp.embedding_model))[0]
|
|
113
|
+
|
|
114
|
+
if inp.include_rag and inp.doc_ids:
|
|
115
|
+
str_doc_ids = [str(d) for d in inp.doc_ids]
|
|
116
|
+
where = (
|
|
117
|
+
{"doc_id": {"$in": str_doc_ids}}
|
|
118
|
+
if len(str_doc_ids) > 1
|
|
119
|
+
else {"doc_id": str_doc_ids[0]}
|
|
120
|
+
)
|
|
121
|
+
rag_hits = self._store.query(
|
|
122
|
+
query_emb, top_k=max(1, inp.top_k_rag * 2), where=where
|
|
123
|
+
)
|
|
124
|
+
for i, hit in enumerate(rag_hits):
|
|
125
|
+
pool[SEGMENT_RAG].append(_Candidate(
|
|
126
|
+
section=SEGMENT_RAG,
|
|
127
|
+
text=hit["document"],
|
|
128
|
+
label=f"rag[{i}]",
|
|
129
|
+
))
|
|
130
|
+
|
|
131
|
+
if inp.include_session and inp.chat_id:
|
|
132
|
+
session_hits = self._store.query(
|
|
133
|
+
query_emb,
|
|
134
|
+
top_k=max(1, inp.top_k_session * 2),
|
|
135
|
+
where={"doc_id": f"session_{inp.chat_id}"},
|
|
136
|
+
)
|
|
137
|
+
for i, hit in enumerate(session_hits):
|
|
138
|
+
pool[SEGMENT_SESSION].append(_Candidate(
|
|
139
|
+
section=SEGMENT_SESSION,
|
|
140
|
+
text=hit["document"],
|
|
141
|
+
label=f"session[{i}]",
|
|
142
|
+
))
|
|
143
|
+
|
|
144
|
+
kept_rag: list[str] = []
|
|
145
|
+
kept_session: list[str] = []
|
|
146
|
+
|
|
147
|
+
for section in self._priorities:
|
|
148
|
+
if section in (SEGMENT_SYSTEM, SEGMENT_PROFILE):
|
|
149
|
+
continue
|
|
150
|
+
if section not in pool or not pool[section]:
|
|
151
|
+
continue
|
|
152
|
+
for cand in pool[section]:
|
|
153
|
+
cost = self._counter.count(cand.text)
|
|
154
|
+
if cost <= remaining:
|
|
155
|
+
if cand.section == SEGMENT_RAG:
|
|
156
|
+
kept_rag.append(cand.text)
|
|
157
|
+
elif cand.section == SEGMENT_SESSION:
|
|
158
|
+
kept_session.append(cand.text)
|
|
159
|
+
remaining -= cost
|
|
160
|
+
report.section_tokens[cand.label] = cost
|
|
161
|
+
else:
|
|
162
|
+
trimmed = self._truncate_to_budget(cand.text, remaining)
|
|
163
|
+
if trimmed:
|
|
164
|
+
if cand.section == SEGMENT_RAG:
|
|
165
|
+
kept_rag.append(trimmed)
|
|
166
|
+
elif cand.section == SEGMENT_SESSION:
|
|
167
|
+
kept_session.append(trimmed)
|
|
168
|
+
report.section_tokens[cand.label] = remaining
|
|
169
|
+
remaining = 0
|
|
170
|
+
else:
|
|
171
|
+
report.dropped_blocks.append(cand.label)
|
|
172
|
+
break
|
|
173
|
+
if remaining <= 0:
|
|
174
|
+
for later_section in self._priorities:
|
|
175
|
+
if later_section <= section:
|
|
176
|
+
continue
|
|
177
|
+
for cand in pool.get(later_section, []):
|
|
178
|
+
report.dropped_blocks.append(cand.label)
|
|
179
|
+
break
|
|
180
|
+
|
|
181
|
+
report.used_tokens = self._max_tokens - remaining
|
|
182
|
+
report.remaining_tokens = remaining
|
|
183
|
+
|
|
184
|
+
parts: list[str] = []
|
|
185
|
+
if inp.system_prompt:
|
|
186
|
+
parts.append(inp.system_prompt)
|
|
187
|
+
if profile_block:
|
|
188
|
+
parts.append(profile_block)
|
|
189
|
+
if kept_rag:
|
|
190
|
+
parts.append("\n\n---\n\n".join(kept_rag))
|
|
191
|
+
if kept_session:
|
|
192
|
+
parts.append("История диалога (сжатая):\n" + "\n---\n".join(kept_session))
|
|
193
|
+
|
|
194
|
+
return ContextOutput(
|
|
195
|
+
system_prompt="\n\n".join(parts),
|
|
196
|
+
rag_blocks=kept_rag,
|
|
197
|
+
session_blocks=kept_session,
|
|
198
|
+
profile_used=bool(profile_block),
|
|
199
|
+
budget_report=report,
|
|
200
|
+
)
|
|
201
|
+
|
|
202
|
+
def _truncate_to_budget(self, text: str, budget: int) -> str:
|
|
203
|
+
"""Cut ``text`` so it fits into ``budget`` tokens, ending on a
|
|
204
|
+
word boundary. Returns empty string if no content fits.
|
|
205
|
+
"""
|
|
206
|
+
if budget <= 0:
|
|
207
|
+
return ""
|
|
208
|
+
words = text.split()
|
|
209
|
+
out: list[str] = []
|
|
210
|
+
used = 0
|
|
211
|
+
for w in words:
|
|
212
|
+
w_cost = self._counter.count(w) + 1
|
|
213
|
+
if used + w_cost > budget:
|
|
214
|
+
break
|
|
215
|
+
out.append(w)
|
|
216
|
+
used += w_cost
|
|
217
|
+
if not out:
|
|
218
|
+
return ""
|
|
219
|
+
return " ".join(out) + "…"
|
protoprompt/llm.py
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Protocol, runtime_checkable
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
@runtime_checkable
|
|
7
|
+
class LLMClientProtocol(Protocol):
|
|
8
|
+
async def chat(self, messages: list[dict], model: str = "", **options: object) -> str:
|
|
9
|
+
...
|
|
10
|
+
|
|
11
|
+
async def embed(self, texts: list[str], model: str = "") -> list[list[float]]:
|
|
12
|
+
...
|
protoprompt/pipeline.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import logging
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
from protoprompt.llm import LLMClientProtocol
|
|
7
|
+
from protoprompt.session.compressor import Compressor
|
|
8
|
+
from protoprompt.session.strategy import HeuristicStrategy, StrategyProtocol
|
|
9
|
+
from protoprompt.session.types import CompressedBlock, Session
|
|
10
|
+
from protoprompt.store.protocol import StoreProtocol
|
|
11
|
+
|
|
12
|
+
logger = logging.getLogger(__name__)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class Pipeline:
|
|
16
|
+
"""Orchestrate session compression and persistence into a vector store.
|
|
17
|
+
|
|
18
|
+
The pipeline decides when compression should run (``should_compress``)
|
|
19
|
+
and is the only place that should write compressed session data to
|
|
20
|
+
the store. The write is performed atomically: the new doc_id is
|
|
21
|
+
written first, then the old one is removed. If the process dies in
|
|
22
|
+
between, the next call overwrites both, so no chunk is lost.
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
def __init__(
|
|
26
|
+
self,
|
|
27
|
+
store: StoreProtocol,
|
|
28
|
+
llm: LLMClientProtocol,
|
|
29
|
+
strategy: StrategyProtocol | None = None,
|
|
30
|
+
compress_every_n: int = 10,
|
|
31
|
+
embedding_model: str = "nomic-embed-text",
|
|
32
|
+
) -> None:
|
|
33
|
+
self._store = store
|
|
34
|
+
self._llm = llm
|
|
35
|
+
self._compressor = Compressor(strategy or HeuristicStrategy())
|
|
36
|
+
self._compress_every_n = compress_every_n
|
|
37
|
+
self._embedding_model = embedding_model
|
|
38
|
+
|
|
39
|
+
async def compress_and_store(self, session: Session) -> list[CompressedBlock]:
|
|
40
|
+
if len(session.messages) < self._compress_every_n:
|
|
41
|
+
return []
|
|
42
|
+
|
|
43
|
+
blocks = await self._compressor.compress(session, self._llm)
|
|
44
|
+
if not blocks:
|
|
45
|
+
return []
|
|
46
|
+
|
|
47
|
+
doc_id = f"session_{session.chat_id}"
|
|
48
|
+
new_doc_id = f"{doc_id}_new"
|
|
49
|
+
texts = [b.text for b in blocks]
|
|
50
|
+
embeddings = await self._llm.embed(texts, model=self._embedding_model)
|
|
51
|
+
|
|
52
|
+
meta: dict[str, Any] = {
|
|
53
|
+
"chat_id": session.chat_id,
|
|
54
|
+
"strategy": session.strategy,
|
|
55
|
+
"message_count": len(session.messages),
|
|
56
|
+
}
|
|
57
|
+
self._store.add(new_doc_id, texts, embeddings, meta)
|
|
58
|
+
self._store.delete(doc_id)
|
|
59
|
+
self._store.add(doc_id, texts, embeddings, meta)
|
|
60
|
+
self._store.delete(new_doc_id)
|
|
61
|
+
|
|
62
|
+
logger.info(
|
|
63
|
+
"Compressed session %s: %d messages -> %d blocks",
|
|
64
|
+
session.chat_id,
|
|
65
|
+
len(session.messages),
|
|
66
|
+
len(blocks),
|
|
67
|
+
)
|
|
68
|
+
return blocks
|
|
69
|
+
|
|
70
|
+
def should_compress(self, message_count: int) -> bool:
|
|
71
|
+
return message_count >= self._compress_every_n
|
|
@@ -0,0 +1,62 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import logging
|
|
5
|
+
|
|
6
|
+
from protoprompt.llm import LLMClientProtocol
|
|
7
|
+
from protoprompt.profile.types import UserProfile
|
|
8
|
+
|
|
9
|
+
logger = logging.getLogger(__name__)
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class ProfileBuilder:
|
|
13
|
+
"""Build a structured user profile by asking the LLM to analyse turns.
|
|
14
|
+
|
|
15
|
+
The LLM is expected to return strict JSON; on any failure the builder
|
|
16
|
+
logs the original exception and returns a minimal profile with
|
|
17
|
+
``summary="Не удалось построить профиль"``.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
def __init__(self, llm: LLMClientProtocol) -> None:
|
|
21
|
+
self._llm = llm
|
|
22
|
+
|
|
23
|
+
async def build(self, user_id: str, messages: list[dict]) -> UserProfile:
|
|
24
|
+
if not messages:
|
|
25
|
+
return UserProfile(user_id=user_id)
|
|
26
|
+
|
|
27
|
+
user_texts = [m["content"] for m in messages if m.get("role") == "user"]
|
|
28
|
+
if not user_texts:
|
|
29
|
+
return UserProfile(user_id=user_id)
|
|
30
|
+
|
|
31
|
+
prompt = (
|
|
32
|
+
"Проанализируй сообщения пользователя. Выдели:\n"
|
|
33
|
+
"1. Стиль общения (кратко/развёрнуто, формально/неформально)\n"
|
|
34
|
+
"2. Предпочтения (любит списки, нарратив, технические детали)\n"
|
|
35
|
+
"3. Уровень экспертизы (новичок, средний, эксперт)\n"
|
|
36
|
+
"Ответ дай строго в формате JSON:\n"
|
|
37
|
+
'{"traits": {"style": "...", "expertise": "..."},'
|
|
38
|
+
'"preferences": {"format": "..."}, "summary": "..."}\n\n'
|
|
39
|
+
"Сообщения пользователя:\n" + "\n".join(user_texts[-20:])
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
try:
|
|
43
|
+
response = await self._llm.chat(
|
|
44
|
+
[{"role": "user", "content": prompt}],
|
|
45
|
+
temperature=0.3,
|
|
46
|
+
max_tokens=300,
|
|
47
|
+
)
|
|
48
|
+
data = json.loads(response)
|
|
49
|
+
return UserProfile(
|
|
50
|
+
user_id=user_id,
|
|
51
|
+
traits=data.get("traits", {}),
|
|
52
|
+
preferences=data.get("preferences", {}),
|
|
53
|
+
summary=data.get("summary", ""),
|
|
54
|
+
)
|
|
55
|
+
except Exception:
|
|
56
|
+
logger.warning(
|
|
57
|
+
"Failed to build profile for user_id=%s", user_id, exc_info=True
|
|
58
|
+
)
|
|
59
|
+
return UserProfile(
|
|
60
|
+
user_id=user_id,
|
|
61
|
+
summary="Не удалось построить профиль",
|
|
62
|
+
)
|
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
{
|
|
2
|
+
"$schema": "http://json-schema.org/draft-07/schema#",
|
|
3
|
+
"type": "object",
|
|
4
|
+
"properties": {
|
|
5
|
+
"traits": {
|
|
6
|
+
"type": "object",
|
|
7
|
+
"properties": {
|
|
8
|
+
"style": { "type": "string" },
|
|
9
|
+
"expertise": { "type": "string", "enum": ["beginner", "intermediate", "expert"] },
|
|
10
|
+
"verbosity": { "type": "string", "enum": ["concise", "balanced", "detailed"] },
|
|
11
|
+
"formality": { "type": "string", "enum": ["casual", "neutral", "formal"] }
|
|
12
|
+
}
|
|
13
|
+
},
|
|
14
|
+
"preferences": {
|
|
15
|
+
"type": "object",
|
|
16
|
+
"properties": {
|
|
17
|
+
"format": { "type": "string", "enum": ["bullets", "narrative", "code_heavy", "mixed"] },
|
|
18
|
+
"language": { "type": "string" },
|
|
19
|
+
"topics": { "type": "array", "items": { "type": "string" } }
|
|
20
|
+
}
|
|
21
|
+
},
|
|
22
|
+
"summary": { "type": "string" }
|
|
23
|
+
}
|
|
24
|
+
}
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass, field
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
@dataclass
|
|
7
|
+
class UserProfile:
|
|
8
|
+
user_id: str = ""
|
|
9
|
+
traits: dict[str, str] = field(default_factory=dict)
|
|
10
|
+
preferences: dict[str, str] = field(default_factory=dict)
|
|
11
|
+
summary: str = ""
|
|
12
|
+
updated_at: str = ""
|
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
from protoprompt.session.compressor import Compressor
|
|
2
|
+
from protoprompt.session.strategy import (
|
|
3
|
+
HeuristicStrategy,
|
|
4
|
+
LLMSummaryStrategy,
|
|
5
|
+
StrategyProtocol,
|
|
6
|
+
)
|
|
7
|
+
from protoprompt.session.types import CompressedBlock, Session
|
|
8
|
+
|
|
9
|
+
__all__ = [
|
|
10
|
+
"Compressor",
|
|
11
|
+
"HeuristicStrategy",
|
|
12
|
+
"LLMSummaryStrategy",
|
|
13
|
+
"StrategyProtocol",
|
|
14
|
+
"Session",
|
|
15
|
+
"CompressedBlock",
|
|
16
|
+
]
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import logging
|
|
4
|
+
|
|
5
|
+
from protoprompt.llm import LLMClientProtocol
|
|
6
|
+
from protoprompt.session.strategy import HeuristicStrategy, StrategyProtocol
|
|
7
|
+
from protoprompt.session.types import CompressedBlock, Session
|
|
8
|
+
|
|
9
|
+
logger = logging.getLogger(__name__)
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class Compressor:
|
|
13
|
+
def __init__(self, strategy: StrategyProtocol | None = None) -> None:
|
|
14
|
+
self._strategy = strategy or HeuristicStrategy()
|
|
15
|
+
|
|
16
|
+
async def compress(self, session: Session, llm: LLMClientProtocol) -> list[CompressedBlock]:
|
|
17
|
+
return await self._strategy.compress(session, llm)
|