schemagate 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.
- schemagate/__init__.py +29 -0
- schemagate/ai/__init__.py +28 -0
- schemagate/ai/describe.py +170 -0
- schemagate/ai/embedder.py +128 -0
- schemagate/ai/providers.py +268 -0
- schemagate/catalog.py +370 -0
- schemagate/cli.py +161 -0
- schemagate/demo_schema.py +260 -0
- schemagate/embedder.py +150 -0
- schemagate/embedders/__init__.py +0 -0
- schemagate/embedders/hf.py +30 -0
- schemagate/identity.py +80 -0
- schemagate/integrations/__init__.py +0 -0
- schemagate/integrations/langchain.py +79 -0
- schemagate/introspect.py +100 -0
- schemagate/mcp_server.py +375 -0
- schemagate/models.py +172 -0
- schemagate/py.typed +0 -0
- schemagate/store.py +20 -0
- schemagate/stores/__init__.py +0 -0
- schemagate/stores/memory.py +44 -0
- schemagate/stores/oracle.py +228 -0
- schemagate/studio.html +1174 -0
- schemagate/studio.py +194 -0
- schemagate-0.1.0.dist-info/METADATA +487 -0
- schemagate-0.1.0.dist-info/RECORD +29 -0
- schemagate-0.1.0.dist-info/WHEEL +4 -0
- schemagate-0.1.0.dist-info/entry_points.txt +2 -0
- schemagate-0.1.0.dist-info/licenses/LICENSE +17 -0
schemagate/__init__.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
"""schemagate -- identity-scoped schema selection for NL2SQL.
|
|
2
|
+
|
|
3
|
+
from schemagate import Catalog, Principal
|
|
4
|
+
cat = Catalog().bootstrap("postgresql://localhost/app")
|
|
5
|
+
sel = cat.select("revenue by month", principal=Principal("okta:jdoe"))
|
|
6
|
+
sel.prompt_fragment()
|
|
7
|
+
"""
|
|
8
|
+
from .identity import Principal, IdentityError
|
|
9
|
+
from .models import Column, ForeignKey, ObjectDoc, Selection, Scored
|
|
10
|
+
from .embedder import HashingEmbedder, cosine_distance, tokenize
|
|
11
|
+
from .stores.memory import MemoryStore
|
|
12
|
+
from .catalog import Catalog
|
|
13
|
+
|
|
14
|
+
__version__ = "0.1.0"
|
|
15
|
+
__all__ = ["Catalog", "Principal", "IdentityError", "ObjectDoc", "Column",
|
|
16
|
+
"ForeignKey", "Selection", "Scored", "HashingEmbedder",
|
|
17
|
+
"MemoryStore", "cosine_distance", "tokenize"]
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def __getattr__(name):
|
|
21
|
+
# Lazy, so importing schemagate never pulls in oracledb or
|
|
22
|
+
# sentence-transformers. Both raise a clear ImportError naming the extra.
|
|
23
|
+
if name == "OracleStore":
|
|
24
|
+
from .stores.oracle import OracleStore
|
|
25
|
+
return OracleStore
|
|
26
|
+
if name == "SentenceTransformerEmbedder":
|
|
27
|
+
from .embedders.hf import SentenceTransformerEmbedder
|
|
28
|
+
return SentenceTransformerEmbedder
|
|
29
|
+
raise AttributeError(f"module 'schemagate' has no attribute {name!r}")
|
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
"""Optional AI features. schemagate works fully without importing this package.
|
|
2
|
+
|
|
3
|
+
Nothing here is required. The default catalog uses an offline, deterministic
|
|
4
|
+
embedder and no network. Import from ``schemagate.ai`` only when you want a model
|
|
5
|
+
to write catalog descriptions or produce embeddings, and bring your own key.
|
|
6
|
+
|
|
7
|
+
from schemagate.ai import SchemaDescriber, AnthropicProvider
|
|
8
|
+
cat.describe(SchemaDescriber(AnthropicProvider(model="claude-sonnet-4-5")))
|
|
9
|
+
"""
|
|
10
|
+
from .describe import SchemaDescriber
|
|
11
|
+
from .embedder import APIEmbedder
|
|
12
|
+
from .providers import (
|
|
13
|
+
AnthropicProvider,
|
|
14
|
+
CallableProvider,
|
|
15
|
+
GeminiProvider,
|
|
16
|
+
OpenAIProvider,
|
|
17
|
+
Provider,
|
|
18
|
+
ProviderError,
|
|
19
|
+
auto_provider,
|
|
20
|
+
available_providers,
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
__all__ = [
|
|
24
|
+
"SchemaDescriber", "APIEmbedder",
|
|
25
|
+
"Provider", "ProviderError", "CallableProvider",
|
|
26
|
+
"AnthropicProvider", "OpenAIProvider", "GeminiProvider",
|
|
27
|
+
"auto_provider", "available_providers",
|
|
28
|
+
]
|
|
@@ -0,0 +1,170 @@
|
|
|
1
|
+
"""AI-generated catalog descriptions.
|
|
2
|
+
|
|
3
|
+
Retrieval quality is limited by how much meaning the schema text carries.
|
|
4
|
+
``CUST_ORD_LN_T`` with columns ``ID``, ``QTY``, ``AMT`` tells a retriever
|
|
5
|
+
almost nothing, which is why the offline embedder scores 50% on questions
|
|
6
|
+
phrased in business words rather than identifier words.
|
|
7
|
+
|
|
8
|
+
A ``SchemaDescriber`` asks a model to write one sentence per object saying
|
|
9
|
+
what it holds and when to use it, and stores that on
|
|
10
|
+
``ObjectDoc.description``. That text is already part of ``embed_text()``,
|
|
11
|
+
so descriptions improve both vector and BM25 retrieval with no other change.
|
|
12
|
+
|
|
13
|
+
**Only schema metadata is sent.** Object names, column names, types,
|
|
14
|
+
nullability, existing comments, and foreign keys. No rows, no sample values,
|
|
15
|
+
no query results, no credentials -- ``ObjectDoc`` does not carry row data,
|
|
16
|
+
and ``_render`` cannot reach any. There is a test asserting this.
|
|
17
|
+
|
|
18
|
+
**A human hint always wins.** ``Catalog.hint()`` is applied after
|
|
19
|
+
descriptions and outranks them everywhere, so a wrong AI description can be
|
|
20
|
+
corrected without regenerating anything.
|
|
21
|
+
|
|
22
|
+
Descriptions cost money, so results are cached by content: an object is
|
|
23
|
+
re-described only when its structure or the model changes.
|
|
24
|
+
"""
|
|
25
|
+
from __future__ import annotations
|
|
26
|
+
|
|
27
|
+
import hashlib
|
|
28
|
+
import json
|
|
29
|
+
import os
|
|
30
|
+
import pathlib
|
|
31
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
32
|
+
from typing import Dict, Iterable, List, Optional, Sequence
|
|
33
|
+
|
|
34
|
+
from ..models import ObjectDoc
|
|
35
|
+
from .providers import Provider, ProviderError
|
|
36
|
+
|
|
37
|
+
_SYSTEM = (
|
|
38
|
+
"You document database schemas for a SQL-generating assistant. "
|
|
39
|
+
"Given one table or view definition, reply with a single sentence, at "
|
|
40
|
+
"most 25 words, saying what the object holds and what question it "
|
|
41
|
+
"answers. Use the business meaning, not the column list. Do not repeat "
|
|
42
|
+
"the object name. Do not speculate about data you cannot see. No "
|
|
43
|
+
"preamble, no markdown, no quotes -- just the sentence."
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def _render(doc: ObjectDoc, max_columns: int = 30) -> str:
|
|
48
|
+
"""The only thing ever sent to a provider. Metadata, never rows."""
|
|
49
|
+
lines = [f"{doc.kind} {doc.qname}"]
|
|
50
|
+
if doc.description:
|
|
51
|
+
lines.append(f"existing comment: {doc.description}")
|
|
52
|
+
for col in doc.columns[:max_columns]:
|
|
53
|
+
bits = f" {col.name} {col.type}"
|
|
54
|
+
if col.pk:
|
|
55
|
+
bits += " PK"
|
|
56
|
+
if col.comment:
|
|
57
|
+
bits += f" -- {col.comment}"
|
|
58
|
+
lines.append(bits)
|
|
59
|
+
if len(doc.columns) > max_columns:
|
|
60
|
+
lines.append(f" ...{len(doc.columns) - max_columns} more columns")
|
|
61
|
+
for fk in doc.foreign_keys:
|
|
62
|
+
lines.append(f" FK {','.join(fk.columns)} -> {fk.ref_table}")
|
|
63
|
+
return "\n".join(lines)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def _fingerprint(doc: ObjectDoc, model: str) -> str:
|
|
67
|
+
"""Cache key: changes when the object's structure or the model changes."""
|
|
68
|
+
payload = json.dumps({
|
|
69
|
+
"q": doc.qname, "k": doc.kind, "m": model,
|
|
70
|
+
"c": [[c.name, c.type, c.pk, c.comment] for c in doc.columns],
|
|
71
|
+
"f": [[fk.columns, fk.ref_table] for fk in doc.foreign_keys],
|
|
72
|
+
"d": doc.description,
|
|
73
|
+
}, sort_keys=True, default=str)
|
|
74
|
+
return hashlib.sha256(payload.encode("utf-8")).hexdigest()[:32]
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
class SchemaDescriber:
|
|
78
|
+
"""Generate one-sentence descriptions for catalog objects.
|
|
79
|
+
|
|
80
|
+
Parameters
|
|
81
|
+
----------
|
|
82
|
+
provider
|
|
83
|
+
Any object with ``complete(system, prompt, max_tokens)``.
|
|
84
|
+
cache_path
|
|
85
|
+
JSON file for generated descriptions. Highly recommended: it makes
|
|
86
|
+
re-runs free and keeps a diffable record of what the model wrote.
|
|
87
|
+
workers
|
|
88
|
+
Parallel requests. Keep it modest; providers rate-limit.
|
|
89
|
+
strict
|
|
90
|
+
``False`` (default) skips objects whose call failed and carries on,
|
|
91
|
+
because a partial catalog still selects. ``True`` re-raises.
|
|
92
|
+
"""
|
|
93
|
+
|
|
94
|
+
def __init__(self, provider: Provider, cache_path: Optional[str] = None,
|
|
95
|
+
max_columns: int = 30, workers: int = 4,
|
|
96
|
+
max_tokens: int = 120, strict: bool = False):
|
|
97
|
+
self.provider = provider
|
|
98
|
+
self.cache_path = pathlib.Path(cache_path) if cache_path else None
|
|
99
|
+
self.max_columns = max_columns
|
|
100
|
+
self.workers = max(1, int(workers))
|
|
101
|
+
self.max_tokens = max_tokens
|
|
102
|
+
self.strict = strict
|
|
103
|
+
self.failures: List[str] = []
|
|
104
|
+
self._cache: Dict[str, str] = self._load_cache()
|
|
105
|
+
|
|
106
|
+
# ---------------- cache ----------------
|
|
107
|
+
|
|
108
|
+
def _load_cache(self) -> Dict[str, str]:
|
|
109
|
+
if not self.cache_path or not self.cache_path.exists():
|
|
110
|
+
return {}
|
|
111
|
+
try:
|
|
112
|
+
return json.loads(self.cache_path.read_text("utf-8"))
|
|
113
|
+
except (json.JSONDecodeError, OSError):
|
|
114
|
+
return {} # a corrupt cache must never break cataloguing
|
|
115
|
+
|
|
116
|
+
def _save_cache(self) -> None:
|
|
117
|
+
if not self.cache_path:
|
|
118
|
+
return
|
|
119
|
+
try:
|
|
120
|
+
self.cache_path.parent.mkdir(parents=True, exist_ok=True)
|
|
121
|
+
tmp = self.cache_path.with_suffix(self.cache_path.suffix + ".tmp")
|
|
122
|
+
tmp.write_text(json.dumps(self._cache, indent=2, sort_keys=True),
|
|
123
|
+
encoding="utf-8")
|
|
124
|
+
os.replace(tmp, self.cache_path)
|
|
125
|
+
except OSError:
|
|
126
|
+
pass # a read-only disk is not a reason to lose the run
|
|
127
|
+
|
|
128
|
+
# ---------------- generation ----------------
|
|
129
|
+
|
|
130
|
+
def _describe_one(self, doc: ObjectDoc) -> Optional[str]:
|
|
131
|
+
key = _fingerprint(doc, getattr(self.provider, "name", "?"))
|
|
132
|
+
if key in self._cache:
|
|
133
|
+
return self._cache[key]
|
|
134
|
+
try:
|
|
135
|
+
text = self.provider.complete(_SYSTEM, _render(doc, self.max_columns),
|
|
136
|
+
max_tokens=self.max_tokens)
|
|
137
|
+
except ProviderError:
|
|
138
|
+
self.failures.append(doc.qname)
|
|
139
|
+
if self.strict:
|
|
140
|
+
raise
|
|
141
|
+
return None
|
|
142
|
+
text = " ".join((text or "").split()).strip().strip('"')
|
|
143
|
+
if not text:
|
|
144
|
+
self.failures.append(doc.qname)
|
|
145
|
+
return None
|
|
146
|
+
self._cache[key] = text
|
|
147
|
+
return text
|
|
148
|
+
|
|
149
|
+
def describe(self, docs: Iterable[ObjectDoc]) -> Dict[str, str]:
|
|
150
|
+
"""Return ``{qname: description}``. Objects that failed are absent."""
|
|
151
|
+
docs = list(docs)
|
|
152
|
+
self.failures = []
|
|
153
|
+
if not docs:
|
|
154
|
+
return {}
|
|
155
|
+
if self.workers == 1:
|
|
156
|
+
results = [self._describe_one(d) for d in docs]
|
|
157
|
+
else:
|
|
158
|
+
with ThreadPoolExecutor(max_workers=self.workers) as pool:
|
|
159
|
+
results = list(pool.map(self._describe_one, docs))
|
|
160
|
+
self._save_cache()
|
|
161
|
+
return {d.qname: t for d, t in zip(docs, results) if t}
|
|
162
|
+
|
|
163
|
+
# what the provider would receive, for review before spending money
|
|
164
|
+
def preview(self, doc: ObjectDoc) -> str:
|
|
165
|
+
return _render(doc, self.max_columns)
|
|
166
|
+
|
|
167
|
+
def estimate_calls(self, docs: Sequence[ObjectDoc]) -> int:
|
|
168
|
+
"""How many billed calls ``describe()`` would make right now."""
|
|
169
|
+
name = getattr(self.provider, "name", "?")
|
|
170
|
+
return sum(1 for d in docs if _fingerprint(d, name) not in self._cache)
|
|
@@ -0,0 +1,128 @@
|
|
|
1
|
+
"""Embeddings from a hosted model.
|
|
2
|
+
|
|
3
|
+
The offline ``HashingEmbedder`` matches subwords, which is strong on
|
|
4
|
+
identifier-shaped questions and weak on pure paraphrase. An API embedder
|
|
5
|
+
trades money, latency and a network dependency for semantic matching.
|
|
6
|
+
|
|
7
|
+
Measure before adopting: on identifier-heavy schema text the offline
|
|
8
|
+
embedder is often competitive, and ``tests/bench.py`` runs against your
|
|
9
|
+
own schema. This is a swap, not an upgrade.
|
|
10
|
+
|
|
11
|
+
from schemagate import Catalog
|
|
12
|
+
from schemagate.ai import APIEmbedder, OpenAIProvider
|
|
13
|
+
|
|
14
|
+
provider = OpenAIProvider(model="gpt-4.1-mini",
|
|
15
|
+
embed_model="text-embedding-3-small")
|
|
16
|
+
cat = Catalog(embedder=APIEmbedder(provider, dim=1536))
|
|
17
|
+
|
|
18
|
+
Vectors are cached in memory for the life of the embedder and, if you pass
|
|
19
|
+
``cache_path``, on disk -- so re-indexing an unchanged schema is free.
|
|
20
|
+
"""
|
|
21
|
+
from __future__ import annotations
|
|
22
|
+
|
|
23
|
+
import hashlib
|
|
24
|
+
import json
|
|
25
|
+
import math
|
|
26
|
+
import os
|
|
27
|
+
import pathlib
|
|
28
|
+
from typing import Dict, List, Optional, Sequence
|
|
29
|
+
|
|
30
|
+
from .providers import Provider, ProviderError
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def _l2(vec: List[float]) -> List[float]:
|
|
34
|
+
n = math.sqrt(sum(v * v for v in vec))
|
|
35
|
+
return [v / n for v in vec] if n else vec
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class APIEmbedder:
|
|
39
|
+
"""Embedder backed by a provider's embedding endpoint.
|
|
40
|
+
|
|
41
|
+
Parameters
|
|
42
|
+
----------
|
|
43
|
+
provider
|
|
44
|
+
Any object with ``embed(texts) -> list[list[float]]``.
|
|
45
|
+
dim
|
|
46
|
+
Expected dimension. Checked on the first response, because a silent
|
|
47
|
+
dimension change invalidates every persisted vector.
|
|
48
|
+
batch_size
|
|
49
|
+
Texts per request.
|
|
50
|
+
cache_path
|
|
51
|
+
Optional JSON cache keyed by SHA-256 of the text plus provider name.
|
|
52
|
+
"""
|
|
53
|
+
|
|
54
|
+
def __init__(self, provider: Provider, dim: int, batch_size: int = 64,
|
|
55
|
+
cache_path: Optional[str] = None, normalize: bool = True):
|
|
56
|
+
if not hasattr(provider, "embed"):
|
|
57
|
+
raise TypeError(
|
|
58
|
+
f"{getattr(provider, 'name', provider)!r} has no embed(); "
|
|
59
|
+
"pass a provider built with an embed model")
|
|
60
|
+
self.provider = provider
|
|
61
|
+
self.dim = int(dim)
|
|
62
|
+
self.batch_size = max(1, int(batch_size))
|
|
63
|
+
self.normalize = normalize
|
|
64
|
+
self.name = f"api:{getattr(provider, 'name', 'provider')}"
|
|
65
|
+
self.cache_path = pathlib.Path(cache_path) if cache_path else None
|
|
66
|
+
self._cache: Dict[str, List[float]] = self._load()
|
|
67
|
+
|
|
68
|
+
# ---------------- cache ----------------
|
|
69
|
+
|
|
70
|
+
def _key(self, text: str) -> str:
|
|
71
|
+
seed = f"{self.name}:{self.dim}:{text}"
|
|
72
|
+
return hashlib.sha256(seed.encode("utf-8")).hexdigest()[:32]
|
|
73
|
+
|
|
74
|
+
def _load(self) -> Dict[str, List[float]]:
|
|
75
|
+
if not self.cache_path or not self.cache_path.exists():
|
|
76
|
+
return {}
|
|
77
|
+
try:
|
|
78
|
+
return json.loads(self.cache_path.read_text("utf-8"))
|
|
79
|
+
except (json.JSONDecodeError, OSError):
|
|
80
|
+
return {}
|
|
81
|
+
|
|
82
|
+
def _save(self) -> None:
|
|
83
|
+
if not self.cache_path:
|
|
84
|
+
return
|
|
85
|
+
try:
|
|
86
|
+
self.cache_path.parent.mkdir(parents=True, exist_ok=True)
|
|
87
|
+
tmp = self.cache_path.with_suffix(self.cache_path.suffix + ".tmp")
|
|
88
|
+
tmp.write_text(json.dumps(self._cache), encoding="utf-8")
|
|
89
|
+
os.replace(tmp, self.cache_path)
|
|
90
|
+
except OSError:
|
|
91
|
+
pass
|
|
92
|
+
|
|
93
|
+
# ---------------- embedding ----------------
|
|
94
|
+
|
|
95
|
+
def embed(self, texts: Sequence[str]) -> List[List[float]]:
|
|
96
|
+
texts = list(texts)
|
|
97
|
+
out: List[Optional[List[float]]] = [None] * len(texts)
|
|
98
|
+
todo: List[int] = []
|
|
99
|
+
for i, text in enumerate(texts):
|
|
100
|
+
hit = self._cache.get(self._key(text))
|
|
101
|
+
if hit is not None:
|
|
102
|
+
out[i] = list(hit)
|
|
103
|
+
else:
|
|
104
|
+
todo.append(i)
|
|
105
|
+
|
|
106
|
+
for start in range(0, len(todo), self.batch_size):
|
|
107
|
+
chunk = todo[start:start + self.batch_size]
|
|
108
|
+
vectors = self.provider.embed([texts[i] for i in chunk])
|
|
109
|
+
if len(vectors) != len(chunk):
|
|
110
|
+
raise ProviderError(
|
|
111
|
+
f"{self.name} returned {len(vectors)} vectors for "
|
|
112
|
+
f"{len(chunk)} inputs")
|
|
113
|
+
for i, vector in zip(chunk, vectors):
|
|
114
|
+
vector = [float(v) for v in vector]
|
|
115
|
+
if len(vector) != self.dim:
|
|
116
|
+
raise ProviderError(
|
|
117
|
+
f"{self.name} returned {len(vector)}-dim vectors but "
|
|
118
|
+
f"dim={self.dim} was declared; a changed embedding "
|
|
119
|
+
"model invalidates every persisted vector, so this "
|
|
120
|
+
"is refused rather than mixed")
|
|
121
|
+
if self.normalize:
|
|
122
|
+
vector = _l2(vector)
|
|
123
|
+
self._cache[self._key(texts[i])] = vector
|
|
124
|
+
out[i] = vector
|
|
125
|
+
|
|
126
|
+
if todo:
|
|
127
|
+
self._save()
|
|
128
|
+
return [v for v in out if v is not None]
|
|
@@ -0,0 +1,268 @@
|
|
|
1
|
+
"""Adapters for hosted AI models. All optional; none is ever required.
|
|
2
|
+
|
|
3
|
+
schemagate works with no AI provider at all -- the default embedder is offline and
|
|
4
|
+
deterministic. A provider buys you two things: better catalog descriptions
|
|
5
|
+
(``schemagate.ai.SchemaDescriber``) and semantic embeddings
|
|
6
|
+
(``schemagate.ai.APIEmbedder``). Both are opt-in and both degrade to the offline
|
|
7
|
+
path if the provider is unavailable.
|
|
8
|
+
|
|
9
|
+
Bring your own key. Nothing here reads a key from anywhere except the
|
|
10
|
+
environment variable of the provider you chose, or the value you pass.
|
|
11
|
+
|
|
12
|
+
``model`` is a required argument on every provider. That is deliberate:
|
|
13
|
+
model identifiers change often, and a library that hardcodes one eventually
|
|
14
|
+
ships a default that 404s for everyone. Pass the model you actually have
|
|
15
|
+
access to.
|
|
16
|
+
|
|
17
|
+
from schemagate.ai import AnthropicProvider, SchemaDescriber
|
|
18
|
+
|
|
19
|
+
provider = AnthropicProvider(model="claude-sonnet-4-5")
|
|
20
|
+
cat.describe(SchemaDescriber(provider))
|
|
21
|
+
|
|
22
|
+
If your provider is not one of the three below, or its SDK changes shape,
|
|
23
|
+
use ``CallableProvider`` and keep control of the call yourself::
|
|
24
|
+
|
|
25
|
+
CallableProvider(lambda system, prompt: my_llm(system, prompt))
|
|
26
|
+
"""
|
|
27
|
+
from __future__ import annotations
|
|
28
|
+
|
|
29
|
+
import os
|
|
30
|
+
from typing import Any, Callable, List, Optional, Protocol, Sequence, runtime_checkable
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
@runtime_checkable
|
|
34
|
+
class Provider(Protocol):
|
|
35
|
+
"""What schemagate needs from a model. Implement either method, or both."""
|
|
36
|
+
|
|
37
|
+
name: str
|
|
38
|
+
|
|
39
|
+
def complete(self, system: str, prompt: str, max_tokens: int = 1024) -> str: ...
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class ProviderError(RuntimeError):
|
|
43
|
+
"""A provider call failed. Callers decide whether that is fatal."""
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
# --------------------------------------------------------------------------
|
|
47
|
+
# Bring your own callable -- the escape hatch that never breaks on an SDK bump
|
|
48
|
+
# --------------------------------------------------------------------------
|
|
49
|
+
|
|
50
|
+
class CallableProvider:
|
|
51
|
+
"""Wrap any function ``f(system, prompt) -> str``.
|
|
52
|
+
|
|
53
|
+
Use this for a provider schemagate does not ship, a gateway, a local model,
|
|
54
|
+
or when an SDK changes and you do not want to wait for a release.
|
|
55
|
+
"""
|
|
56
|
+
|
|
57
|
+
def __init__(self, fn: Callable[[str, str], str], name: str = "callable",
|
|
58
|
+
embed_fn: Optional[Callable[[Sequence[str]], List[List[float]]]] = None):
|
|
59
|
+
self._fn = fn
|
|
60
|
+
self._embed_fn = embed_fn
|
|
61
|
+
self.name = name
|
|
62
|
+
|
|
63
|
+
def complete(self, system: str, prompt: str, max_tokens: int = 1024) -> str:
|
|
64
|
+
return self._fn(system, prompt)
|
|
65
|
+
|
|
66
|
+
def embed(self, texts: Sequence[str]) -> List[List[float]]:
|
|
67
|
+
if self._embed_fn is None:
|
|
68
|
+
raise ProviderError(f"{self.name} was built without an embed_fn")
|
|
69
|
+
return self._embed_fn(texts)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
# --------------------------------------------------------------------------
|
|
73
|
+
# Anthropic (Claude)
|
|
74
|
+
# --------------------------------------------------------------------------
|
|
75
|
+
|
|
76
|
+
class AnthropicProvider:
|
|
77
|
+
"""Claude via the official ``anthropic`` SDK.
|
|
78
|
+
|
|
79
|
+
``pip install anthropic``. Key from ``api_key=`` or ``ANTHROPIC_API_KEY``.
|
|
80
|
+
"""
|
|
81
|
+
|
|
82
|
+
env_var = "ANTHROPIC_API_KEY"
|
|
83
|
+
|
|
84
|
+
def __init__(self, model: str, api_key: Optional[str] = None,
|
|
85
|
+
client: Any = None, timeout: float = 60.0):
|
|
86
|
+
if not model:
|
|
87
|
+
raise ValueError("model is required, e.g. model='claude-sonnet-4-5'")
|
|
88
|
+
self.model = model
|
|
89
|
+
self.name = f"anthropic:{model}"
|
|
90
|
+
if client is not None:
|
|
91
|
+
self._client = client
|
|
92
|
+
return
|
|
93
|
+
try:
|
|
94
|
+
import anthropic
|
|
95
|
+
except ImportError as e:
|
|
96
|
+
raise ImportError("pip install 'schemagate[anthropic]' to use "
|
|
97
|
+
"AnthropicProvider") from e
|
|
98
|
+
key = api_key or os.environ.get(self.env_var)
|
|
99
|
+
if not key:
|
|
100
|
+
raise ValueError(f"no API key: pass api_key= or set {self.env_var}")
|
|
101
|
+
self._client = anthropic.Anthropic(api_key=key, timeout=timeout)
|
|
102
|
+
|
|
103
|
+
def complete(self, system: str, prompt: str, max_tokens: int = 1024) -> str:
|
|
104
|
+
try:
|
|
105
|
+
msg = self._client.messages.create(
|
|
106
|
+
model=self.model, max_tokens=max_tokens, system=system,
|
|
107
|
+
messages=[{"role": "user", "content": prompt}],
|
|
108
|
+
)
|
|
109
|
+
except Exception as e:
|
|
110
|
+
raise ProviderError(f"{self.name}: {e}") from e
|
|
111
|
+
return "".join(b.text for b in msg.content if getattr(b, "type", "") == "text")
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
# --------------------------------------------------------------------------
|
|
115
|
+
# OpenAI (GPT) -- also covers Azure OpenAI and any OpenAI-compatible gateway
|
|
116
|
+
# --------------------------------------------------------------------------
|
|
117
|
+
|
|
118
|
+
class OpenAIProvider:
|
|
119
|
+
"""GPT via the official ``openai`` SDK.
|
|
120
|
+
|
|
121
|
+
``pip install openai``. Key from ``api_key=`` or ``OPENAI_API_KEY``.
|
|
122
|
+
Point ``base_url`` at any OpenAI-compatible endpoint (Azure, vLLM,
|
|
123
|
+
OpenRouter, Ollama) to use this adapter with something else.
|
|
124
|
+
"""
|
|
125
|
+
|
|
126
|
+
env_var = "OPENAI_API_KEY"
|
|
127
|
+
|
|
128
|
+
def __init__(self, model: str, api_key: Optional[str] = None,
|
|
129
|
+
client: Any = None, base_url: Optional[str] = None,
|
|
130
|
+
embed_model: Optional[str] = None, timeout: float = 60.0):
|
|
131
|
+
if not model:
|
|
132
|
+
raise ValueError("model is required, e.g. model='gpt-4.1-mini'")
|
|
133
|
+
self.model = model
|
|
134
|
+
self.embed_model = embed_model
|
|
135
|
+
self.name = f"openai:{model}"
|
|
136
|
+
if client is not None:
|
|
137
|
+
self._client = client
|
|
138
|
+
return
|
|
139
|
+
try:
|
|
140
|
+
import openai
|
|
141
|
+
except ImportError as e:
|
|
142
|
+
raise ImportError("pip install 'schemagate[openai]' to use "
|
|
143
|
+
"OpenAIProvider") from e
|
|
144
|
+
key = api_key or os.environ.get(self.env_var)
|
|
145
|
+
if not key:
|
|
146
|
+
raise ValueError(f"no API key: pass api_key= or set {self.env_var}")
|
|
147
|
+
self._client = openai.OpenAI(api_key=key, base_url=base_url,
|
|
148
|
+
timeout=timeout)
|
|
149
|
+
|
|
150
|
+
def complete(self, system: str, prompt: str, max_tokens: int = 1024) -> str:
|
|
151
|
+
try:
|
|
152
|
+
resp = self._client.chat.completions.create(
|
|
153
|
+
model=self.model, max_tokens=max_tokens,
|
|
154
|
+
messages=[{"role": "system", "content": system},
|
|
155
|
+
{"role": "user", "content": prompt}],
|
|
156
|
+
)
|
|
157
|
+
except Exception as e:
|
|
158
|
+
raise ProviderError(f"{self.name}: {e}") from e
|
|
159
|
+
return resp.choices[0].message.content or ""
|
|
160
|
+
|
|
161
|
+
def embed(self, texts: Sequence[str]) -> List[List[float]]:
|
|
162
|
+
if not self.embed_model:
|
|
163
|
+
raise ProviderError(
|
|
164
|
+
"pass embed_model= to use OpenAIProvider for embeddings, "
|
|
165
|
+
"e.g. embed_model='text-embedding-3-small'")
|
|
166
|
+
try:
|
|
167
|
+
resp = self._client.embeddings.create(model=self.embed_model,
|
|
168
|
+
input=list(texts))
|
|
169
|
+
except Exception as e:
|
|
170
|
+
raise ProviderError(f"{self.name}: {e}") from e
|
|
171
|
+
return [list(d.embedding) for d in resp.data]
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
# --------------------------------------------------------------------------
|
|
175
|
+
# Google (Gemini)
|
|
176
|
+
# --------------------------------------------------------------------------
|
|
177
|
+
|
|
178
|
+
class GeminiProvider:
|
|
179
|
+
"""Gemini via the ``google-genai`` SDK.
|
|
180
|
+
|
|
181
|
+
``pip install google-genai``. Key from ``api_key=``, ``GEMINI_API_KEY``
|
|
182
|
+
or ``GOOGLE_API_KEY``.
|
|
183
|
+
"""
|
|
184
|
+
|
|
185
|
+
env_var = "GEMINI_API_KEY"
|
|
186
|
+
alt_env_var = "GOOGLE_API_KEY"
|
|
187
|
+
|
|
188
|
+
def __init__(self, model: str, api_key: Optional[str] = None,
|
|
189
|
+
client: Any = None, embed_model: Optional[str] = None):
|
|
190
|
+
if not model:
|
|
191
|
+
raise ValueError("model is required, e.g. model='gemini-2.5-flash'")
|
|
192
|
+
self.model = model
|
|
193
|
+
self.embed_model = embed_model
|
|
194
|
+
self.name = f"gemini:{model}"
|
|
195
|
+
if client is not None:
|
|
196
|
+
self._client = client
|
|
197
|
+
return
|
|
198
|
+
try:
|
|
199
|
+
from google import genai
|
|
200
|
+
except ImportError as e:
|
|
201
|
+
raise ImportError("pip install 'schemagate[gemini]' to use "
|
|
202
|
+
"GeminiProvider") from e
|
|
203
|
+
key = (api_key or os.environ.get(self.env_var)
|
|
204
|
+
or os.environ.get(self.alt_env_var))
|
|
205
|
+
if not key:
|
|
206
|
+
raise ValueError(
|
|
207
|
+
f"no API key: pass api_key= or set {self.env_var}")
|
|
208
|
+
self._client = genai.Client(api_key=key)
|
|
209
|
+
|
|
210
|
+
def complete(self, system: str, prompt: str, max_tokens: int = 1024) -> str:
|
|
211
|
+
try:
|
|
212
|
+
resp = self._client.models.generate_content(
|
|
213
|
+
model=self.model, contents=f"{system}\n\n{prompt}")
|
|
214
|
+
except Exception as e:
|
|
215
|
+
raise ProviderError(f"{self.name}: {e}") from e
|
|
216
|
+
return resp.text or ""
|
|
217
|
+
|
|
218
|
+
def embed(self, texts: Sequence[str]) -> List[List[float]]:
|
|
219
|
+
if not self.embed_model:
|
|
220
|
+
raise ProviderError(
|
|
221
|
+
"pass embed_model= to use GeminiProvider for embeddings")
|
|
222
|
+
try:
|
|
223
|
+
resp = self._client.models.embed_content(model=self.embed_model,
|
|
224
|
+
contents=list(texts))
|
|
225
|
+
except Exception as e:
|
|
226
|
+
raise ProviderError(f"{self.name}: {e}") from e
|
|
227
|
+
return [list(e.values) for e in resp.embeddings]
|
|
228
|
+
|
|
229
|
+
|
|
230
|
+
# --------------------------------------------------------------------------
|
|
231
|
+
# Detection
|
|
232
|
+
# --------------------------------------------------------------------------
|
|
233
|
+
|
|
234
|
+
#: checked in order; first key present wins
|
|
235
|
+
_AUTO_ORDER = [
|
|
236
|
+
("ANTHROPIC_API_KEY", AnthropicProvider),
|
|
237
|
+
("OPENAI_API_KEY", OpenAIProvider),
|
|
238
|
+
("GEMINI_API_KEY", GeminiProvider),
|
|
239
|
+
("GOOGLE_API_KEY", GeminiProvider),
|
|
240
|
+
]
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
def available_providers(env: Optional[dict] = None) -> List[str]:
|
|
244
|
+
"""Names of providers whose API key is present. Never returns a key."""
|
|
245
|
+
env = os.environ if env is None else env
|
|
246
|
+
seen, out = set(), []
|
|
247
|
+
for var, cls in _AUTO_ORDER:
|
|
248
|
+
if env.get(var) and cls.__name__ not in seen:
|
|
249
|
+
seen.add(cls.__name__)
|
|
250
|
+
out.append(cls.__name__)
|
|
251
|
+
return out
|
|
252
|
+
|
|
253
|
+
|
|
254
|
+
def auto_provider(model: str, env: Optional[dict] = None, **kwargs):
|
|
255
|
+
"""Build a provider from whichever API key is in the environment.
|
|
256
|
+
|
|
257
|
+
Convenience only. Construct the provider directly when you care which
|
|
258
|
+
one you get -- this picks by key presence, not by capability.
|
|
259
|
+
"""
|
|
260
|
+
env = os.environ if env is None else env
|
|
261
|
+
for var, cls in _AUTO_ORDER:
|
|
262
|
+
if env.get(var):
|
|
263
|
+
return cls(model=model, api_key=env[var], **kwargs)
|
|
264
|
+
raise ValueError(
|
|
265
|
+
"no provider API key found; set one of "
|
|
266
|
+
+ ", ".join(v for v, _ in _AUTO_ORDER)
|
|
267
|
+
+ " or construct a provider directly. schemagate works without one."
|
|
268
|
+
)
|