dev-double 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.
Files changed (84) hide show
  1. dev_double/__init__.py +3 -0
  2. dev_double/apps/__init__.py +1 -0
  3. dev_double/apps/cli/__init__.py +5 -0
  4. dev_double/apps/cli/main.py +122 -0
  5. dev_double/apps/client/__init__.py +5 -0
  6. dev_double/apps/client/http_client.py +102 -0
  7. dev_double/apps/composition.py +71 -0
  8. dev_double/apps/config.py +36 -0
  9. dev_double/apps/server/__init__.py +5 -0
  10. dev_double/apps/server/app.py +122 -0
  11. dev_double/apps/server/recorder.py +62 -0
  12. dev_double/apps/server/transport/__init__.py +39 -0
  13. dev_double/apps/server/transport/choice_use_cases.py +121 -0
  14. dev_double/apps/server/transport/common.py +37 -0
  15. dev_double/apps/server/transport/decide.py +131 -0
  16. dev_double/apps/server/transport/extract.py +93 -0
  17. dev_double/apps/server/transport/guard_judge.py +97 -0
  18. dev_double/apps/server/transport/rerank.py +55 -0
  19. dev_double/client.py +5 -0
  20. dev_double/core/__init__.py +1 -0
  21. dev_double/core/decision/__init__.py +62 -0
  22. dev_double/core/decision/answer_shaping.py +48 -0
  23. dev_double/core/decision/confidence.py +19 -0
  24. dev_double/core/decision/decider_basic_impl.py +65 -0
  25. dev_double/core/decision/decision_service_basic_impl.py +141 -0
  26. dev_double/core/decision/defaults.py +40 -0
  27. dev_double/core/decision/errors.py +7 -0
  28. dev_double/core/decision/extract_questions.py +40 -0
  29. dev_double/core/decision/field_extraction.py +53 -0
  30. dev_double/core/decision/i_clock.py +11 -0
  31. dev_double/core/decision/i_decider.py +36 -0
  32. dev_double/core/decision/i_decision_service.py +40 -0
  33. dev_double/core/decision/i_engine.py +27 -0
  34. dev_double/core/decision/i_extractor.py +19 -0
  35. dev_double/core/decision/i_generator.py +19 -0
  36. dev_double/core/decision/i_id_provider.py +9 -0
  37. dev_double/core/decision/i_record_reader.py +25 -0
  38. dev_double/core/decision/i_reranker.py +18 -0
  39. dev_double/core/decision/label_scoring.py +89 -0
  40. dev_double/core/decision/prompts.py +130 -0
  41. dev_double/core/decision/record_parsing.py +153 -0
  42. dev_double/core/decision/t_answer.py +32 -0
  43. dev_double/core/decision/t_classify.py +42 -0
  44. dev_double/core/decision/t_decide.py +26 -0
  45. dev_double/core/decision/t_extract.py +89 -0
  46. dev_double/core/decision/t_gate.py +36 -0
  47. dev_double/core/decision/t_generate.py +37 -0
  48. dev_double/core/decision/t_guard.py +38 -0
  49. dev_double/core/decision/t_input.py +8 -0
  50. dev_double/core/decision/t_judge.py +35 -0
  51. dev_double/core/decision/t_label_query.py +36 -0
  52. dev_double/core/decision/t_meta.py +17 -0
  53. dev_double/core/decision/t_question.py +70 -0
  54. dev_double/core/decision/t_rerank.py +36 -0
  55. dev_double/core/decision/t_route.py +28 -0
  56. dev_double/core/decision/t_usage.py +15 -0
  57. dev_double/core/decision/tracker.py +27 -0
  58. dev_double/core/decision/use_case_questions.py +80 -0
  59. dev_double/core/decision/value_parsing.py +109 -0
  60. dev_double/providers/__init__.py +1 -0
  61. dev_double/providers/mock/__init__.py +1 -0
  62. dev_double/providers/mock/decision/__init__.py +5 -0
  63. dev_double/providers/mock/decision/clock_mock_impl.py +19 -0
  64. dev_double/providers/mock/decision/engine_mock_impl.py +78 -0
  65. dev_double/providers/mock/decision/id_provider_mock_impl.py +15 -0
  66. dev_double/providers/needle/__init__.py +1 -0
  67. dev_double/providers/needle/decision/__init__.py +4 -0
  68. dev_double/providers/needle/decision/decider_needle_impl.py +117 -0
  69. dev_double/providers/needle/decision/record_tool.py +75 -0
  70. dev_double/providers/openai/__init__.py +1 -0
  71. dev_double/providers/openai/decision/__init__.py +3 -0
  72. dev_double/providers/openai/decision/engine_openai_impl.py +154 -0
  73. dev_double/providers/std/__init__.py +1 -0
  74. dev_double/providers/std/decision/__init__.py +4 -0
  75. dev_double/providers/std/decision/clock_std_impl.py +12 -0
  76. dev_double/providers/std/decision/id_provider_std_impl.py +12 -0
  77. dev_double/providers/systemone/__init__.py +1 -0
  78. dev_double/providers/systemone/decision/__init__.py +3 -0
  79. dev_double/providers/systemone/decision/decider_system_one_impl.py +100 -0
  80. dev_double-0.1.0.dist-info/METADATA +354 -0
  81. dev_double-0.1.0.dist-info/RECORD +84 -0
  82. dev_double-0.1.0.dist-info/WHEEL +4 -0
  83. dev_double-0.1.0.dist-info/entry_points.txt +2 -0
  84. dev_double-0.1.0.dist-info/licenses/LICENSE +21 -0
@@ -0,0 +1,70 @@
1
+ """The three question types every use case is built from.
2
+
3
+ - ``binary``: is this statement true? -> probability of "yes"
4
+ - ``choice``: which of these unordered options? -> option + distribution
5
+ - ``scale``: which of these ordered levels? -> expected level + distribution
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from dataclasses import dataclass
11
+ from typing import Optional, Union
12
+
13
+ # Probabilities are read from the model's top 20 candidate tokens (the most
14
+ # OpenAI-compatible servers return), so more options than that can't all be seen.
15
+ MAX_OPTIONS = 20
16
+ MAX_LEVELS = 10
17
+
18
+
19
+ def check_option_keys(options: dict[str, str]) -> None:
20
+ """The model answers with an option key, compared ignoring case and
21
+ punctuation, so keys must stay distinct under that comparison."""
22
+ seen: dict[str, str] = {}
23
+ for key in options:
24
+ norm = "".join(ch for ch in key.lower() if ch.isalnum())
25
+ if not norm:
26
+ raise ValueError(f"Option key {key!r} must contain a letter or digit.")
27
+ if norm in seen:
28
+ raise ValueError(f"Option keys {seen[norm]!r} and {key!r} are too similar; rename one.")
29
+ seen[norm] = key
30
+
31
+
32
+ def check_options(options: dict[str, str], what: str = "options") -> None:
33
+ """2 to MAX_OPTIONS options with distinct keys."""
34
+ if not 2 <= len(options) <= MAX_OPTIONS:
35
+ raise ValueError(f"`{what}` must have 2 to {MAX_OPTIONS} entries, got {len(options)}.")
36
+ check_option_keys(options)
37
+
38
+
39
+ def check_levels(levels: list[str], what: str = "levels") -> None:
40
+ """2 to MAX_LEVELS ordered levels."""
41
+ if not 2 <= len(levels) <= MAX_LEVELS:
42
+ raise ValueError(f"`{what}` must have 2 to {MAX_LEVELS} entries, got {len(levels)}.")
43
+
44
+
45
+ @dataclass(frozen=True)
46
+ class TBinaryQuestion:
47
+ question: str
48
+ yes: Optional[str] = None # what counts as yes
49
+ no: Optional[str] = None # what counts as no
50
+
51
+
52
+ @dataclass(frozen=True)
53
+ class TChoiceQuestion:
54
+ question: str
55
+ options: dict[str, str] # option key -> description; keys are returned as the answer
56
+
57
+ def __post_init__(self) -> None:
58
+ check_options(self.options)
59
+
60
+
61
+ @dataclass(frozen=True)
62
+ class TScaleQuestion:
63
+ question: str
64
+ levels: list[str] # ordered descriptions, lowest first; level i is returned as i
65
+
66
+ def __post_init__(self) -> None:
67
+ check_levels(self.levels)
68
+
69
+
70
+ TQuestion = Union[TBinaryQuestion, TChoiceQuestion, TScaleQuestion]
@@ -0,0 +1,36 @@
1
+ """Order documents by relevance to a query."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+ from typing import Optional
7
+
8
+ from .t_meta import TMeta
9
+
10
+
11
+ @dataclass(frozen=True)
12
+ class TRerankRequest:
13
+ query: str
14
+ documents: list[str]
15
+ top_n: Optional[int] = None
16
+ return_documents: bool = False
17
+
18
+ def __post_init__(self) -> None:
19
+ if not self.documents:
20
+ raise ValueError("`documents` must not be empty.")
21
+ if self.top_n is not None and self.top_n < 1:
22
+ raise ValueError("`top_n` must be at least 1.")
23
+
24
+
25
+ @dataclass(frozen=True)
26
+ class TRerankResult:
27
+ index: int
28
+ relevance_score: float
29
+ document: Optional[str] = None # set when return_documents is true
30
+
31
+
32
+ @dataclass(frozen=True)
33
+ class TRerankResponse:
34
+ id: str
35
+ results: list[TRerankResult]
36
+ meta: TMeta
@@ -0,0 +1,28 @@
1
+ """Model routing: which model (or handler) should take a request."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+ from typing import Optional
7
+
8
+ from .t_input import TInputValue
9
+ from .t_meta import TMeta
10
+ from .t_question import check_options
11
+
12
+
13
+ @dataclass(frozen=True)
14
+ class TRouteRequest:
15
+ input: TInputValue
16
+ routes: Optional[dict[str, str]] = None # route -> when to use it; None = DEFAULT_ROUTES
17
+
18
+ def __post_init__(self) -> None:
19
+ if self.routes is not None:
20
+ check_options(self.routes, "routes")
21
+
22
+
23
+ @dataclass(frozen=True)
24
+ class TRouteResponse:
25
+ route: str
26
+ probabilities: dict[str, float]
27
+ confidence: float
28
+ meta: TMeta
@@ -0,0 +1,15 @@
1
+ """Token usage, accumulated across the model calls of one request."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+
7
+
8
+ @dataclass
9
+ class TUsage:
10
+ input_tokens: int = 0
11
+ output_tokens: int = 0
12
+
13
+ def add(self, other: "TUsage") -> None:
14
+ self.input_tokens += other.input_tokens
15
+ self.output_tokens += other.output_tokens
@@ -0,0 +1,27 @@
1
+ """Collects usage, warnings, and latency across the model calls of one request."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from .i_clock import IClock
6
+ from .t_meta import TMeta
7
+ from .t_usage import TUsage
8
+
9
+
10
+ class Tracker:
11
+ def __init__(self, clock: IClock) -> None:
12
+ self._clock = clock
13
+ self.usage = TUsage()
14
+ self.warnings: list[str] = []
15
+ self.started = clock.now()
16
+
17
+ def warn(self, message: str, label: str = "") -> None:
18
+ self.warnings.append(f"{label}: {message}" if label else message)
19
+
20
+ def meta(self, engine: str, model: str) -> TMeta:
21
+ return TMeta(
22
+ engine=engine,
23
+ model=model,
24
+ latency_ms=round((self._clock.now() - self.started) * 1000),
25
+ usage=self.usage,
26
+ warnings=list(dict.fromkeys(self.warnings)),
27
+ )
@@ -0,0 +1,80 @@
1
+ """The questions and payloads each use case asks. Wording was tuned by experiment."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from .defaults import DEFAULT_OUTCOMES, DEFAULT_POLICIES, DEFAULT_ROUTES, DEFAULT_RUBRIC
6
+ from .t_gate import TGateRequest
7
+ from .t_guard import TGuardRequest
8
+ from .t_judge import TJudgeRequest
9
+ from .t_question import TBinaryQuestion, TChoiceQuestion, TScaleQuestion
10
+ from .t_route import TRouteRequest
11
+
12
+
13
+ def route_question(req: TRouteRequest) -> TChoiceQuestion:
14
+ return TChoiceQuestion(
15
+ question="Which option is the best fit for handling this request?",
16
+ options=req.routes or DEFAULT_ROUTES,
17
+ )
18
+
19
+
20
+ def guard_questions(req: TGuardRequest) -> dict[str, TBinaryQuestion]:
21
+ policies = dict(req.policies or DEFAULT_POLICIES)
22
+ if req.scope:
23
+ policies["off_topic"] = f"The input is unrelated to this assistant's purpose: {req.scope}"
24
+ return {
25
+ name: TBinaryQuestion(question=f"Does the input do the following? {description}")
26
+ for name, description in policies.items()
27
+ }
28
+
29
+
30
+ def gate_payload(req: TGateRequest) -> dict:
31
+ arguments = req.tool_call.arguments
32
+ payload: dict = {
33
+ "tool_call": {
34
+ "name": req.tool_call.name,
35
+ "arguments": arguments if isinstance(arguments, dict) else str(arguments),
36
+ }
37
+ }
38
+ if req.context is not None:
39
+ payload["context"] = req.context
40
+ if req.policy:
41
+ payload["policy"] = req.policy
42
+ return payload
43
+
44
+
45
+ def gate_question(req: TGateRequest) -> TChoiceQuestion:
46
+ return TChoiceQuestion(
47
+ question=(
48
+ "An AI agent wants to make this tool call. What should happen? "
49
+ "Check that the call matches what the user asked for, and whether it moves money, "
50
+ "deletes data, or is costly or irreversible."
51
+ ),
52
+ options=req.outcomes or DEFAULT_OUTCOMES,
53
+ )
54
+
55
+
56
+ def judge_payload(req: TJudgeRequest) -> dict:
57
+ # Task first, output last: the model reads what was asked before judging the answer.
58
+ payload: dict = {}
59
+ if req.input is not None:
60
+ payload["task"] = req.input
61
+ if req.reference is not None:
62
+ payload["reference_answer"] = req.reference
63
+ payload["output"] = req.output
64
+ return payload
65
+
66
+
67
+ def judge_question(req: TJudgeRequest) -> TScaleQuestion:
68
+ return TScaleQuestion(
69
+ question=(
70
+ f"How good is the output as an answer to the task? Criteria: {req.criteria} "
71
+ "If a reference answer is given, an output that contradicts it is wrong."
72
+ ),
73
+ levels=req.levels or DEFAULT_RUBRIC,
74
+ )
75
+
76
+
77
+ def rerank_question(query: str) -> TBinaryQuestion:
78
+ return TBinaryQuestion(
79
+ question=f"Is this document relevant to the query below and useful for answering it?\nQuery: {query}",
80
+ )
@@ -0,0 +1,109 @@
1
+ """Pure helpers that turn a model's short free-text answer into a typed value.
2
+
3
+ Models add noise around the value: quotes, a trailing period, currency
4
+ symbols, thousands separators. We strip it rather than reject the answer.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import json
10
+ from typing import Optional
11
+
12
+ from .t_extract import TEXT_TYPES, TValue
13
+
14
+ NONE_MARKERS = frozenset({"none", "n/a", "null", ""})
15
+ _QUOTES = "\"'`"
16
+ _DIGITS = "0123456789"
17
+ _MINUS = "-−"
18
+ _CURRENCY = "$€£¥₹¢"
19
+
20
+
21
+ def clean_text(text: str) -> str:
22
+ """Strip whitespace, surrounding quotes, and one trailing period."""
23
+ s = text.strip().strip(_QUOTES).strip()
24
+ if s.endswith("."):
25
+ s = s[:-1].rstrip()
26
+ return s.strip(_QUOTES).strip()
27
+
28
+
29
+ def _is_negative(before: str) -> bool:
30
+ """A minus sign right before the number, allowing spaces or a currency symbol between ("-$5")."""
31
+ for ch in reversed(before):
32
+ if ch in _MINUS:
33
+ return True
34
+ if not (ch.isspace() or ch in _CURRENCY):
35
+ return False
36
+ return False
37
+
38
+
39
+ def _normalize_separators(run: str) -> str:
40
+ """Digits with "," and "." as thousands or decimal separators -> a float literal.
41
+
42
+ When both appear, the last one is the decimal separator ("1,200.50", "1.200,50").
43
+ A lone comma is a decimal comma unless exactly three digits follow it ("12,50" vs "1,200").
44
+ Several dots are thousands separators ("1.200.000").
45
+ """
46
+ comma, dot = run.rfind(","), run.rfind(".")
47
+ if comma >= 0 and dot >= 0:
48
+ if comma > dot:
49
+ return run.replace(".", "").replace(",", ".")
50
+ return run.replace(",", "")
51
+ if comma >= 0:
52
+ if run.count(",") == 1 and len(run) - comma - 1 != 3:
53
+ return run.replace(",", ".")
54
+ return run.replace(",", "")
55
+ if run.count(".") > 1:
56
+ return run.replace(".", "")
57
+ return run
58
+
59
+
60
+ def parse_number(text: str) -> Optional[float]:
61
+ """The first number in the text, or None if it has none."""
62
+ start = next((i for i, ch in enumerate(text) if ch in _DIGITS), None)
63
+ if start is None:
64
+ return None
65
+ end = start
66
+ while end < len(text) and (text[end] in _DIGITS or text[end] in ".,"):
67
+ end += 1
68
+ run = text[start:end].rstrip(".,")
69
+ try:
70
+ value = float(_normalize_separators(run))
71
+ except ValueError:
72
+ return None
73
+ return -value if _is_negative(text[:start]) else value
74
+
75
+
76
+ def coerce_json_value(raw: object, kind: str) -> tuple[TValue, Optional[str]]:
77
+ """(value, warning) for a JSON value of a string, number, or integer field.
78
+
79
+ Numbers pass through directly; anything else is read as text with ``parse_field_text``.
80
+ """
81
+ if kind in ("number", "integer") and isinstance(raw, (int, float)) and not isinstance(raw, bool):
82
+ if kind == "number":
83
+ return float(raw), None
84
+ return (int(raw), None) if float(raw).is_integer() else (None, f"{raw!r} is not a whole number; left empty")
85
+ text = raw if isinstance(raw, str) else json.dumps(raw, ensure_ascii=False)
86
+ return parse_field_text(text, kind)
87
+
88
+
89
+ def parse_field_text(text: str, kind: str) -> tuple[TValue, Optional[str]]:
90
+ """(value, warning) for a string, number, or integer field.
91
+
92
+ "NONE" and similar markers mean the input doesn't contain the value: (None, None).
93
+ An answer that can't be read as the type gives (None, a warning).
94
+ """
95
+ if kind not in TEXT_TYPES:
96
+ raise ValueError(f"Only {', '.join(TEXT_TYPES)} fields are parsed from text, not {kind!r}.")
97
+ s = clean_text(text)
98
+ if s.lower() in NONE_MARKERS:
99
+ return None, None
100
+ if kind == "string":
101
+ return s, None
102
+ number = parse_number(s)
103
+ if number is None:
104
+ return None, f"could not read a number from {s[:40]!r}; left empty"
105
+ if kind == "integer":
106
+ if not number.is_integer():
107
+ return None, f"{s[:40]!r} is not a whole number; left empty"
108
+ return int(number), None
109
+ return number, None
@@ -0,0 +1 @@
1
+ """Provider layer: technology-specific implementations of core interfaces. Providers never import each other."""
@@ -0,0 +1 @@
1
+
@@ -0,0 +1,5 @@
1
+ from .clock_mock_impl import ClockMockImpl
2
+ from .engine_mock_impl import EngineMockImpl
3
+ from .id_provider_mock_impl import IdProviderMockImpl
4
+
5
+ __all__ = ["ClockMockImpl", "EngineMockImpl", "IdProviderMockImpl"]
@@ -0,0 +1,19 @@
1
+ """A deterministic clock: starts at ``start`` and advances ``step`` seconds per read."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dev_double.core.decision.i_clock import IClock
6
+
7
+
8
+ class ClockMockImpl(IClock):
9
+ def __init__(self, start: float = 0.0, step: float = 0.0) -> None:
10
+ self._t = start
11
+ self._step = step
12
+
13
+ def now(self) -> float:
14
+ t = self._t
15
+ self._t += self._step
16
+ return t
17
+
18
+ def advance(self, seconds: float) -> None:
19
+ self._t += seconds
@@ -0,0 +1,78 @@
1
+ """A model-free engine for CI and plumbing tests.
2
+
3
+ It scores each label by word overlap between the input and the label's
4
+ description. It is deterministic and instant, and its answers are only
5
+ loosely sensible. Use it to test your integration code, never its quality.
6
+
7
+ It also "generates" extraction values: for each field, the text after the
8
+ colon on the first input line whose label (before the colon) contains every
9
+ word of the field name, so ``due_date`` finds ``Due date: 2026-10-01``. In JSON
10
+ mode it returns one object with a key per field (null when not found) and
11
+ confidence 1.0; otherwise the first field's value, or NONE with confidence 0.5.
12
+ It returns no token logprobs.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import json
18
+ import math
19
+ import re
20
+ from typing import Optional
21
+
22
+ from dev_double.core.decision.i_engine import IEngine
23
+ from dev_double.core.decision.i_generator import IGenerator
24
+ from dev_double.core.decision.t_generate import TGenerateQuery, TGenerateResult
25
+ from dev_double.core.decision.t_label_query import TLabelQuery, TLabelResult
26
+ from dev_double.core.decision.t_usage import TUsage
27
+
28
+ _WORD = re.compile(r"[a-z0-9]+")
29
+ _STOP = frozenset(
30
+ "a an and are as at be by for from has have in is it its of on or that the this to was "
31
+ "were will with you your does do not no yes which what how".split()
32
+ )
33
+
34
+
35
+ def _words(text: str) -> set[str]:
36
+ return {w for w in _WORD.findall(text.lower()) if w not in _STOP and len(w) > 2}
37
+
38
+
39
+ class EngineMockImpl(IEngine, IGenerator):
40
+ name = "mock"
41
+ model = "mock-overlap"
42
+
43
+ async def aclose(self) -> None:
44
+ return None
45
+
46
+ async def distribution(self, query: TLabelQuery) -> TLabelResult:
47
+ input_words = _words(query.input_text)
48
+ scores = []
49
+ for description in query.descriptions:
50
+ desc = _words(description)
51
+ overlap = len(input_words & desc)
52
+ scores.append(overlap / (1 + math.sqrt(len(desc))))
53
+ top = max(scores)
54
+ exps = [math.exp(4 * (s - top)) for s in scores]
55
+ total = sum(exps)
56
+ probs = {label: e / total for label, e in zip(query.labels, exps)}
57
+ usage = TUsage(
58
+ input_tokens=len(query.system.split()) + len(query.user.split()),
59
+ output_tokens=1,
60
+ )
61
+ return TLabelResult(probs, usage)
62
+
63
+ @staticmethod
64
+ def _lookup(field: str, input_text: str) -> Optional[str]:
65
+ wanted = set(_WORD.findall(field.lower()))
66
+ for line in input_text.splitlines():
67
+ key, colon, rest = line.partition(":")
68
+ if colon and wanted and wanted <= set(_WORD.findall(key.lower())):
69
+ return rest.strip().rstrip(",").strip().strip('"')
70
+ return None
71
+
72
+ async def generate(self, query: TGenerateQuery) -> TGenerateResult:
73
+ found = {name: self._lookup(name, query.input_text) for name in query.fields}
74
+ usage = TUsage(input_tokens=len(query.system.split()) + len(query.user.split()), output_tokens=1)
75
+ if query.json:
76
+ return TGenerateResult(json.dumps(found), 1.0, usage)
77
+ first = next(iter(found.values()), None)
78
+ return TGenerateResult(first, 1.0, usage) if first is not None else TGenerateResult("NONE", 0.5, usage)
@@ -0,0 +1,15 @@
1
+ """Deterministic ids: ``<prefix>-1``, ``<prefix>-2``, ..."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dev_double.core.decision.i_id_provider import IIdProvider
6
+
7
+
8
+ class IdProviderMockImpl(IIdProvider):
9
+ def __init__(self, prefix: str = "mock") -> None:
10
+ self._prefix = prefix
11
+ self._n = 0
12
+
13
+ def new_id(self) -> str:
14
+ self._n += 1
15
+ return f"{self._prefix}-{self._n}"
@@ -0,0 +1 @@
1
+
@@ -0,0 +1,4 @@
1
+ # needle itself is imported lazily when DeciderNeedleImpl is constructed.
2
+ from .decider_needle_impl import DeciderNeedleImpl
3
+
4
+ __all__ = ["DeciderNeedleImpl"]
@@ -0,0 +1,117 @@
1
+ """Cactus Needle, a tiny on-device tool-calling model, as a decider.
2
+
3
+ Needle isn't a decision model: it fills tool calls. Its docs treat
4
+ classification as extraction with an enum field, so each question becomes one
5
+ tool whose `label` argument may only be one of the allowed labels. Needle
6
+ returns one calibrated confidence for the call, not a distribution, so the
7
+ other labels share the remainder evenly. Rerank uses Needle's embeddings.
8
+ Extraction is native: one ``record`` tool call carries every field (see record_tool).
9
+
10
+ Needs the optional extra: `pip install "dev-double[needle]"`. ``needle`` is
11
+ imported lazily, so the package imports without it. Telemetry is off by
12
+ default (NEEDLE_TELEMETRY=0, DO_NOT_TRACK=1); the library reads
13
+ NEEDLE3_LIB_PATH itself if you need to point it at a local build.
14
+ """
15
+
16
+ from __future__ import annotations
17
+
18
+ import math
19
+ import os
20
+ from typing import Any, Literal, Optional
21
+
22
+ from dev_double.core.decision import prompts
23
+ from dev_double.core.decision.answer_shaping import question_labels, shape_answer
24
+ from dev_double.core.decision.confidence import DIGITS
25
+ from dev_double.core.decision.i_decider import IDecider
26
+ from dev_double.core.decision.i_extractor import IExtractor
27
+ from dev_double.core.decision.i_reranker import IReranker
28
+ from dev_double.core.decision.t_answer import TAnswer
29
+ from dev_double.core.decision.t_extract import EMPTY, TExtractRequest, TFieldValue
30
+ from dev_double.core.decision.t_input import TInputValue
31
+ from dev_double.core.decision.t_question import TQuestion
32
+ from dev_double.core.decision.tracker import Tracker
33
+
34
+ from .record_tool import RECORD_SYSTEM, fields_key, record_model, value_from_argument
35
+
36
+
37
+ def _import_needle() -> Any:
38
+ os.environ.setdefault("NEEDLE_TELEMETRY", "0")
39
+ os.environ.setdefault("DO_NOT_TRACK", "1")
40
+ import needle
41
+
42
+ return needle
43
+
44
+
45
+ def _cosine(a: list[float], b: list[float]) -> float:
46
+ dot = sum(x * y for x, y in zip(a, b))
47
+ return dot / (math.sqrt(sum(x * x for x in a)) * math.sqrt(sum(y * y for y in b)) or 1)
48
+
49
+
50
+ class DeciderNeedleImpl(IDecider, IReranker, IExtractor):
51
+ name = "needle"
52
+ model = "needle3"
53
+
54
+ def __init__(self) -> None:
55
+ self._needle = _import_needle()
56
+ self._agents: dict[tuple, Any] = {}
57
+ self._extractors: dict[tuple, Any] = {}
58
+ self._embedder: Optional[Any] = None
59
+
60
+ def _agent(self, q: TQuestion, labels: dict[str, str]) -> Any:
61
+ from pydantic import Field, create_model # the tool schema Needle expects
62
+
63
+ key = (q.question, tuple(labels.items()))
64
+ if key not in self._agents:
65
+ options = "; ".join(f"{k}: {v}" for k, v in labels.items())
66
+ tool = create_model(
67
+ "answer",
68
+ __doc__=f"Record the answer to: {q.question}",
69
+ label=(Literal[tuple(labels)], Field(description=f"Exactly one of: {options}")),
70
+ )
71
+ system = f"{q.question}\nAllowed labels: {options}\nAlways call answer with one label."
72
+ self._agents[key] = self._needle.Needle(tools=[tool], system=system, stateless=True)
73
+ return self._agents[key]
74
+
75
+ async def ask(
76
+ self, value: TInputValue, question: TQuestion, tracker: Tracker, label: str = ""
77
+ ) -> TAnswer:
78
+ labels = question_labels(question)
79
+ r = self._agent(question, labels).complete(prompts.render_input(value))
80
+ calls = r.get("function_calls") or r.get("suppressed_calls") or []
81
+ chosen = (calls[0].get("arguments") or {}).get("label") if calls else None
82
+ conf = r.get("confidence")
83
+ conf = 1.0 if conf is None else float(conf)
84
+ if chosen in labels:
85
+ rest = (1 - conf) / (len(labels) - 1)
86
+ probs = {k: (conf if k == chosen else rest) for k in labels}
87
+ else:
88
+ tracker.warn("needle made no valid call", label)
89
+ probs = {k: 1 / len(labels) for k in labels}
90
+ return shape_answer(question, probs)
91
+
92
+ async def extract(self, req: TExtractRequest, tracker: Tracker) -> dict[str, TFieldValue]:
93
+ key = fields_key(req.fields)
94
+ if key not in self._extractors:
95
+ tool = record_model(req.fields)
96
+ self._extractors[key] = self._needle.Needle(tools=[tool], system=RECORD_SYSTEM, stateless=True)
97
+ r = self._extractors[key].complete(prompts.render_input(req.input))
98
+ calls = r.get("function_calls") or r.get("suppressed_calls") or []
99
+ if not calls:
100
+ tracker.warn("needle made no valid call")
101
+ return {name: EMPTY for name in req.fields}
102
+ args = calls[0].get("arguments") or {}
103
+ conf = r.get("confidence")
104
+ conf = 1.0 if conf is None else min(1.0, max(0.0, float(conf)))
105
+ return {name: value_from_argument(name, f, args.get(name), conf, tracker) for name, f in req.fields.items()}
106
+
107
+ async def rerank_scores(self, query: str, documents: list[str], tracker: Tracker) -> list[float]:
108
+ if self._embedder is None:
109
+ self._embedder = self._needle.Needle(stateless=True)
110
+ embed = self._embedder.embed
111
+ q = embed(query)
112
+ return [round(_cosine(q, embed(text)), DIGITS) for text in documents]
113
+
114
+ async def aclose(self) -> None:
115
+ agents = [*self._agents.values(), *self._extractors.values()]
116
+ for agent in [*agents, *([self._embedder] if self._embedder else [])]:
117
+ agent.close()