Coverage for gws-app/gws/lib/xmlx/validator.py: 77%

111 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-10-05 13:35 +0200

1"""XML schema validator for tests.""" 

2 

3import re 

4import os 

5import lxml.etree 

6import requests 

7 

8import gws 

9 

10from . import util 

11 

12 

13class Error(gws.Error): 

14 """Validation or schema error, with ``message`` and ``lineno``.""" 

15 

16 def __init__(self, *args, **kwargs): 

17 """Create the error from a message and a line number.""" 

18 

19 super().__init__(*args, **kwargs) 

20 self.message = args[0] 

21 self.lineno = args[1] 

22 

23 

24def validate(xml: str | bytes) -> bool: 

25 """Validate a document against the schemas listed in its ``xsi:schemaLocation``. 

26 

27 Remote schemas are downloaded and cached under ``gws.c.CACHE_DIR``; URLs containing ``.loc`` or ``local`` 

28 are downloaded without caching. 

29 

30 Args: 

31 xml: The document as a string or bytes. 

32 

33 Returns: 

34 ``True`` if the document is valid. 

35 

36 Raises: 

37 Error: If the document or a schema cannot be parsed, or the document is invalid. 

38 """ 

39 

40 try: 

41 parser = lxml.etree.XMLParser(resolve_entities=False, no_network=True) 

42 parser.resolvers.add(_CachingResolver()) 

43 

44 schema_locations = _extract_schema_locations(xml) 

45 xsd = _create_combined_xsd(schema_locations) 

46 

47 xml_tree = _etree(xml, parser) 

48 schema_tree = _etree(xsd, parser) 

49 schema = lxml.etree.XMLSchema(schema_tree) 

50 except lxml.etree.Error as exc: 

51 raise _error(exc) from exc 

52 

53 try: 

54 schema.assertValid(xml_tree) 

55 return True 

56 except Exception as exc: 

57 raise _error(exc) from exc 

58 

59 

60def _extract_schema_locations(xml: str | bytes) -> dict: 

61 """Read the ``schemaLocation`` of the root element as a dict URI -> location.""" 

62 

63 tree = _etree(xml, None) 

64 root = tree.getroot() 

65 

66 xsi_ns = '{http://www.w3.org/2001/XMLSchema-instance}' 

67 attr = root.get(f'{xsi_ns}schemaLocation') 

68 if not attr: 

69 attr = root.get('schemaLocation') 

70 if not attr: 

71 return {} 

72 

73 d = {} 

74 

75 parts = attr.strip().split() 

76 while len(parts) >= 2: 

77 namespace = parts.pop(0) 

78 location = parts.pop(0) 

79 d[namespace] = location 

80 

81 return d 

82 

83 

84def _create_combined_xsd(schema_locations: dict) -> str: 

85 """Create a schema that imports all given schemas.""" 

86 

87 xml = [] 

88 xml.append('<?xml version="1.0" encoding="UTF-8"?>') 

89 xml.append('<xs:schema xmlns:xs="http://www.w3.org/2001/XMLSchema">') 

90 

91 for ns, loc in schema_locations.items(): 

92 xml.append(f'<xs:import namespace="{util.escape_attribute(ns)}" schemaLocation="{util.escape_attribute(loc)}"/>') 

93 

94 xml.append('</xs:schema>\n') 

95 

96 return '\n'.join(xml) 

97 

98 

99def _etree(xml: str | bytes, parser: lxml.etree.XMLParser | None) -> lxml.etree.ElementTree: 

100 """Parse a document with lxml.""" 

101 

102 if isinstance(xml, str): 

103 xml = xml.encode('utf-8') 

104 return lxml.etree.ElementTree(lxml.etree.fromstring(xml, parser)) 

105 

106 

107def _error(exc): 

108 """Convert an lxml exception to an ``Error``.""" 

109 

110 # exc is either {'message': ..., 'lineno': ...} 

111 # or {'error_log': '<string>:17:0:ERROR:...} 

112 

113 cls = exc.__class__.__name__ 

114 

115 s = getattr(exc, 'error_log', None) 

116 if s: 

117 try: 

118 lineno = int(s.split(':')[1]) 

119 except Exception: 

120 lineno = 0 

121 return Error(f'{cls}: {s}', lineno) 

122 

123 lineno = getattr(exc, 'lineno', 0) 

124 return Error(f'{cls}: {exc}', lineno) 

125 

126 

127class _CachingResolver(lxml.etree.Resolver): 

128 """lxml resolver that downloads remote schemas, with a file cache.""" 

129 

130 def resolve(self, url, id, context): 

131 if url.startswith(('http://', 'https://')): 

132 if '.loc' in url or 'local' in url: 

133 buf = _download_url(url, with_cache=False) 

134 else: 

135 buf = _download_url(url, with_cache=True) 

136 return self.resolve_string(buf, context, base_url=url) 

137 

138 return super().resolve(url, id, context) 

139 

140 

141def _download_url(url: str, with_cache: bool) -> bytes: 

142 """Download a URL, optionally using the file cache.""" 

143 

144 if not with_cache: 

145 return _raw_download_url(url) 

146 

147 cache_dir = gws.u.ensure_dir(gws.c.CACHE_DIR + '/xmlx') 

148 cache_path = _cache_path(cache_dir, url) 

149 

150 if os.path.exists(cache_path): 

151 return gws.u.read_file_b(cache_path) 

152 

153 content = _raw_download_url(url) 

154 gws.u.write_file_b(cache_path, content) 

155 return content 

156 

157 

158def _raw_download_url(url: str) -> bytes: 

159 """Download a URL, raise ``ValueError`` unless the status is 200.""" 

160 

161 gws.log.debug(f'xmlx.validator: downloading {url!r}') 

162 response = requests.get(url, timeout=10) 

163 if response.status_code != 200: 

164 raise ValueError(f'Failed to download {url!r}: {response.status_code}') 

165 return response.content 

166 

167 

168def _cache_path(cache_dir: str, url: str) -> str: 

169 """Get the cache file path for a URL, creating its directory.""" 

170 

171 u = url.strip().split('//')[-1] 

172 if '?' in u: 

173 u = u.split('?', 1)[0] 

174 fname = 'index.xml' 

175 parts = u.split('/') 

176 

177 if u.endswith('/'): 

178 parts.pop() 

179 else: 

180 m = re.search(r'[^/]+\.[a-z]+$', parts[-1]) 

181 if m: 

182 fname = m.group(0) 

183 parts.pop() 

184 

185 d = '/'.join(_to_dirname(p) for p in parts) 

186 if not d: 

187 return cache_dir + '/' + fname 

188 d = gws.u.ensure_dir(cache_dir + '/' + d) 

189 return d + '/' + fname 

190 

191 

192def _to_dirname(s: str) -> str: 

193 """Convert a URL path component to a safe directory name.""" 

194 

195 s = s.lower().strip().lstrip('.') 

196 s = re.sub(r'[^a-zA-Z0-9.]+', '_', s).strip('_') 

197 return s