-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathtest_adapter_phase.py
More file actions
167 lines (136 loc) · 5.09 KB
/
Copy pathtest_adapter_phase.py
File metadata and controls
167 lines (136 loc) · 5.09 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
"""Tests for ``AdapterPhase``.
Uses an in-process fake ``PlatformClient`` to avoid real HTTP. Verifies:
- happy path: source has N adapters, target gets N POSTs, all remapped
- idempotency: re-run with target already populated → zero POSTs, all adopted
- dry-run: zero POSTs, all skipped
- on_name_conflict='abort' raises on existing
"""
from __future__ import annotations
import pytest
from unstract.clone.context import (
CloneContext,
CloneOptions,
RemapTable,
)
from unstract.clone.exceptions import NameConflictError
from unstract.clone.phases.adapter import AdapterPhase
from unstract.clone.report import CloneReport
class FakeClient:
"""Minimal in-memory stand-in for ``PlatformClient``."""
# Mirrors DRF OPTIONS actions.POST writable fields for adapter.
POST_SCHEMA = frozenset(
{"adapter_id", "adapter_name", "adapter_type", "adapter_metadata", "description"}
)
def __init__(self, adapters: list[dict] | None = None):
# Stored as a list of dicts; mutated by create_adapter.
self.adapters: list[dict] = list(adapters or [])
self.posts: list[dict] = []
self._next_id = 1
def get_post_schema(self, entity_path):
return self.POST_SCHEMA
def list_adapters(self, *, name=None, adapter_type=None):
result = self.adapters
if name is not None:
result = [a for a in result if a["adapter_name"] == name]
if adapter_type is not None:
result = [a for a in result if a["adapter_type"] == adapter_type]
# Mimic AdapterListSerializer — strip adapter_metadata from list output.
return [{k: v for k, v in a.items() if k != "adapter_metadata"} for a in result]
def get_adapter(self, adapter_pk):
for a in self.adapters:
if a["id"] == adapter_pk:
return a
raise KeyError(adapter_pk)
def create_adapter(self, payload):
new = dict(payload)
new["id"] = f"tgt-{self._next_id:08d}-0000-0000-0000-000000000000"
self._next_id += 1
self.adapters.append(new)
self.posts.append(new)
return new
def _src_adapter(id_, name, atype="LLM"):
return {
"id": id_,
"adapter_id": "openai-llm-v2",
"adapter_name": name,
"adapter_type": atype,
"adapter_metadata": {"api_key": "sk-secret", "model": "gpt-4"},
"description": f"{name} desc",
}
def _ctx(source: FakeClient, target: FakeClient, **opt_overrides):
ctx = CloneContext(
source=source,
target=target,
options=CloneOptions(**opt_overrides),
remap=RemapTable(),
)
return ctx
def test_happy_path_creates_all_and_records_remap():
src = FakeClient(
[
_src_adapter("src-a", "OpenAI Prod"),
_src_adapter("src-b", "Mistral Stg", atype="EMBEDDING"),
]
)
tgt = FakeClient()
ctx = _ctx(src, tgt)
report = CloneReport()
result = AdapterPhase(ctx).run(report)
assert result.created == 2
assert result.adopted == 0
assert result.failed == 0
assert len(tgt.posts) == 2
assert ctx.remap.resolve("adapter", "src-a") == tgt.posts[0]["id"]
assert ctx.remap.resolve("adapter", "src-b") == tgt.posts[1]["id"]
def test_idempotency_zero_creates_on_rerun():
src_adapters = [_src_adapter("src-a", "OpenAI Prod")]
src = FakeClient(src_adapters)
# Target pre-populated with the same name+type — simulates a prior run.
tgt = FakeClient(
[
{
"id": "preexisting",
"adapter_id": "openai-llm-v2",
"adapter_name": "OpenAI Prod",
"adapter_type": "LLM",
"adapter_metadata": {},
}
]
)
ctx = _ctx(src, tgt, on_name_conflict="adopt")
report = CloneReport()
result = AdapterPhase(ctx).run(report)
assert result.created == 0
assert result.adopted == 1
assert tgt.posts == [] # no new POSTs
assert ctx.remap.resolve("adapter", "src-a") == "preexisting"
def test_dry_run_makes_no_posts():
src = FakeClient([_src_adapter("src-a", "OpenAI Prod")])
tgt = FakeClient()
ctx = _ctx(src, tgt, dry_run=True)
report = CloneReport()
result = AdapterPhase(ctx).run(report)
# Dry-run predicts the create (count matches a real run) but writes nothing
# and records a synthetic remap so dependent phases can plan.
assert result.created == 1
assert result.skipped == 0
assert tgt.posts == []
planned = ctx.remap.resolve("adapter", "src-a")
assert planned is not None and ctx.remap.is_planned(planned)
def test_abort_on_name_conflict_raises():
src = FakeClient([_src_adapter("src-a", "OpenAI Prod")])
tgt = FakeClient(
[
{
"id": "preexisting",
"adapter_id": "openai-llm-v2",
"adapter_name": "OpenAI Prod",
"adapter_type": "LLM",
"adapter_metadata": {},
}
]
)
ctx = _ctx(src, tgt, on_name_conflict="abort")
report = CloneReport()
with pytest.raises(NameConflictError):
AdapterPhase(ctx).run(report)