Coverage for src/time_agnostic_library/sparql.py: 100%
184 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-09-03 21:17 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-09-03 21:17 +0000
1# SPDX-FileCopyrightText: 2021-2026 Arcangelo Massari <arcangelo.massari@unibo.it>
2#
3# SPDX-License-Identifier: ISC
6import atexit
7import threading
8import zipfile
10from rdflib import Dataset
11from rdflib.term import Literal, URIRef
12from sparqlite import SPARQLClient
14from time_agnostic_library.prov_entity import ProvEntity
16__all__ = [
17 "Sparql",
18 "_binding_to_n3",
19 "_n3_to_binding",
20 "_n3_value",
21]
23CONFIG_PATH = "./config.json"
25_PROV_PROPERTY_STRINGS: tuple[str, ...] = tuple(ProvEntity.get_prov_properties())
27_client_cache: dict[tuple[str, int, int, float, float | None], SPARQLClient] = {}
28_client_lock = threading.Lock()
31def _get_client(
32 url: str,
33 max_retries: int = 5,
34 backoff_factor: float = 0.5,
35 timeout: float | None = None,
36) -> SPARQLClient:
37 key = (url, threading.get_ident(), max_retries, backoff_factor, timeout)
38 with _client_lock:
39 client = _client_cache.get(key)
40 if client is None:
41 client = SPARQLClient(
42 url,
43 max_retries=max_retries,
44 backoff_factor=backoff_factor,
45 timeout=timeout,
46 )
47 _client_cache[key] = client
48 return client
51def _close_all_clients() -> None:
52 with _client_lock:
53 for client in _client_cache.values():
54 client.close()
55 _client_cache.clear()
58atexit.register(_close_all_clients)
61def _escape_n3(v: str) -> str:
62 return (
63 v.replace("\\", "\\\\")
64 .replace('"', '\\"')
65 .replace("\n", "\\n")
66 .replace("\r", "\\r")
67 )
70def _binding_to_n3(val: dict) -> str:
71 if val["type"] == "uri":
72 return f"<{val['value']}>"
73 if val["type"] == "bnode":
74 return f"_:{val['value']}"
75 escaped = _escape_n3(val["value"])
76 if "datatype" in val:
77 return f'"{escaped}"^^<{val["datatype"]}>'
78 if "xml:lang" in val:
79 return f'"{escaped}"@{val["xml:lang"]}'
80 return f'"{escaped}"'
83def _find_closing_quote(n3: str) -> int:
84 pos = n3.find('"', 1)
85 while pos > 0:
86 num_backslashes = 0
87 check = pos - 1
88 while check >= 1 and n3[check] == "\\":
89 num_backslashes += 1
90 check -= 1
91 if num_backslashes % 2 == 0:
92 return pos
93 pos = n3.find('"', pos + 1)
94 return -1
97def _unescape_n3(raw: str) -> str:
98 out: list[str] = []
99 i = 0
100 while i < len(raw):
101 if raw[i] == "\\" and i + 1 < len(raw):
102 nxt = raw[i + 1]
103 if nxt == "n":
104 out.append("\n")
105 elif nxt == "r":
106 out.append("\r")
107 elif nxt == '"':
108 out.append('"')
109 elif nxt == "\\":
110 out.append("\\")
111 else:
112 out.append(raw[i])
113 out.append(nxt)
114 i += 2
115 else:
116 out.append(raw[i])
117 i += 1
118 return "".join(out)
121def _parse_n3_literal(n3: str) -> tuple[str, str]:
122 quote_end = _find_closing_quote(n3)
123 if quote_end == -1:
124 return n3, ""
125 raw = n3[1:quote_end]
126 return _unescape_n3(raw), n3[quote_end + 1 :]
129def _n3_value(n3: str) -> str:
130 if n3.startswith("<") and n3.endswith(">"):
131 return n3[1:-1]
132 if n3.startswith("_:"):
133 return n3[2:]
134 value, _ = _parse_n3_literal(n3)
135 return value
138def _n3_to_binding(n3: str) -> dict:
139 if n3.startswith("<") and n3.endswith(">"):
140 return {"type": "uri", "value": n3[1:-1]}
141 if n3.startswith("_:"):
142 return {"type": "bnode", "value": n3[2:]}
143 value, rest = _parse_n3_literal(n3)
144 if rest.startswith("^^<") and rest.endswith(">"):
145 return {"type": "literal", "value": value, "datatype": rest[3:-1]}
146 if rest.startswith("@"):
147 return {"type": "literal", "value": value, "xml:lang": rest[1:]}
148 return {"type": "literal", "value": value}
151class Sparql:
152 def __init__(self, query: str, config: dict):
153 self.query = query
154 self.config = config
155 if any(uri in query for uri in _PROV_PROPERTY_STRINGS):
156 self.storer: dict = config["provenance"]
157 else:
158 self.storer: dict = config["dataset"]
160 def _client(self, url: str) -> SPARQLClient:
161 max_retries = (
162 self.config["sparql_max_retries"]
163 if "sparql_max_retries" in self.config
164 else 5
165 )
166 backoff_factor = (
167 self.config["sparql_backoff_factor"]
168 if "sparql_backoff_factor" in self.config
169 else 0.5
170 )
171 timeout = (
172 self.config["sparql_timeout"] if "sparql_timeout" in self.config else None
173 )
174 return _get_client(url, max_retries, backoff_factor, timeout)
176 def run_select_query(self) -> dict:
177 output = {"head": {"vars": []}, "results": {"bindings": []}}
178 if self.storer["file_paths"]:
179 output = self._get_results_from_files(output)
180 if self.storer["triplestore_urls"]:
181 output = self._get_results_from_triplestores(output)
182 return output
184 def _get_results_from_files(self, output: dict) -> dict:
185 storer: list[str] = self.storer["file_paths"]
186 for file_path in storer:
187 file_cg = Dataset(default_union=True)
188 if file_path.endswith(".zip"):
189 with (
190 zipfile.ZipFile(file_path, "r") as z,
191 z.open(z.namelist()[0]) as file,
192 ):
193 file_cg.parse(file=file, format="json-ld") # type: ignore[arg-type]
194 else:
195 file_cg.parse(location=file_path, format="json-ld")
196 query_results = file_cg.query(self.query)
197 vars_list = [str(var) for var in query_results.vars or []]
198 output["head"]["vars"] = vars_list
199 for result in query_results:
200 binding = {}
201 for var in vars_list:
202 value = result[var] # type: ignore[index]
203 if value is not None:
204 binding[var] = self._format_result_value(value)
205 output["results"]["bindings"].append(binding)
206 return output
208 def _get_results_from_triplestores(self, output: dict) -> dict:
209 storer = self.storer["triplestore_urls"]
210 for url in storer:
211 results = self._client(url).query(self.query)
212 if not output["head"]["vars"]:
213 output["head"]["vars"] = results["head"]["vars"]
214 output["results"]["bindings"].extend(results["results"]["bindings"])
215 return output
217 @staticmethod
218 def _format_result_value(value) -> dict:
219 if isinstance(value, URIRef):
220 return {"type": "uri", "value": str(value)}
221 if isinstance(value, Literal):
222 result = {"type": "literal", "value": str(value)}
223 if value.datatype:
224 result["datatype"] = str(value.datatype)
225 if value.language:
226 result["xml:lang"] = value.language
227 return result
228 return {"type": "literal", "value": str(value)}
230 def run_select_to_quad_set(self) -> set[tuple[str, ...]]:
231 results = self.run_select_query()
232 output: set[tuple[str, ...]] = set()
233 vars_list = results["head"]["vars"]
234 for binding in results["results"]["bindings"]:
235 components: list[str] = []
236 skip = False
237 for var in vars_list:
238 if var not in binding:
239 skip = True
240 break
241 components.append(_binding_to_n3(binding[var]))
242 if not skip:
243 output.add(tuple(components))
244 return output
246 def run_ask_query(self) -> bool:
247 storer = self.storer["triplestore_urls"]
248 for url in storer:
249 return self._client(url).ask(self.query)
250 return False
252 @classmethod
253 def _get_tuples_set(cls, result_dict: dict, output: set, vars_list: list) -> None:
254 results_list = []
255 for var in vars_list:
256 if str(var) in result_dict:
257 val = result_dict[str(var)]
258 if isinstance(val, dict) and "value" in val:
259 results_list.append(str(val["value"]))
260 else:
261 results_list.append(str(val))
262 else:
263 results_list.append(None)
264 output.add(tuple(results_list))