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

1# SPDX-FileCopyrightText: 2021-2026 Arcangelo Massari <arcangelo.massari@unibo.it> 

2# 

3# SPDX-License-Identifier: ISC 

4 

5 

6import atexit 

7import threading 

8import zipfile 

9 

10from rdflib import Dataset 

11from rdflib.term import Literal, URIRef 

12from sparqlite import SPARQLClient 

13 

14from time_agnostic_library.prov_entity import ProvEntity 

15 

16__all__ = [ 

17 "Sparql", 

18 "_binding_to_n3", 

19 "_n3_to_binding", 

20 "_n3_value", 

21] 

22 

23CONFIG_PATH = "./config.json" 

24 

25_PROV_PROPERTY_STRINGS: tuple[str, ...] = tuple(ProvEntity.get_prov_properties()) 

26 

27_client_cache: dict[tuple[str, int, int, float, float | None], SPARQLClient] = {} 

28_client_lock = threading.Lock() 

29 

30 

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 

49 

50 

51def _close_all_clients() -> None: 

52 with _client_lock: 

53 for client in _client_cache.values(): 

54 client.close() 

55 _client_cache.clear() 

56 

57 

58atexit.register(_close_all_clients) 

59 

60 

61def _escape_n3(v: str) -> str: 

62 return ( 

63 v.replace("\\", "\\\\") 

64 .replace('"', '\\"') 

65 .replace("\n", "\\n") 

66 .replace("\r", "\\r") 

67 ) 

68 

69 

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}"' 

81 

82 

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 

95 

96 

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) 

119 

120 

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 :] 

127 

128 

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 

136 

137 

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} 

149 

150 

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"] 

159 

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) 

175 

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 

183 

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 

207 

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 

216 

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)} 

229 

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 

245 

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 

251 

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))