relrae 0.9.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.
relrae/RELRaE.py ADDED
@@ -0,0 +1,91 @@
1
+ from rdflib import Graph, Namespace
2
+ from rdflib.namespace import RDF, RDFS, OWL
3
+ import xmlschema
4
+ from pathlib import Path
5
+
6
+
7
+ from .modules.RuBREx import RuBREx
8
+ from .modules.LLM_Ref import LLMRefinement
9
+ from .modules.human_fix import HumanFix
10
+ from .utils import get_now
11
+
12
+
13
+ class RELRaE:
14
+
15
+ def __init__(self, onto_name, schema, namespace, prefix, config,
16
+ components_root="relrae_components", output_root="output"):
17
+ self.components_root = Path(components_root)
18
+ self.output_root = Path(output_root)
19
+ schema_path = self.components_root / "schema" / schema
20
+ self.schema = xmlschema.XMLSchema(schema_path)
21
+ self.namespace = Namespace(namespace)
22
+ self.prefix = prefix
23
+ self.onto_name = onto_name
24
+ self.onto = Graph()
25
+ self.errors = []
26
+ self.config = config
27
+ self.configs = [{"Pipeline": config}]
28
+ self.logs = {}
29
+
30
+ self.onto.bind('rdf', RDF)
31
+ self.onto.bind('rdfs', RDFS)
32
+ self.onto.bind('owl', OWL)
33
+ self.onto.bind(self.prefix, self.namespace)
34
+
35
+ def RuBREx(self):
36
+ m_rubrex = RuBREx(
37
+ self.schema, self.onto, self.prefix, self.namespace,
38
+ self.components_root,
39
+ )
40
+ m_rubrex.match_concepts()
41
+ m_rubrex.schema_coverage()
42
+
43
+ self.configs.append({"RuBREx": m_rubrex.config})
44
+ self.logs["RuBREx"] = m_rubrex.log
45
+ # NOTE: Consider the module complete
46
+ self.onto = self.onto + m_rubrex.onto
47
+ self.errors.append(m_rubrex.errors)
48
+
49
+ def LLM_Refinement_Loop(self):
50
+ m_LLM_ref = LLMRefinement(
51
+ self.schema, self.onto, self.prefix, self.namespace,
52
+ self.components_root)
53
+ m_LLM_ref.evaluate_relations()
54
+ self.errors.append(m_LLM_ref.errors)
55
+
56
+ def human_fix(self):
57
+ # TODO:
58
+ m_human_fix = HumanFix(
59
+ self.schema, self.onto, self.prefix, self.namespace,
60
+ self.components_root, self.errors, self.config["modules"]
61
+ )
62
+ m_human_fix.set_user_info()
63
+ m_human_fix.fix_errors()
64
+ m_human_fix.concepts_to_add_output(path)
65
+
66
+ def write_logs(self, path):
67
+ for key, values in self.logs.items():
68
+ with (path / f"{key}.txt").open("w", encoding="utf-8") as f:
69
+ for line in values:
70
+ f.write(line + "\n")
71
+
72
+ def write_metadata(self, path):
73
+ with (path / "metadata.txt").open("w", encoding="utf-8") as f:
74
+ for config in self.configs:
75
+ for section, values in config.items():
76
+ f.write(f"{section}\n\n")
77
+ for key, value in values.items():
78
+ f.write(f"{key}: {value}\n")
79
+ f.write("\n\n")
80
+
81
+ def serialise(self):
82
+ timestamp = get_now().replace(":", "-")
83
+ main_path = self.output_root / f"{self.onto_name}{timestamp}"
84
+ log_path = main_path / "logs"
85
+ log_path.mkdir(parents=True)
86
+ self.write_logs(log_path)
87
+ self.write_metadata(main_path)
88
+ self.onto.serialize(
89
+ destination=main_path / f"{self.onto_name}.ttl",
90
+ format="ttl",
91
+ )
relrae/__init__.py ADDED
@@ -0,0 +1,6 @@
1
+ """Public package interface for RELRaE."""
2
+
3
+ from ._version import __version__
4
+ from .RELRaE import RELRaE
5
+
6
+ __all__ = ["RELRaE", "__version__"]
relrae/__main__.py ADDED
@@ -0,0 +1,4 @@
1
+ from .cli import main
2
+
3
+ if __name__ == "__main__":
4
+ main()
relrae/_version.py ADDED
@@ -0,0 +1,3 @@
1
+ """The single source of truth for the RELRaE package version."""
2
+
3
+ __version__ = "0.9.0"
relrae/cli.py ADDED
@@ -0,0 +1,65 @@
1
+ import argparse
2
+ from importlib.resources import files
3
+ from pathlib import Path
4
+ import shutil
5
+
6
+
7
+ COMPONENT_DIRECTORIES = ("config", "rules", "schema", "output")
8
+
9
+
10
+ def _copy_defaults(source, destination):
11
+ """Copy packaged defaults without replacing user-edited files."""
12
+ for resource in source.iterdir():
13
+ # Finder metadata is not part of a RELRaE project, even if it happens
14
+ # to be present in a source checkout.
15
+ if resource.name.startswith("."):
16
+ continue
17
+
18
+ target = destination / resource.name
19
+ if resource.is_dir():
20
+ target.mkdir(exist_ok=True)
21
+ _copy_defaults(resource, target)
22
+ elif not target.exists():
23
+ with resource.open("rb") as source_file, target.open("xb") as target_file:
24
+ shutil.copyfileobj(source_file, target_file)
25
+
26
+
27
+ def initialise_project(project_root=None):
28
+ """Create a RELRaE project beneath *project_root* (the cwd by default).
29
+
30
+ Existing files are intentionally left untouched, making ``relrae init``
31
+ safe to run again after configuration files have been customised.
32
+ """
33
+ project_root = Path.cwd() if project_root is None else Path(project_root)
34
+ components_root = project_root / "relrae_components"
35
+
36
+ project_root.mkdir(parents=True, exist_ok=True)
37
+ for directory in COMPONENT_DIRECTORIES:
38
+ (components_root / directory).mkdir(parents=True, exist_ok=True)
39
+
40
+ packaged_defaults = files("relrae").joinpath("relrae_components")
41
+ _copy_defaults(packaged_defaults, components_root)
42
+
43
+ print(f"Initialised RELRaE project in {components_root}")
44
+ return components_root
45
+
46
+
47
+ def main():
48
+ parser = argparse.ArgumentParser()
49
+ subparsers = parser.add_subparsers(dest="command")
50
+
51
+ init_parser = subparsers.add_parser(
52
+ "init", help="create a RELRaE project in the current directory"
53
+ )
54
+ init_parser.add_argument(
55
+ "directory",
56
+ nargs="?",
57
+ default=Path.cwd(),
58
+ type=Path,
59
+ help="project directory (default: current directory)",
60
+ )
61
+
62
+ args = parser.parse_args()
63
+
64
+ if args.command == "init":
65
+ initialise_project(args.directory)
@@ -0,0 +1,525 @@
1
+ import json
2
+ import os
3
+ import re
4
+ import configparser
5
+ from pathlib import Path
6
+ from urllib import request
7
+ from urllib.error import URLError
8
+ from typing import Literal
9
+ from pydantic import BaseModel, Field, config
10
+ from ..utils import get_now, namespace_to_prefix
11
+ from random import randint
12
+ import numpy as np
13
+ from rdflib.namespace import RDFS
14
+ from rdflib import URIRef
15
+ import rdflib
16
+
17
+ class LLMRefinement:
18
+
19
+ def __init__(self, schema, onto, prefix, namespace,
20
+ components_root="relrae_components"):
21
+ self.log = []
22
+ cfg = configparser.ConfigParser()
23
+ cfg.read(Path(components_root) / "config" / "LLM_ref_conf.txt")
24
+ self.config = cfg["MAIN"]
25
+ self.errors = []
26
+ self.schema = schema
27
+ self.onto = onto
28
+ self.prefix = prefix
29
+ self.namespace = namespace
30
+ self.set_models()
31
+
32
+ # NOTE: This might not be neccessary
33
+ def set_models(self):
34
+ pass
35
+
36
+ def get_relations(self):
37
+ query = """
38
+ PREFIX rdf: <http://www.w3.org/1999/02/22-rdf-syntax-ns#>
39
+ PREFIX owl: <http://www.w3.org/2002/07/owl#>
40
+
41
+ SELECT ?property ?propertyType
42
+ WHERE {
43
+ VALUES ?propertyType {
44
+ rdf:Property
45
+ owl:ObjectProperty
46
+ owl:DatatypeProperty
47
+ owl:AnnotationProperty
48
+ }
49
+
50
+ ?property rdf:type ?propertyType .
51
+ }
52
+ """
53
+ qres = self.onto.query(query)
54
+ property_list = []
55
+ for r in qres:
56
+ property_list.append(r.property)
57
+ self.log.append(f"Retrieved {len(property_list)} properties")
58
+ return property_list
59
+
60
+ def get_relation_info(self, relation):
61
+ query = f"""
62
+ PREFIX rdf: <http://www.w3.org/1999/02/22-rdf-syntax-ns#>
63
+ PREFIX rdfs: <http://www.w3.org/2000/01/rdf-schema#>
64
+ PREFIX owl: <http://www.w3.org/2002/07/owl#>
65
+
66
+ SELECT ?subject ?predicate ?object
67
+ WHERE {{
68
+ {{
69
+ BIND(<{relation}> AS ?subject)
70
+ ?subject ?predicate ?object .
71
+ }}
72
+ UNION
73
+ {{
74
+ ?subject ?predicate <{relation}> .
75
+ BIND(<{relation}> AS ?object)
76
+ }}
77
+ }}
78
+ """
79
+ qres = self.onto.query(query)
80
+ prefixed_relation = namespace_to_prefix(
81
+ relation, self.namespace, self.prefix)
82
+ rel_info = {prefixed_relation: []}
83
+ for r in qres:
84
+ rel_info[prefixed_relation].append([
85
+ namespace_to_prefix(r.predicate, self.namespace, self.prefix),
86
+ namespace_to_prefix(r.object, self.namespace, self.prefix)
87
+ ])
88
+ self.log.append(f"Retrieved info for relation {str(relation)}")
89
+ # print(rel_info)
90
+ return rel_info
91
+
92
+ def normalise_relation_label(self, label):
93
+ label = namespace_to_prefix(label, self.namespace, self.prefix).strip()
94
+ prefix_marker = f"{self.prefix}:"
95
+
96
+ if ":" in label:
97
+ prefix, label = label.split(":", 1)
98
+ if prefix != self.prefix:
99
+ prefix = self.prefix
100
+ else:
101
+ prefix = self.prefix
102
+
103
+ words = re.findall(r"[A-Za-z0-9]+", label)
104
+ if not words:
105
+ raise ValueError(f"Cannot normalise empty relation label: {label}")
106
+
107
+ camel_label = words[0].lower()
108
+ for word in words[1:]:
109
+ camel_label += word[:1].upper() + word[1:].lower()
110
+
111
+ return f"{prefix_marker}{camel_label}"
112
+
113
+ def get_examples(self, model_role, strat):
114
+ if model_role == "eval":
115
+ examples = self.config["eval_examples"]
116
+ else:
117
+ examples = self.config["ref_examples"]
118
+
119
+ # print(examples)
120
+ shots = {"zero": None,
121
+ "one": 0,
122
+ "few": 4}
123
+
124
+ if strat == "zero":
125
+ return []
126
+
127
+ examples_list = examples[0:shots[strat]]
128
+ return examples_list
129
+
130
+ def average_eval(self, responses):
131
+ evals = []
132
+ confs = []
133
+ for r in responses:
134
+ if r["evaluation"] == "Yes":
135
+ evals.append(1)
136
+ else:
137
+ evals.append(0)
138
+ confs.append(r["confidence"])
139
+ mean_evals = round((sum(evals)/len(evals)), 0)
140
+ if int(mean_evals) == 1:
141
+ eval = "Yes"
142
+ else:
143
+ eval = "No"
144
+ avg_eval = {
145
+ "evaluation": eval,
146
+ "confidence": round(sum(confs)/len(confs), 3),
147
+ "variance": round(float(np.var(evals)), 3)
148
+ }
149
+ return avg_eval
150
+
151
+ def evaluation_loop(self, relation, evaluator, refiner, info):
152
+ accepted = False
153
+ rejected_labels = []
154
+ loops = 0
155
+ active_label = relation
156
+
157
+ while not accepted and loops < int(self.config["refinement_loops"]):
158
+ loops += 1
159
+ eval_model = evaluator[0]
160
+ eval_prompt = eval_model.build_eval_prompt(
161
+ namespace_to_prefix(active_label, self.namespace, self.prefix),
162
+ self.config["domain"],
163
+ self.config["source"],
164
+ info,
165
+ rejected_labels,
166
+ self.config["eval_messages"],
167
+ evaluator[1])
168
+ # print(eval_prompt)
169
+ full_eval = eval_model.run_prompt(eval_prompt)
170
+ if full_eval[0] == "Error":
171
+ self.log.append("Error: LLM could not generate valid response")
172
+ else:
173
+ avg_eval = self.average_eval(full_eval[0])
174
+ print(active_label)
175
+ print("\n=======================\n")
176
+ print(avg_eval)
177
+ print("\n=======================\n")
178
+ print(full_eval[1])
179
+
180
+ self.log.append(full_eval[1])
181
+ if avg_eval["evaluation"] == "Yes":
182
+ self.log.append("Label accepted")
183
+ break
184
+
185
+ self.log.append("Label rejected")
186
+ rejected_labels.append(namespace_to_prefix(
187
+ active_label, self.namespace, self.prefix))
188
+ print("\n ==== Rejected -> Refining ==== \n")
189
+ ref_model = refiner[0]
190
+ ref_prompt = ref_model.build_ref_prompt(
191
+ namespace_to_prefix(active_label, self.namespace, self.prefix),
192
+ self.config["domain"],
193
+ self.config["source"],
194
+ info,
195
+ rejected_labels,
196
+ self.config["ref_messages"],
197
+ refiner[1])
198
+ full_ref = ref_model.run_prompt(ref_prompt)
199
+ ref = full_ref[0][0]
200
+ print(ref)
201
+ print(rejected_labels)
202
+ print("\n=======================\n")
203
+ print(full_ref[1])
204
+ print("\n=======================\n")
205
+ self.log.append(full_ref[1])
206
+ active_label = self.normalise_relation_label(
207
+ ref["relationship_label"])
208
+
209
+ if accepted:
210
+ self.replace_relation(relation, active_label)
211
+ return ["refined", rejected_labels]
212
+ else:
213
+ return ["unrefined", rejected_labels]
214
+
215
+ def evaluate_relations(self):
216
+ relations = self.get_relations()
217
+
218
+ eval_m = LLM(self.config["eval_llm"],
219
+ EvaluatorResponse,
220
+ self.config["eval_repeats"],
221
+ self.namespace,
222
+ self.prefix)
223
+ eval_ex = self.get_examples("eval", self.config["eval_prompt_strat"])
224
+
225
+ ref_m = LLM(self.config["ref_llm"],
226
+ RefinerResponse,
227
+ self.config["ref_repeats"],
228
+ self.namespace,
229
+ self.prefix)
230
+ ref_ex = self.get_examples("ref", self.config["ref_prompt_strat"])
231
+
232
+ for rel in relations:
233
+ rel_info = self.get_relation_info(rel)
234
+ response = self.evaluation_loop(rel, [eval_m, eval_ex], [
235
+ ref_m, ref_ex], rel_info)
236
+ if response[0] == "unrefined":
237
+ self.errors.append([rel, rel_info, response[1]])
238
+
239
+ def replace_relation(self, original, new):
240
+ query = f"""
241
+ PREFIX rdf: <http://www.w3.org/1999/02/22-rdf-syntax-ns#>
242
+ PREFIX rdfs: <http://www.w3.org/2000/01/rdf-schema#>
243
+ PREFIX owl: <http://www.w3.org/2002/07/owl#>
244
+
245
+ SELECT ?s ?p ?o WHERE{{
246
+ ?s ?p ?o
247
+ FITLER(
248
+ ?s = <{self.namespace}{original}> ||
249
+ ?p = <{self.namespace}{original}> ||
250
+ ?o = <{self.namespace}{original}>
251
+ )
252
+ }}
253
+ """
254
+ qres = self.onto.query(query)
255
+
256
+ for r in qres:
257
+
258
+ if r.p == str(RDFS.label):
259
+ n_sub = URIRef(f"{self.namespace}{new}")
260
+ n_obj = rdflib.Literal(new)
261
+ self.onto.remove((r.s, r.p, r.o))
262
+ self.onto.add((n_sub, r.p, n_obj))
263
+ elif r.s == f"{self.namespace}{original}":
264
+ n_sub = URIRef(f"{self.namespace}{new}")
265
+ self.onto.remove((r.s, r.p, r.o))
266
+ self.onto.add((n_sub, r.p, r.o))
267
+ elif r.p == f"{self.namespace}{original}":
268
+ n_pred = URIRef(f"{self.namespace}{new}")
269
+ self.onto.remove((r.s, r.p, r.o))
270
+ self.onto.add((r.s, n_pred, r.o))
271
+ elif r.o == f"{self.namespace}{original}":
272
+ n_obj = URIRef(f"{self.namespace}{new}")
273
+ self.onto.remove((r.s, r.p, r.o))
274
+ self.onto.add((r.s, r.p, n_obj))
275
+
276
+ s = URIRef(f"{self.namespace}{new}")
277
+ p = URIRef(f"{self.namespace}editedBy")
278
+ o = rdflib.Literal(f"{self.config["eval_llm"]} @ {get_now()}")
279
+ self.onto.add([s, p, o])
280
+ self.log.append(f"Relationship {original} updated to {new}.")
281
+
282
+ class EvaluatorResponse(BaseModel):
283
+ evaluation: Literal["Yes", "No"]
284
+ justification: str
285
+ confidence: float = Field(ge=0.00, le=100.00)
286
+
287
+
288
+ class RefinerResponse(BaseModel):
289
+ relationship_label: str
290
+ justification: str
291
+
292
+
293
+ class LLM:
294
+
295
+ def __init__(self, settings, format, repeats, ns, pr):
296
+ self.repeats = repeats
297
+ self.model = settings[0]
298
+ self.api_key = settings[1]
299
+ self.params = settings[2]
300
+ self.provider = self.get_provider(settings)
301
+ self.r_format = format
302
+ self.onto_ns = ns
303
+ self.onto_pr = pr
304
+
305
+ def gen_seeds(self, n):
306
+ seeds = []
307
+ for i in range(int(n)):
308
+ seeds.append(randint(1, 9999))
309
+ return seeds
310
+
311
+ def get_provider(self, settings):
312
+ if len(settings) > 3:
313
+ return str(settings[3]).lower()
314
+ if isinstance(self.params, dict) and self.params.get("provider"):
315
+ return str(self.params["provider"]).lower()
316
+ if str(self.api_key).lower() in ("", "none", "ollama"):
317
+ return "ollama"
318
+ if str(self.model).lower().startswith(("gpt", "o1", "o3", "o4")):
319
+ return "openai"
320
+ if str(self.model).lower().startswith(("gemini", "models/gemini")):
321
+ return "google"
322
+ return "ollama"
323
+
324
+ def build_messages(self, prompt, examples, active):
325
+ message_list = [{"role": "system", "content": prompt}]
326
+ for example in examples:
327
+ for e in example:
328
+ message_list.append(e)
329
+ message_list.append(active)
330
+ return message_list
331
+
332
+ def process_context(self, context):
333
+ context_list = []
334
+ reject_relations = [
335
+ "http://www.w3.org/2000/01/rdf-schema#label",
336
+ f"{self.onto_pr}:generatedBy",
337
+ f"{self.onto_pr}:hasXSDSource"
338
+ ]
339
+ for l in context.keys():
340
+ for r in context[l]:
341
+ if r[0] in reject_relations:
342
+ continue
343
+ triple = [l, r[0], r[1]]
344
+ context_list.append(triple)
345
+ print(context_list)
346
+ return context_list
347
+
348
+ def build_eval_prompt(self, relation, domain, source, info, rejected_label, prompt, examples):
349
+ active = {"role": "user", "content": f"relation: {relation}, context: {
350
+ self.process_context(info)}, domain: {domain}, source: {source}, rejected labels: {rejected_label}"}
351
+
352
+ messages = self.build_messages(prompt[0], examples, active)
353
+ return messages
354
+
355
+ def build_ref_prompt(self, relation, domain, source, info, rejected_label, prompt, examples):
356
+ active = {"role": "user", "content": f"relation: {relation}, context: {
357
+ self.process_context(info)}, domain: {domain}, source: {source}, rejected_labels: {rejected_label}"}
358
+
359
+ messages = self.build_messages(prompt[0], examples, active)
360
+ return messages
361
+
362
+ def run_prompt(self, messages):
363
+ results = []
364
+ logs = []
365
+ seeds = []
366
+ seed = 0
367
+
368
+ while seed < int(self.repeats):
369
+ retrys = 10
370
+ valid = False
371
+ while not valid and retrys > 0:
372
+ c_seed = randint(1, 9999)
373
+ try:
374
+ raw_response = self.call_provider(messages, c_seed)
375
+ parsed_response = self.parse_response(raw_response)
376
+ results.append(parsed_response)
377
+ logs.append({
378
+ "provider": self.provider,
379
+ "model": self.model,
380
+ "seed": c_seed,
381
+ "response": parsed_response,
382
+ })
383
+ valid = True
384
+ seeds.append(c_seed)
385
+ seed += 1
386
+ retrys -= 1
387
+ except Exception as exc:
388
+ logs.append({
389
+ "provider": self.provider,
390
+ "model": self.model,
391
+ "seed": c_seed,
392
+ "error": str(exc),
393
+ })
394
+ retrys -= 1
395
+
396
+ if not results:
397
+ results = "Error"
398
+
399
+ return [results, logs]
400
+
401
+ def call_provider(self, messages, seed):
402
+ if self.provider == "ollama":
403
+ return self.call_ollama(messages, seed)
404
+ if self.provider == "openai":
405
+ return self.call_openai(messages, seed)
406
+ if self.provider in ("google", "googleai", "gemini"):
407
+ return self.call_google(messages, seed)
408
+ raise ValueError(f"Unsupported LLM provider: {self.provider}")
409
+
410
+ def get_params(self, seed, provider=None):
411
+ params = {}
412
+ if isinstance(self.params, dict):
413
+ params.update(self.params)
414
+ params.pop("provider", None)
415
+ params["seed"] = seed
416
+ if provider == "openai":
417
+ supported = {
418
+ "temperature",
419
+ "top_p",
420
+ "max_tokens",
421
+ "max_completion_tokens",
422
+ "presence_penalty",
423
+ "frequency_penalty",
424
+ "seed",
425
+ "stop",
426
+ }
427
+ params = {
428
+ key: value for key, value in params.items()
429
+ if key in supported
430
+ }
431
+ return params
432
+
433
+ def get_api_key(self, env_var):
434
+ if str(self.api_key).lower() not in ("", "none"):
435
+ return self.api_key
436
+ return os.environ.get(env_var)
437
+
438
+ def parse_response(self, raw_response):
439
+ if isinstance(raw_response, self.r_format):
440
+ return raw_response.model_dump()
441
+ if isinstance(raw_response, dict):
442
+ return self.r_format.model_validate(raw_response).model_dump()
443
+ return self.r_format.model_validate_json(raw_response).model_dump()
444
+
445
+ def messages_to_text(self, messages):
446
+ prompt_parts = []
447
+ for message in messages:
448
+ prompt_parts.append(f"{message['role']}: {message['content']}")
449
+ return "\n\n".join(prompt_parts)
450
+
451
+ def call_ollama(self, messages, seed):
452
+ params = self.get_params(seed, "ollama")
453
+ payload = {
454
+ "model": self.model,
455
+ "messages": messages,
456
+ "stream": False,
457
+ "format": self.r_format.model_json_schema(),
458
+ "options": params,
459
+ }
460
+ data = json.dumps(payload).encode("utf-8")
461
+ req = request.Request(
462
+ "http://localhost:11434/api/chat",
463
+ data=data,
464
+ headers={"Content-Type": "application/json"},
465
+ method="POST",
466
+ )
467
+
468
+ try:
469
+ with request.urlopen(req, timeout=120) as response:
470
+ body = json.loads(response.read().decode("utf-8"))
471
+ except URLError as exc:
472
+ raise RuntimeError(
473
+ "Could not reach Ollama at localhost:11434") from exc
474
+
475
+ return body["message"]["content"]
476
+
477
+ def call_openai(self, messages, seed):
478
+ try:
479
+ from openai import OpenAI
480
+ except ImportError as exc:
481
+ raise ImportError(
482
+ "Install the openai package to use OpenAI models") from exc
483
+
484
+ api_key = self.get_api_key("OPENAI_API_KEY")
485
+ if not api_key:
486
+ raise ValueError("OpenAI API key missing")
487
+
488
+ params = self.get_params(seed, "openai")
489
+ client = OpenAI(api_key=api_key)
490
+ completion = client.chat.completions.parse(
491
+ model=self.model,
492
+ messages=messages,
493
+ response_format=self.r_format,
494
+ **params,
495
+ )
496
+ return completion.choices[0].message.parsed
497
+
498
+ def call_google(self, messages, seed):
499
+ try:
500
+ from google import genai
501
+ except ImportError as exc:
502
+ raise ImportError(
503
+ "Install the google-genai package to use Google AI models"
504
+ ) from exc
505
+
506
+ api_key = self.get_api_key("GEMINI_API_KEY")
507
+ if not api_key:
508
+ raise ValueError("Google AI API key missing")
509
+
510
+ params = self.get_params(seed, "google")
511
+ client = genai.Client(api_key=api_key)
512
+ response = client.models.generate_content(
513
+ model=self.model,
514
+ contents=self.messages_to_text(messages),
515
+ config={
516
+ **params,
517
+ "response_format": {
518
+ "text": {
519
+ "mime_type": "application/json",
520
+ "schema": self.r_format.model_json_schema(),
521
+ }
522
+ },
523
+ },
524
+ )
525
+ return response.text