dusha 0.0.21__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.
- dusha/__init__.py +3 -0
- dusha/affect.py +348 -0
- dusha/affect_semantic.py +71 -0
- dusha/api.py +155 -0
- dusha/api_openai.py +276 -0
- dusha/api_state.py +473 -0
- dusha/cli.py +392 -0
- dusha/config.py +550 -0
- dusha/context.py +212 -0
- dusha/database.py +226 -0
- dusha/decision.py +264 -0
- dusha/decision_worker.py +40 -0
- dusha/defaults.py +11 -0
- dusha/embedding.py +83 -0
- dusha/emotions.py +405 -0
- dusha/evergreen.py +397 -0
- dusha/identity.py +26 -0
- dusha/memory.py +274 -0
- dusha/memory_plugin.py +88 -0
- dusha/memory_provider.py +292 -0
- dusha/memory_worker.py +79 -0
- dusha/plugin_runtime.py +38 -0
- dusha/proactive.py +317 -0
- dusha/prompts.py +217 -0
- dusha/resources/defaults.json +81 -0
- dusha/resources/emotions.json +133 -0
- dusha/resources/prompts.json +49 -0
- dusha/resources.py +40 -0
- dusha/schema.py +144 -0
- dusha/serialization.py +37 -0
- dusha/service.py +552 -0
- dusha/timeutil.py +29 -0
- dusha-0.0.21.dist-info/METADATA +126 -0
- dusha-0.0.21.dist-info/RECORD +37 -0
- dusha-0.0.21.dist-info/WHEEL +4 -0
- dusha-0.0.21.dist-info/entry_points.txt +2 -0
- dusha-0.0.21.dist-info/licenses/LICENSE +201 -0
dusha/__init__.py
ADDED
dusha/affect.py
ADDED
|
@@ -0,0 +1,348 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import logging
|
|
5
|
+
import math
|
|
6
|
+
import threading
|
|
7
|
+
from collections.abc import Callable
|
|
8
|
+
from datetime import datetime
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
from . import emotions as _emotions
|
|
12
|
+
from . import prompts as _prompts
|
|
13
|
+
from .affect_semantic import SemanticAppraisal
|
|
14
|
+
from .config import AffectConfig
|
|
15
|
+
from .database import Database
|
|
16
|
+
from .serialization import compact_json
|
|
17
|
+
from .timeutil import isoformat, parse_time, utc_now
|
|
18
|
+
|
|
19
|
+
logger = logging.getLogger("dusha")
|
|
20
|
+
|
|
21
|
+
Decider = Callable[..., "str | None"]
|
|
22
|
+
PhraseMatcher = Callable[[str], dict[str, float] | None]
|
|
23
|
+
|
|
24
|
+
class AffectEngine:
|
|
25
|
+
def __init__(
|
|
26
|
+
self,
|
|
27
|
+
database: Database,
|
|
28
|
+
config: AffectConfig,
|
|
29
|
+
prompts: dict[str, Any] | None = None,
|
|
30
|
+
decision_increment: float | None = None,
|
|
31
|
+
appraisal: SemanticAppraisal | None = None,
|
|
32
|
+
):
|
|
33
|
+
self.database = database
|
|
34
|
+
self.config = config
|
|
35
|
+
self.appraisal = appraisal
|
|
36
|
+
self._lock = threading.RLock()
|
|
37
|
+
self.emotions = _emotions.resolve_emotions(config, decision_increment=decision_increment)
|
|
38
|
+
self.emotions_version = str(self.emotions["emotion_version"])
|
|
39
|
+
self.emotions_fingerprint = _emotions.fingerprint(self.emotions)
|
|
40
|
+
self.spec = {name: values.copy() for name, values in self.emotions["dimensions"].items()}
|
|
41
|
+
value_range = self.emotions["value_range"]
|
|
42
|
+
self.value_min = float(value_range["min"])
|
|
43
|
+
self.value_max = float(value_range["max"])
|
|
44
|
+
self.value_span = self.value_max - self.value_min
|
|
45
|
+
self.affect_presentation = dict(
|
|
46
|
+
(prompts or _prompts.default_prompts())["affect_presentation"]
|
|
47
|
+
)
|
|
48
|
+
if appraisal is not None:
|
|
49
|
+
appraisal.configure(self.emotions["appraisal"])
|
|
50
|
+
|
|
51
|
+
def initial_state(self) -> dict[str, Any]:
|
|
52
|
+
base = {name: values["neutral"] for name, values in self.spec.items()}
|
|
53
|
+
return {"base": base.copy(), "mood": base.copy()}
|
|
54
|
+
|
|
55
|
+
def _ensure(self, db: Any, now: datetime) -> None:
|
|
56
|
+
state = self.initial_state()
|
|
57
|
+
db.execute(
|
|
58
|
+
"""INSERT OR IGNORE INTO affect_state
|
|
59
|
+
(id, state_json, last_updated_at, last_interaction_at)
|
|
60
|
+
VALUES(1,?,?,?)""",
|
|
61
|
+
(compact_json(state, ensure_ascii=True), isoformat(now), isoformat(now)),
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
def _load_row(self, db: Any, now: datetime) -> tuple[Any, dict[str, Any]]:
|
|
65
|
+
self._ensure(db, now)
|
|
66
|
+
row = db.execute("SELECT * FROM affect_state WHERE id=1").fetchone()
|
|
67
|
+
stored = json.loads(row["state_json"])
|
|
68
|
+
neutral = self.initial_state()
|
|
69
|
+
return row, {
|
|
70
|
+
layer: {name: stored.get(layer, {}).get(name, default) for name, default in defaults.items()}
|
|
71
|
+
for layer, defaults in neutral.items()
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
def _clamp(self, value: float, floor: float | None = None) -> float:
|
|
75
|
+
return min(self.value_max, max(self.value_min if floor is None else floor, value))
|
|
76
|
+
|
|
77
|
+
def increment_for(self, emotion: str) -> float:
|
|
78
|
+
return float(self.spec[emotion].get("increment", self.emotions["decision"]["increment"]))
|
|
79
|
+
|
|
80
|
+
def _apply_deltas(self, state: dict[str, Any], deltas: dict[str, float], scale: float = 1.0) -> None:
|
|
81
|
+
base = state["base"]
|
|
82
|
+
mood = state["mood"]
|
|
83
|
+
negative = set(self.emotions["negative_dimensions"])
|
|
84
|
+
impact = float(self.emotions["impact_scale"])
|
|
85
|
+
for name, nominal in deltas.items():
|
|
86
|
+
if not isinstance(name, str) or name not in base or isinstance(nominal, bool):
|
|
87
|
+
continue
|
|
88
|
+
try:
|
|
89
|
+
delta = float(nominal) * scale
|
|
90
|
+
except (TypeError, ValueError):
|
|
91
|
+
continue
|
|
92
|
+
if not math.isfinite(delta):
|
|
93
|
+
continue
|
|
94
|
+
current = float(base[name])
|
|
95
|
+
room = self.value_max - current if delta > 0 else current - self.value_min
|
|
96
|
+
effective = delta * impact * room / self.value_span
|
|
97
|
+
next_value = max(current + effective, self.spec[name]["floor"])
|
|
98
|
+
if delta < 0 and name in negative:
|
|
99
|
+
next_value = max(next_value, min(current, float(mood[name])))
|
|
100
|
+
base[name] = self._clamp(next_value, self.spec[name]["floor"])
|
|
101
|
+
|
|
102
|
+
def _increase(self, state: dict[str, Any], emotion: str) -> None:
|
|
103
|
+
spec = self.spec[emotion]
|
|
104
|
+
current = float(state["base"][emotion])
|
|
105
|
+
state["base"][emotion] = self._clamp(
|
|
106
|
+
current + self.increment_for(emotion), spec["floor"]
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
def _semantic_emotion(self, message: str) -> str | None:
|
|
110
|
+
try:
|
|
111
|
+
weights = self.appraisal.weights(message)
|
|
112
|
+
except Exception:
|
|
113
|
+
return None
|
|
114
|
+
if not isinstance(weights, dict) or not weights:
|
|
115
|
+
return None
|
|
116
|
+
prototypes = self.emotions["appraisal"]["prototypes"]
|
|
117
|
+
blended: dict[str, float] = {}
|
|
118
|
+
for label, weight in weights.items():
|
|
119
|
+
if isinstance(weight, bool) or not isinstance(weight, (int, float)):
|
|
120
|
+
continue
|
|
121
|
+
for dimension, delta in prototypes.get(label, {}).get("deltas", {}).items():
|
|
122
|
+
blended[dimension] = blended.get(dimension, 0.0) + float(weight) * delta
|
|
123
|
+
positive = {name: value for name, value in blended.items() if value > 0 and name in self.spec}
|
|
124
|
+
if not positive:
|
|
125
|
+
return None
|
|
126
|
+
return max(positive, key=positive.get)
|
|
127
|
+
|
|
128
|
+
def _advance_values(self, state: dict[str, Any], hours: float) -> None:
|
|
129
|
+
if hours <= 0:
|
|
130
|
+
return
|
|
131
|
+
gain_cfg = self.emotions["mood_follow_gain"]
|
|
132
|
+
mood_follow_hours = self.emotions["affect"]["mood_follow_hours"]
|
|
133
|
+
mood_return_hours = self.emotions["affect"]["mood_return_hours"]
|
|
134
|
+
for name, params in self.spec.items():
|
|
135
|
+
base = float(state["base"].get(name, params["neutral"]))
|
|
136
|
+
mood = float(state["mood"].get(name, params["neutral"]))
|
|
137
|
+
deviation = abs(base - mood)
|
|
138
|
+
gain = max(
|
|
139
|
+
gain_cfg["min"], min(gain_cfg["max"], gain_cfg["factor"] * deviation / self.value_span)
|
|
140
|
+
)
|
|
141
|
+
follow = 1 - math.exp(-hours * gain / mood_follow_hours)
|
|
142
|
+
mood += (base - mood) * follow
|
|
143
|
+
mood = params["neutral"] + (mood - params["neutral"]) * math.exp(
|
|
144
|
+
-hours / mood_return_hours
|
|
145
|
+
)
|
|
146
|
+
floor = params["floor"]
|
|
147
|
+
state["mood"][name] = self._clamp(mood, floor)
|
|
148
|
+
state["base"][name] = self._clamp(
|
|
149
|
+
state["mood"][name] + (base - state["mood"][name]) * math.exp(-hours / params["tau"]),
|
|
150
|
+
floor,
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
def _advance(self, row: Any, state: dict[str, Any], now: datetime) -> None:
|
|
154
|
+
last_updated = parse_time(row["last_updated_at"])
|
|
155
|
+
hours = max(0.0, (now - last_updated).total_seconds() / 3600)
|
|
156
|
+
self._advance_values(state, hours)
|
|
157
|
+
last_user = parse_time(row["last_user_message_at"]) if row["last_user_message_at"] else None
|
|
158
|
+
if last_user and now > last_updated:
|
|
159
|
+
silence_start = max(last_updated, last_user)
|
|
160
|
+
silence_hours = max(0.0, (now - silence_start).total_seconds() / 3600)
|
|
161
|
+
total_silence = max(0.0, (now - last_user).total_seconds() / 3600)
|
|
162
|
+
for name, rule in self.emotions["silence"]["rules"].items():
|
|
163
|
+
if total_silence < rule["gate_hours"]:
|
|
164
|
+
continue
|
|
165
|
+
current = float(state["base"][name])
|
|
166
|
+
drifted = current + rule["rate_per_hour"] * silence_hours
|
|
167
|
+
state["base"][name] = max(current, min(self.spec[name]["neutral"] + rule["cap"], drifted))
|
|
168
|
+
|
|
169
|
+
def _save(self, db: Any, state: dict[str, Any], now: datetime, **fields: Any) -> None:
|
|
170
|
+
assignments = ["state_json=?", "last_updated_at=?", "revision=revision+1"]
|
|
171
|
+
values: list[Any] = [compact_json(state, ensure_ascii=True), isoformat(now)]
|
|
172
|
+
for name, value in fields.items():
|
|
173
|
+
assignments.append(f"{name}=?")
|
|
174
|
+
values.append(value)
|
|
175
|
+
db.execute(f"UPDATE affect_state SET {', '.join(assignments)} WHERE id=1", values)
|
|
176
|
+
|
|
177
|
+
def record_user_message(
|
|
178
|
+
self,
|
|
179
|
+
*,
|
|
180
|
+
message: str,
|
|
181
|
+
source_message_id: int | None,
|
|
182
|
+
decider: Decider | None,
|
|
183
|
+
instruction: str,
|
|
184
|
+
phrase_matcher: PhraseMatcher | None = None,
|
|
185
|
+
now: datetime | None = None,
|
|
186
|
+
) -> dict[str, Any] | None:
|
|
187
|
+
requested = now if now is not None else utc_now()
|
|
188
|
+
with self._lock, self.database.connect() as db:
|
|
189
|
+
row, state = self._load_row(db, requested)
|
|
190
|
+
current = max(requested, parse_time(row["last_updated_at"]))
|
|
191
|
+
self._advance(row, state, current)
|
|
192
|
+
db.execute(
|
|
193
|
+
"UPDATE proactive_events SET status='cancelled', updated_at=? "
|
|
194
|
+
"WHERE status IN ('pending','leased')",
|
|
195
|
+
(isoformat(current),),
|
|
196
|
+
)
|
|
197
|
+
self._save(db, state, current, last_interaction_at=isoformat(current))
|
|
198
|
+
db.execute(
|
|
199
|
+
"""UPDATE affect_state SET
|
|
200
|
+
last_user_message_at=CASE
|
|
201
|
+
WHEN last_user_message_at IS NULL OR last_user_message_at < ? THEN ?
|
|
202
|
+
ELSE last_user_message_at END,
|
|
203
|
+
unanswered_proactive=0
|
|
204
|
+
WHERE id=1""",
|
|
205
|
+
(isoformat(current), isoformat(current)),
|
|
206
|
+
)
|
|
207
|
+
|
|
208
|
+
updated = db.execute("SELECT * FROM affect_state WHERE id=1").fetchone()
|
|
209
|
+
snapshot = self._public_state(state, updated, current)
|
|
210
|
+
|
|
211
|
+
emotion = None
|
|
212
|
+
if decider is not None:
|
|
213
|
+
try:
|
|
214
|
+
emotion = decider(
|
|
215
|
+
message=message,
|
|
216
|
+
emotions=self.spec,
|
|
217
|
+
state=snapshot,
|
|
218
|
+
instruction=instruction,
|
|
219
|
+
)
|
|
220
|
+
except Exception:
|
|
221
|
+
emotion = None
|
|
222
|
+
if not isinstance(emotion, str) or emotion not in self.spec:
|
|
223
|
+
emotion = None
|
|
224
|
+
|
|
225
|
+
if emotion is None and self.appraisal is not None:
|
|
226
|
+
emotion = self._semantic_emotion(message)
|
|
227
|
+
|
|
228
|
+
decision_id = None
|
|
229
|
+
public = snapshot
|
|
230
|
+
if emotion is not None:
|
|
231
|
+
with self._lock, self.database.connect() as db:
|
|
232
|
+
row, state = self._load_row(db, requested)
|
|
233
|
+
phase_now = now if now is not None else utc_now()
|
|
234
|
+
current = max(phase_now, parse_time(row["last_updated_at"]))
|
|
235
|
+
self._advance(row, state, current)
|
|
236
|
+
self._increase(state, emotion)
|
|
237
|
+
self._save(db, state, current)
|
|
238
|
+
cursor = db.execute(
|
|
239
|
+
"""INSERT INTO affect_decisions
|
|
240
|
+
(source_message_id, emotion, increment, occurred_at)
|
|
241
|
+
VALUES(?,?,?,?)""",
|
|
242
|
+
(source_message_id, emotion, self.increment_for(emotion), isoformat(current)),
|
|
243
|
+
)
|
|
244
|
+
if cursor.lastrowid is not None:
|
|
245
|
+
decision_id = int(cursor.lastrowid)
|
|
246
|
+
updated = db.execute("SELECT * FROM affect_state WHERE id=1").fetchone()
|
|
247
|
+
public = self._public_state(state, updated, current)
|
|
248
|
+
|
|
249
|
+
if phrase_matcher is not None:
|
|
250
|
+
try:
|
|
251
|
+
deltas = phrase_matcher(message)
|
|
252
|
+
if isinstance(deltas, dict):
|
|
253
|
+
deltas = dict(deltas.items())
|
|
254
|
+
else:
|
|
255
|
+
deltas = None
|
|
256
|
+
except Exception:
|
|
257
|
+
deltas = None
|
|
258
|
+
unknown = sorted(str(name) for name in deltas or {} if name not in self.spec)
|
|
259
|
+
if unknown:
|
|
260
|
+
logger.warning("phrase deltas name unknown emotion dimension(s): %s", unknown)
|
|
261
|
+
if deltas:
|
|
262
|
+
with self._lock, self.database.connect() as db:
|
|
263
|
+
row, state = self._load_row(db, requested)
|
|
264
|
+
phase_now = now if now is not None else utc_now()
|
|
265
|
+
current = max(phase_now, parse_time(row["last_updated_at"]))
|
|
266
|
+
self._advance(row, state, current)
|
|
267
|
+
self._apply_deltas(state, deltas)
|
|
268
|
+
self._save(db, state, current)
|
|
269
|
+
updated = db.execute("SELECT * FROM affect_state WHERE id=1").fetchone()
|
|
270
|
+
public = self._public_state(state, updated, current)
|
|
271
|
+
if emotion is None:
|
|
272
|
+
return None
|
|
273
|
+
return {
|
|
274
|
+
"emotion": emotion,
|
|
275
|
+
"increment": self.increment_for(emotion),
|
|
276
|
+
"decision_id": decision_id,
|
|
277
|
+
"state": public,
|
|
278
|
+
}
|
|
279
|
+
|
|
280
|
+
def status(self, now: datetime | None = None) -> dict[str, Any]:
|
|
281
|
+
current = now or utc_now()
|
|
282
|
+
with self._lock, self.database.connect() as db:
|
|
283
|
+
row, state = self._load_row(db, current)
|
|
284
|
+
self._advance(row, state, current)
|
|
285
|
+
self._save(db, state, current)
|
|
286
|
+
updated = db.execute("SELECT * FROM affect_state WHERE id=1").fetchone()
|
|
287
|
+
return self._public_state(state, updated, current)
|
|
288
|
+
|
|
289
|
+
def on_proactive_sent(self, now: datetime | None = None) -> None:
|
|
290
|
+
current = now or utc_now()
|
|
291
|
+
with self._lock, self.database.connect() as db:
|
|
292
|
+
row, state = self._load_row(db, current)
|
|
293
|
+
self._advance(row, state, current)
|
|
294
|
+
self._apply_deltas(state, self.emotions["proactive_sent_deltas"])
|
|
295
|
+
self._save(
|
|
296
|
+
db,
|
|
297
|
+
state,
|
|
298
|
+
current,
|
|
299
|
+
last_proactive_sent_at=isoformat(current),
|
|
300
|
+
unanswered_proactive=int(row["unanswered_proactive"]) + 1,
|
|
301
|
+
)
|
|
302
|
+
|
|
303
|
+
def describe(self, snapshot: dict[str, Any]) -> str:
|
|
304
|
+
values = snapshot["base"]
|
|
305
|
+
prompt = self.emotions["prompt"]
|
|
306
|
+
presentation = self.affect_presentation
|
|
307
|
+
deviations = sorted(
|
|
308
|
+
((abs(values[name] - self.spec[name]["neutral"]), name, values[name]) for name in values),
|
|
309
|
+
reverse=True,
|
|
310
|
+
)
|
|
311
|
+
selected = [
|
|
312
|
+
(name, value)
|
|
313
|
+
for deviation, name, value in deviations
|
|
314
|
+
if deviation >= prompt["deviation_threshold"]
|
|
315
|
+
][: prompt["top_n"]]
|
|
316
|
+
for name, minimum in prompt["always_show"].items():
|
|
317
|
+
if values[name] >= minimum and not any(shown == name for shown, _ in selected):
|
|
318
|
+
selected.append((name, values[name]))
|
|
319
|
+
if not selected:
|
|
320
|
+
return presentation["baseline"]
|
|
321
|
+
labels = []
|
|
322
|
+
for name, value in selected:
|
|
323
|
+
level = (
|
|
324
|
+
presentation["level_high"]
|
|
325
|
+
if value >= prompt["level_high"]
|
|
326
|
+
else presentation["level_elevated"]
|
|
327
|
+
if value >= prompt["level_elevated"]
|
|
328
|
+
else presentation["level_noticeable"]
|
|
329
|
+
)
|
|
330
|
+
labels.append(f"{name}{presentation['level_connector']}{level}")
|
|
331
|
+
return (
|
|
332
|
+
presentation["prefix"]
|
|
333
|
+
+ presentation["separator"].join(labels)
|
|
334
|
+
+ presentation["suffix"]
|
|
335
|
+
)
|
|
336
|
+
|
|
337
|
+
def prompt_context(self, now: datetime | None = None) -> str:
|
|
338
|
+
return self.describe(self.status(now))
|
|
339
|
+
|
|
340
|
+
def _public_state(self, state: dict[str, Any], row: Any, now: datetime) -> dict[str, Any]:
|
|
341
|
+
return {
|
|
342
|
+
"base": {key: round(float(value), 4) for key, value in state["base"].items()},
|
|
343
|
+
"mood": {key: round(float(value), 4) for key, value in state["mood"].items()},
|
|
344
|
+
"last_updated_at": isoformat(now),
|
|
345
|
+
"last_user_message_at": row["last_user_message_at"] if row else None,
|
|
346
|
+
"last_proactive_sent_at": row["last_proactive_sent_at"] if row else None,
|
|
347
|
+
"unanswered_proactive": int(row["unanswered_proactive"]) if row else 0,
|
|
348
|
+
}
|
dusha/affect_semantic.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
"""Small semantic appraisal layer; no vector store or additional model runtime."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import logging
|
|
6
|
+
import math
|
|
7
|
+
import threading
|
|
8
|
+
import time
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
import httpx
|
|
12
|
+
|
|
13
|
+
from . import emotions as _emotions
|
|
14
|
+
from .embedding import OpenAIEmbeddingClient
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class SemanticAppraisal:
|
|
18
|
+
def __init__(self, client: OpenAIEmbeddingClient, spec: dict[str, Any] | None = None):
|
|
19
|
+
self.client = client
|
|
20
|
+
self._anchors: list[list[float]] = []
|
|
21
|
+
self._retry_at = 0.0
|
|
22
|
+
self._lock = threading.Lock()
|
|
23
|
+
self.configure(spec if spec is not None else _emotions.default_emotions()["appraisal"])
|
|
24
|
+
|
|
25
|
+
def configure(self, spec: dict[str, Any]) -> None:
|
|
26
|
+
with self._lock:
|
|
27
|
+
self.prototypes = {label: item["text"] for label, item in spec["prototypes"].items()}
|
|
28
|
+
self.min_similarity = float(spec["min_similarity"])
|
|
29
|
+
self.fallback_label = str(spec["fallback_label"])
|
|
30
|
+
self._anchors = []
|
|
31
|
+
|
|
32
|
+
def weights(self, text: str) -> dict[str, float] | None:
|
|
33
|
+
"""Return bounded blend weights, or None when the provider is unavailable."""
|
|
34
|
+
with self._lock:
|
|
35
|
+
if not self.prototypes or time.monotonic() < self._retry_at:
|
|
36
|
+
return None
|
|
37
|
+
try:
|
|
38
|
+
if not self._anchors:
|
|
39
|
+
anchors = []
|
|
40
|
+
texts = list(self.prototypes.values())
|
|
41
|
+
size = max(1, self.client.config.batch_size)
|
|
42
|
+
for start in range(0, len(texts), size):
|
|
43
|
+
anchors.extend(self.client.embed(texts[start : start + size]))
|
|
44
|
+
self._anchors = anchors
|
|
45
|
+
query = self.client.embed([text])[0]
|
|
46
|
+
floor = self.min_similarity
|
|
47
|
+
if any(len(anchor) != len(query) for anchor in self._anchors):
|
|
48
|
+
self._anchors = []
|
|
49
|
+
raise ValueError("affect embedding dimensions changed")
|
|
50
|
+
scores = {
|
|
51
|
+
label: max(
|
|
52
|
+
0.0,
|
|
53
|
+
min(
|
|
54
|
+
1.0,
|
|
55
|
+
(math.fsum(a * b for a, b in zip(query, anchor, strict=True)) - floor)
|
|
56
|
+
/ (1 - floor),
|
|
57
|
+
),
|
|
58
|
+
)
|
|
59
|
+
for label, anchor in zip(self.prototypes, self._anchors, strict=True)
|
|
60
|
+
}
|
|
61
|
+
total = sum(scores.values())
|
|
62
|
+
# Weak matches leave their remaining mass to the fallback label.
|
|
63
|
+
weights = {label: score / max(1.0, total) for label, score in scores.items() if score > 0}
|
|
64
|
+
if self.fallback_label:
|
|
65
|
+
fallback = self.fallback_label
|
|
66
|
+
weights[fallback] = weights.get(fallback, 0.0) + max(0.0, 1 - total)
|
|
67
|
+
return weights
|
|
68
|
+
except (httpx.HTTPError, ValueError, TypeError, KeyError, IndexError):
|
|
69
|
+
self._retry_at = time.monotonic() + max(0, self.client.config.failure_cooldown_seconds)
|
|
70
|
+
logging.getLogger(__name__).warning("Semantic affect unavailable; using phrase matching")
|
|
71
|
+
return None
|
dusha/api.py
ADDED
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
import contextlib
|
|
5
|
+
import logging
|
|
6
|
+
import os
|
|
7
|
+
import secrets
|
|
8
|
+
from contextlib import asynccontextmanager
|
|
9
|
+
from importlib import metadata as _metadata
|
|
10
|
+
from typing import Any
|
|
11
|
+
|
|
12
|
+
from fastapi import FastAPI, Header, HTTPException, Request
|
|
13
|
+
from fastapi.responses import JSONResponse
|
|
14
|
+
|
|
15
|
+
from . import schema as _schema
|
|
16
|
+
from .api_openai import create_openai_router
|
|
17
|
+
from .api_state import create_state_router
|
|
18
|
+
from .config import AppConfig, load_config
|
|
19
|
+
from .proactive import ProactiveEngine
|
|
20
|
+
from .service import CompanionService, StorageUnavailableError
|
|
21
|
+
|
|
22
|
+
logger = logging.getLogger("dusha")
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class AuthConfigError(ValueError):
|
|
26
|
+
pass
|
|
27
|
+
|
|
28
|
+
def _resolve_token(cfg: AppConfig) -> str:
|
|
29
|
+
if not cfg.api_token_env:
|
|
30
|
+
return ""
|
|
31
|
+
value = os.getenv(cfg.api_token_env, "")
|
|
32
|
+
if not value:
|
|
33
|
+
raise AuthConfigError(
|
|
34
|
+
f"api_token_env is set to {cfg.api_token_env!r} but that environment variable "
|
|
35
|
+
"is missing or empty. Export it or clear api_token_env to disable auth"
|
|
36
|
+
)
|
|
37
|
+
return value
|
|
38
|
+
|
|
39
|
+
async def _scheduler(proactive: ProactiveEngine, interval: int) -> None:
|
|
40
|
+
while True:
|
|
41
|
+
try:
|
|
42
|
+
await asyncio.to_thread(proactive.evaluate)
|
|
43
|
+
except asyncio.CancelledError:
|
|
44
|
+
raise
|
|
45
|
+
except Exception:
|
|
46
|
+
logger.exception("proactive evaluation failed")
|
|
47
|
+
await asyncio.sleep(max(5, interval))
|
|
48
|
+
|
|
49
|
+
async def _ingest_scheduler(service: CompanionService, interval: int) -> None:
|
|
50
|
+
while True:
|
|
51
|
+
try:
|
|
52
|
+
await asyncio.to_thread(service.catch_up_plugin)
|
|
53
|
+
except asyncio.CancelledError:
|
|
54
|
+
raise
|
|
55
|
+
except Exception:
|
|
56
|
+
logger.exception("memory plugin ingestion catch-up failed")
|
|
57
|
+
await asyncio.sleep(max(1, interval))
|
|
58
|
+
|
|
59
|
+
async def _index_backfill_scheduler(service: CompanionService, interval: int) -> None:
|
|
60
|
+
while True:
|
|
61
|
+
try:
|
|
62
|
+
result = await asyncio.to_thread(service.memory_index_backfill)
|
|
63
|
+
if result.get("error") and not result.get("cooling_down"):
|
|
64
|
+
logger.warning("memory index backfill unavailable: %s", result["error"])
|
|
65
|
+
except asyncio.CancelledError:
|
|
66
|
+
raise
|
|
67
|
+
except Exception:
|
|
68
|
+
logger.exception("memory index backfill failed")
|
|
69
|
+
await asyncio.sleep(max(1, interval))
|
|
70
|
+
|
|
71
|
+
def create_app(config: AppConfig | None = None) -> FastAPI:
|
|
72
|
+
cfg = config or load_config()
|
|
73
|
+
expected_token = _resolve_token(cfg)
|
|
74
|
+
auth_enabled = bool(expected_token)
|
|
75
|
+
service = CompanionService(cfg)
|
|
76
|
+
proactive = ProactiveEngine(service, cfg)
|
|
77
|
+
|
|
78
|
+
@asynccontextmanager
|
|
79
|
+
async def lifespan(_: FastAPI):
|
|
80
|
+
tasks = [
|
|
81
|
+
asyncio.create_task(
|
|
82
|
+
_scheduler(proactive, cfg.proactive.poll_interval_seconds),
|
|
83
|
+
name="proactive-evaluator",
|
|
84
|
+
)
|
|
85
|
+
]
|
|
86
|
+
if service.memory.enabled and cfg.storage.enabled:
|
|
87
|
+
tasks.append(
|
|
88
|
+
asyncio.create_task(
|
|
89
|
+
_ingest_scheduler(
|
|
90
|
+
service, cfg.memory_plugin.ingest_backfill_interval_seconds
|
|
91
|
+
),
|
|
92
|
+
name="memory-plugin-ingest",
|
|
93
|
+
)
|
|
94
|
+
)
|
|
95
|
+
if service.memory.enabled:
|
|
96
|
+
tasks.append(
|
|
97
|
+
asyncio.create_task(
|
|
98
|
+
_index_backfill_scheduler(service, cfg.embedding.backfill_interval_seconds),
|
|
99
|
+
name="memory-plugin-backfill",
|
|
100
|
+
)
|
|
101
|
+
)
|
|
102
|
+
try:
|
|
103
|
+
yield
|
|
104
|
+
finally:
|
|
105
|
+
for task in tasks:
|
|
106
|
+
task.cancel()
|
|
107
|
+
for task in tasks:
|
|
108
|
+
with contextlib.suppress(asyncio.CancelledError):
|
|
109
|
+
await task
|
|
110
|
+
service.close()
|
|
111
|
+
|
|
112
|
+
app = FastAPI(
|
|
113
|
+
title="Dusha",
|
|
114
|
+
version=_metadata.version("dusha"),
|
|
115
|
+
description="https://github.com/Somme4096/dusha",
|
|
116
|
+
lifespan=lifespan,
|
|
117
|
+
docs_url=None if auth_enabled else "/docs",
|
|
118
|
+
redoc_url=None if auth_enabled else "/redoc",
|
|
119
|
+
openapi_url=None if auth_enabled else "/openapi.json",
|
|
120
|
+
openapi_tags=[
|
|
121
|
+
{"name": "health", "description": "https://github.com/Somme4096/dusha"},
|
|
122
|
+
{"name": "state", "description": "https://github.com/Somme4096/dusha"},
|
|
123
|
+
{"name": "proxy", "description": "https://github.com/Somme4096/dusha"},
|
|
124
|
+
],
|
|
125
|
+
)
|
|
126
|
+
app.state.config = cfg
|
|
127
|
+
app.state.service = service
|
|
128
|
+
app.state.proactive = proactive
|
|
129
|
+
|
|
130
|
+
async def authorized(x_companion_token: str = Header(default="")) -> None:
|
|
131
|
+
if not auth_enabled:
|
|
132
|
+
return
|
|
133
|
+
if not x_companion_token or not secrets.compare_digest(
|
|
134
|
+
x_companion_token.encode("utf-8"), expected_token.encode("utf-8")
|
|
135
|
+
):
|
|
136
|
+
raise HTTPException(status_code=401, detail="invalid companion token")
|
|
137
|
+
|
|
138
|
+
@app.exception_handler(StorageUnavailableError)
|
|
139
|
+
async def storage_unavailable(_: Request, error: StorageUnavailableError) -> JSONResponse:
|
|
140
|
+
return JSONResponse(status_code=503, content={"detail": str(error)})
|
|
141
|
+
|
|
142
|
+
@app.get("/health", tags=["health"], response_model=_schema.HealthResponse)
|
|
143
|
+
async def health() -> dict[str, Any]:
|
|
144
|
+
return {
|
|
145
|
+
"status": "ok",
|
|
146
|
+
"companion": cfg.home.name if cfg.home else "",
|
|
147
|
+
"database": await asyncio.to_thread(service.database.integrity_check),
|
|
148
|
+
"upstream_configured": bool(cfg.upstream.base_url),
|
|
149
|
+
"memory_index": await asyncio.to_thread(service.memory_index_status),
|
|
150
|
+
}
|
|
151
|
+
|
|
152
|
+
app.include_router(create_state_router(service, proactive, authorized))
|
|
153
|
+
if cfg.api_openai.enabled:
|
|
154
|
+
app.include_router(create_openai_router(cfg, service, authorized))
|
|
155
|
+
return app
|