providers.py 10 KB

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