providers.py 9.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291
  1. import json
  2. import logging
  3. import os
  4. from abc import ABC, abstractmethod
  5. from pathlib import Path
  6. from threading import Lock
  7. from typing import Optional
  8. import jsonschema
  9. import requests
  10. from metadata_manager.vehicles_manager import DEFAULT_VEHICLES
  11. from metadata_manager.firmware_server import ManifestJSON
  12. from metadata_manager.firmware_server.models import ReleaseRecord
  13. from .models import ForkRemoteSpec, RemoteInfo, VersionInfo
  14. OFFICIAL_REMOTE_NAME = "ardupilot"
  15. OFFICIAL_REMOTE_URL = "https://github.com/ardupilot/ardupilot.git"
  16. def _vehicle_id_for_tag_segment(segment: str) -> Optional[str]:
  17. for vehicle in DEFAULT_VEHICLES:
  18. if vehicle.name == segment or vehicle.id == segment:
  19. return vehicle.id
  20. return None
  21. DEFAULT_WHITELISTED_FORK_REMOTES = [
  22. ForkRemoteSpec(owner="tridge", repo="ardupilot"),
  23. ForkRemoteSpec(owner="peterbarker", repo="ardupilot"),
  24. ForkRemoteSpec(owner="rmackay9", repo="rmackay9-ardupilot"),
  25. ForkRemoteSpec(owner="shiv-tyagi", repo="ardupilot"),
  26. ForkRemoteSpec(owner="andyp1per", repo="ardupilot"),
  27. ]
  28. class VersionsProvider(ABC):
  29. @property
  30. @abstractmethod
  31. def name(self) -> str:
  32. ...
  33. @abstractmethod
  34. def get_versions(self, vehicle_id: str) -> list[VersionInfo]:
  35. ...
  36. @abstractmethod
  37. def get_remotes(self) -> list[RemoteInfo]:
  38. ...
  39. def refresh(self) -> None:
  40. """Optional hook to refresh provider data."""
  41. @property
  42. def is_available(self) -> bool:
  43. return True
  44. class ManifestJsonVersionsProvider(VersionsProvider):
  45. name = "official-ardupilot"
  46. def __init__(self, manifest_json: ManifestJSON):
  47. self._manifest_json = manifest_json
  48. self._remote = RemoteInfo(OFFICIAL_REMOTE_NAME, OFFICIAL_REMOTE_URL)
  49. @property
  50. def is_available(self) -> bool:
  51. return self._manifest_json.is_available
  52. def refresh(self) -> None:
  53. self._manifest_json.refresh()
  54. def get_remotes(self) -> list[RemoteInfo]:
  55. if not self.is_available:
  56. return []
  57. return [self._remote]
  58. def get_versions(self, vehicle_id: str) -> list[VersionInfo]:
  59. if not self.is_available:
  60. return []
  61. return [
  62. VersionInfo.from_release_record(self._remote, release)
  63. for release in self._manifest_json.get_releases(vehicle_id)
  64. ]
  65. class WhitelistedForkTagVersionsProvider(VersionsProvider):
  66. """
  67. Discover custom-build tags from whitelisted GitHub forks.
  68. Only tags with the ``custom-build/`` prefix are included. Vehicle-specific
  69. tag segments may use a vehicle name (e.g. ``Copter``) or id (e.g. ``copter``).
  70. Tag format examples:
  71. - ``custom-build/my-feature`` — listed for all vehicles
  72. - ``custom-build/Copter/my-feature`` — listed for Copter only
  73. - ``custom-build/copter/my-feature`` — listed for Copter only
  74. """
  75. name = "whitelisted-fork-tags"
  76. def __init__(
  77. self,
  78. fork_remotes: Optional[list[ForkRemoteSpec]] = None,
  79. ):
  80. self.logger = logging.getLogger(__name__)
  81. self._fork_remotes = fork_remotes or DEFAULT_WHITELISTED_FORK_REMOTES
  82. self._vehicle_ids = [v.id for v in DEFAULT_VEHICLES]
  83. self._versions_by_remote_vehicle: dict[str, dict[str, list[ReleaseRecord]]] = {}
  84. self._available = False
  85. @property
  86. def is_available(self) -> bool:
  87. return self._available
  88. def get_remotes(self) -> list[RemoteInfo]:
  89. if not self.is_available:
  90. return []
  91. fork_by_owner = {spec.owner: spec for spec in self._fork_remotes}
  92. return [
  93. RemoteInfo(name=remote_name, url=fork_by_owner[remote_name].url)
  94. for remote_name in self._versions_by_remote_vehicle
  95. if remote_name != OFFICIAL_REMOTE_NAME and remote_name in fork_by_owner
  96. ]
  97. def refresh(self) -> None:
  98. versions_map = {
  99. spec.owner: {vehicle_id: [] for vehicle_id in self._vehicle_ids}
  100. for spec in self._fork_remotes
  101. }
  102. for spec in self._fork_remotes:
  103. if spec.owner == OFFICIAL_REMOTE_NAME:
  104. continue
  105. try:
  106. tag_objs = self._fetch_tags_from_github(spec.github_repo)
  107. except Exception as exc:
  108. self.logger.warning(
  109. "Skipping remote %s (%s): %s",
  110. spec.owner,
  111. spec.github_repo,
  112. exc,
  113. )
  114. continue
  115. for tag_info in tag_objs:
  116. ref = tag_info["ref"].replace("refs/tags/", "")
  117. parts = ref.split("/", 3)
  118. if len(parts) <= 1 or parts[0] != "custom-build":
  119. continue
  120. tag_vehicle_id = _vehicle_id_for_tag_segment(parts[1])
  121. if tag_vehicle_id is not None:
  122. if len(parts) == 2:
  123. continue
  124. vehicles_for_tag = [tag_vehicle_id]
  125. else:
  126. vehicles_for_tag = self._vehicle_ids
  127. for vid in vehicles_for_tag:
  128. versions_map[spec.owner][vid].append(
  129. ReleaseRecord(
  130. vehicle_id=vid,
  131. release_type="tag",
  132. version_number=parts[-1],
  133. commit_reference=tag_info["object"]["sha"],
  134. )
  135. )
  136. self._versions_by_remote_vehicle = versions_map
  137. self._available = True
  138. def get_versions(self, vehicle_id: str) -> list[VersionInfo]:
  139. if not self.is_available:
  140. return []
  141. fork_by_owner = {spec.owner: spec for spec in self._fork_remotes}
  142. versions = []
  143. for remote_name, vehicles_map in self._versions_by_remote_vehicle.items():
  144. if remote_name == OFFICIAL_REMOTE_NAME:
  145. continue
  146. spec = fork_by_owner.get(remote_name)
  147. if spec is None:
  148. continue
  149. remote = RemoteInfo(name=spec.owner, url=spec.url)
  150. for release in vehicles_map.get(vehicle_id, []):
  151. versions.append(VersionInfo.from_release_record(remote, release))
  152. return versions
  153. def _fetch_tags_from_github(self, github_repo: str) -> list:
  154. url = f"https://api.github.com/repos/{github_repo}/git/refs/tags"
  155. headers = {
  156. "X-GitHub-Api-Version": "2022-11-28",
  157. "Accept": "application/vnd.github+json",
  158. }
  159. token = os.getenv("CBS_GITHUB_ACCESS_TOKEN")
  160. if token:
  161. headers["Authorization"] = f"Bearer {token}"
  162. response = requests.get(url=url, headers=headers, timeout=60)
  163. response.raise_for_status()
  164. return response.json()
  165. class RemotesJsonVersionsProvider(VersionsProvider):
  166. name = "remotes-json"
  167. def __init__(self, remotes_json_path: str, schema_path: str):
  168. self._remotes_json_path = remotes_json_path
  169. self._schema_path = schema_path
  170. self._lock = Lock()
  171. self._metadata: list = []
  172. @property
  173. def is_available(self) -> bool:
  174. return bool(self._metadata)
  175. def refresh(self) -> None:
  176. self.reload()
  177. def reload(self) -> None:
  178. path = Path(self._remotes_json_path)
  179. if not path.is_file():
  180. with self._lock:
  181. self._metadata = []
  182. return
  183. content = path.read_text(encoding="utf-8")
  184. if not content.strip():
  185. with self._lock:
  186. self._metadata = []
  187. return
  188. metadata = json.loads(content)
  189. schema = json.loads(Path(self._schema_path).read_text(encoding="utf-8"))
  190. jsonschema.validate(instance=metadata, schema=schema)
  191. with self._lock:
  192. self._metadata = metadata
  193. def get_remotes(self) -> list[RemoteInfo]:
  194. with self._lock:
  195. metadata = list(self._metadata)
  196. return [
  197. RemoteInfo(name=remote.get("name"), url=remote.get("url"))
  198. for remote in metadata
  199. if remote.get("name") and remote.get("url")
  200. ]
  201. def get_versions(self, vehicle_id: str) -> list[VersionInfo]:
  202. versions = []
  203. with self._lock:
  204. metadata = list(self._metadata)
  205. for remote in metadata:
  206. remote_info = RemoteInfo(
  207. name=remote.get("name"),
  208. url=remote.get("url"),
  209. )
  210. for remote_vehicle in remote.get("vehicles", []):
  211. if remote_vehicle.get("id") != vehicle_id:
  212. continue
  213. for release in remote_vehicle.get("releases", []):
  214. versions.append(
  215. VersionInfo(
  216. remote_info=remote_info,
  217. commit_ref=release.get("commit_reference"),
  218. release_type=release.get("release_type"),
  219. version_number=release.get("version_number"),
  220. )
  221. )
  222. return versions
  223. def build_default_providers(
  224. manifest_json: ManifestJSON,
  225. remotes_json_path: str,
  226. schema_path: str,
  227. fork_remotes: Optional[list[ForkRemoteSpec]] = None,
  228. ) -> list[VersionsProvider]:
  229. return [
  230. ManifestJsonVersionsProvider(manifest_json),
  231. WhitelistedForkTagVersionsProvider(fork_remotes=fork_remotes),
  232. RemotesJsonVersionsProvider(
  233. remotes_json_path=remotes_json_path,
  234. schema_path=schema_path,
  235. ),
  236. ]