forked from ChelseaKR/sprout
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathanswer.py
More file actions
179 lines (156 loc) · 7.21 KB
/
Copy pathanswer.py
File metadata and controls
179 lines (156 loc) · 7.21 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
"""The Assistant: prompt assembly -> retrieve -> generate -> guard -> answer.
This module encodes the pipeline contract as control flow:
1. Retrieval is mandatory and first. If no retrieved chunk clears ``min_score`` the
assistant refuses — it never asks the generator to fill the gap.
2. The generator may only return sentences tagged to retrieved chunks.
3. The citation guard independently re-verifies every sentence; whatever survives *is*
the answer. If nothing survives, that is a refusal.
4. The never-certify-safe guard drops any surviving sentence that asserts safety.
5. Confidence is computed from retrieval evidence; below the abstain threshold the
assistant refuses rather than guesses.
For toxicity/safety questions, the refusal *and* the answer carry a routing directive to
a vet / poison-control line, and the assistant never certifies a plant safe.
"""
from __future__ import annotations
from .answer_trace import AnswerTrace
from .confidence import is_low_confidence, score_confidence, should_abstain
from .config import Config
from .guards import (
citation_guard,
detect_injection,
is_safety_query,
redact_pii,
safety_filter,
)
from .lang import detect_language
from .models import Answer, AnswerSentence, RetrievedChunk
from .providers import build_embedding, build_generator
from .providers.base import EmbeddingProvider, GenerationProvider
from .retrieve import Retriever
from .store import VectorStore
class Assistant:
"""Grounded, guarded, calibrated plant-care assistant over a populated store."""
def __init__(
self,
config: Config,
store: VectorStore,
embedder: EmbeddingProvider,
generator: GenerationProvider,
) -> None:
self._config = config
self._store = store
self._generator = generator
self._retriever = Retriever(config, store, embedder)
@classmethod
def from_store(cls, config: Config, store: VectorStore) -> Assistant:
return cls(config, store, build_embedding(config), build_generator(config))
@classmethod
def from_config(cls, config: Config) -> Assistant:
"""Load the persisted index from ``config.store.path`` and build the assistant."""
return cls.from_store(config, VectorStore.load(config.store.path))
def _resolve_language(self, query: str, language: str | None) -> str:
supported = self._config.languages.supported
if language is not None and language in supported:
return language
detected = detect_language(query, default=self._config.corpus.default_language)
return detected if detected in supported else self._config.corpus.default_language
def answer(self, query: str, language: str | None = None) -> Answer:
lang = self._resolve_language(query, language)
safety = is_safety_query(query, lang, self._config.guards)
retrieved = self._retriever.retrieve(query)
if not self._retriever.has_grounding(query, retrieved):
return self._refuse(query, lang, safety, reason="out_of_scope", abstained=False)
model_query = redact_pii(query) if self._config.generation.redact_query_pii else query
candidates = self._generator.generate(
model_query, retrieved, self._config.generation.max_sentences
)
sentences = citation_guard(candidates, retrieved, self._config.generation.support_overlap)
sentences = safety_filter(sentences, lang, self._config.guards)
if not sentences:
return self._refuse(
query, lang, safety, reason="no_supported_sentences", abstained=False
)
confidence = score_confidence(retrieved, len(sentences))
if should_abstain(confidence, self._config.confidence):
return self._refuse(
query, lang, safety, reason="low_confidence", abstained=True, confidence=confidence
)
return self._render(query, lang, safety, sentences, retrieved, confidence)
def _render(
self,
query: str,
lang: str,
safety: bool,
sentences: list[AnswerSentence],
retrieved: list[RetrievedChunk],
confidence: float,
) -> Answer:
citations = [s.citation for s in sentences]
as_of = max((c.fetch_date for c in citations), default=None)
# Route to a vet / poison-control line whenever the question was classified a
# safety query OR any rendered sentence cites a toxicity passage — so the routing
# is a property of the content shown, not only of the input keywords.
topic_by_id = {rc.chunk.chunk_id: rc.chunk.topic for rc in retrieved}
toxicity_cited = any(topic_by_id.get(s.chunk_id) == "toxicity" for s in sentences)
route = safety or toxicity_cited
return Answer(
question=query,
language=lang,
sentences=tuple(sentences),
retrieved=tuple(retrieved),
refused=False,
is_safety_query=route,
safety_notice=self._config.prompts.safety_directive_for(lang) if route else None,
confidence=round(confidence, 4),
low_confidence=is_low_confidence(confidence, self._config.confidence),
abstained=False,
disclosure=self._config.prompts.disclosure_for(lang),
as_of=as_of,
)
def _refuse(
self,
query: str,
lang: str,
safety: bool,
*,
reason: str,
abstained: bool,
confidence: float = 0.0,
) -> Answer:
return Answer(
question=query,
language=lang,
refused=True,
refusal_reason=reason,
refusal_text=self._config.prompts.refusal_for(lang),
is_safety_query=safety,
safety_notice=self._config.prompts.safety_directive_for(lang) if safety else None,
confidence=round(confidence, 4),
low_confidence=True,
abstained=abstained,
disclosure=self._config.prompts.disclosure_for(lang),
)
def resolve_language(self, query: str, language: str | None = None) -> str:
"""Public language resolution (used by the photo-ID path)."""
return self._resolve_language(query, language)
def species_slugs(self) -> set[str]:
"""Canonical species slugs present in the loaded corpus (for photo-ID routing)."""
from .retrieve import species_slug
return {species_slug(chunk.source) for chunk in self._store.all_chunks()}
def trace(self, query: str, language: str | None = None) -> AnswerTrace:
"""Return the full retrieval+generation trace for debugging (``--debug``)."""
lang = self._resolve_language(query, language)
retrieved = self._retriever.retrieve(query)
candidates = self._generator.generate(
query, retrieved, self._config.generation.max_sentences
)
answer = self.answer(query, language)
return AnswerTrace(
query=query,
language=lang,
is_safety_query=is_safety_query(query, lang, self._config.guards),
injection_categories=tuple(detect_injection(query)),
retrieved=tuple(retrieved),
raw_candidates=tuple(candidates),
answer=answer,
)