rivian-python-client 0.2.4__tar.gz → 1.0.1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: rivian-python-client
3
- Version: 0.2.4
3
+ Version: 1.0.1
4
4
  Summary: Rivian API Client (Unofficial)
5
5
  License: MIT
6
6
  Author: Brian Retterer
@@ -12,7 +12,8 @@ Classifier: Programming Language :: Python :: 3.9
12
12
  Classifier: Programming Language :: Python :: 3.10
13
13
  Classifier: Programming Language :: Python :: 3.11
14
14
  Requires-Dist: aiohttp (>=3.0.0)
15
- Requires-Dist: yarl (>=1.6.0)
15
+ Requires-Dist: backports-strenum (>=1.2.4,<2.0.0) ; python_version < "3.11"
16
+ Requires-Dist: cryptography (>=41.0.1,<42.0.0)
16
17
  Description-Content-Type: text/markdown
17
18
 
18
19
  # Python: Rivian API Client
@@ -1,6 +1,6 @@
1
1
  [tool.poetry]
2
2
  name = "rivian-python-client"
3
- version = "0.2.4"
3
+ version = "1.0.1"
4
4
  description = "Rivian API Client (Unofficial)"
5
5
  readme = "README.md"
6
6
  authors = ["Brian Retterer <bretterer@gmail.com>"]
@@ -12,17 +12,21 @@ packages = [
12
12
  [tool.poetry.dependencies]
13
13
  python = "^3.9"
14
14
  aiohttp = ">=3.0.0"
15
- yarl = ">=1.6.0"
15
+ cryptography = "^41.0.1"
16
+ backports-strenum = { version = "^1.2.4", python = "<3.11" }
16
17
 
17
- [tool.poetry.dev-dependencies]
18
+ [tool.poetry.group.dev.dependencies]
18
19
  pytest = "^7.1.2"
19
20
  pytest-asyncio = "^0.18.3"
20
21
  python-dotenv = "^0.20.0"
21
22
  aresponses = "^2.1.5"
22
23
 
24
+ [tool.poetry-dynamic-versioning]
25
+ enable = true
26
+ vcs = "git"
27
+ style = "semver"
28
+ pattern = "default-unprefixed"
23
29
 
24
30
  [build-system]
25
- requires = ["poetry-core>=1.0.0"]
26
- build-backend = "poetry.core.masonry.api"
27
-
28
-
31
+ requires = ["poetry-core>=1.0.0", "poetry-dynamic-versioning"]
32
+ build-backend = "poetry_dynamic_versioning.backend"
@@ -1,5 +1,6 @@
1
1
  """Asynchronous Python client for the Rivian API."""
2
2
 
3
+ from .const import VehicleCommand
3
4
  from .rivian import Rivian
4
5
 
5
- __all__ = ["Rivian"]
6
+ __all__ = ["Rivian", "VehicleCommand"]
@@ -1,8 +1,14 @@
1
1
  """Rivian constants."""
2
2
  from __future__ import annotations
3
3
 
4
+ import sys
4
5
  from typing import Final
5
6
 
7
+ if sys.version_info >= (3, 11):
8
+ from enum import StrEnum
9
+ else:
10
+ from backports.strenum import StrEnum
11
+
6
12
  LIVE_SESSION_PROPERTIES: Final[set[str]] = {
7
13
  "chargerId",
8
14
  "currentCurrency",
@@ -120,3 +126,73 @@ VEHICLE_STATE_PROPERTIES: Final[set[str]] = {
120
126
  "windowRearRightClosed",
121
127
  "wiperFluidState",
122
128
  }
129
+
130
+
131
+ class VehicleCommand(StrEnum):
132
+ """Supported vehicle commands."""
133
+
134
+ WAKE_VEHICLE = "WAKE_VEHICLE"
135
+ HONK_AND_FLASH_LIGHTS = "HONK_AND_FLASH_LIGHTS"
136
+ UNLOCK_USER_PREFERENCES_AND_DISABLE_ALARM = (
137
+ "UNLOCK_USER_PREFERENCES_AND_DISABLE_ALARM"
138
+ )
139
+
140
+ # Charging
141
+ CHARGING_LIMITS = "CHARGING_LIMITS"
142
+ START_CHARGING = "START_CHARGING"
143
+ STOP_CHARGING = "STOP_CHARGING"
144
+
145
+ # Climate
146
+ CABIN_HVAC_DEFROST_DEFOG = "CABIN_HVAC_DEFROST_DEFOG"
147
+ CABIN_HVAC_LEFT_SEAT_HEAT = "CABIN_HVAC_LEFT_SEAT_HEAT"
148
+ CABIN_HVAC_LEFT_SEAT_VENT = "CABIN_HVAC_LEFT_SEAT_VENT"
149
+ CABIN_HVAC_REAR_LEFT_SEAT_HEAT = "CABIN_HVAC_REAR_LEFT_SEAT_HEAT"
150
+ CABIN_HVAC_REAR_RIGHT_SEAT_HEAT = "CABIN_HVAC_REAR_RIGHT_SEAT_HEAT"
151
+ CABIN_HVAC_RIGHT_SEAT_HEAT = "CABIN_HVAC_RIGHT_SEAT_HEAT"
152
+ CABIN_HVAC_RIGHT_SEAT_VENT = "CABIN_HVAC_RIGHT_SEAT_VENT"
153
+ CABIN_HVAC_STEERING_HEAT = "CABIN_HVAC_STEERING_HEAT"
154
+ CABIN_PRECONDITIONING_SET_TEMP = "CABIN_PRECONDITIONING_SET_TEMP"
155
+ VEHICLE_CABIN_PRECONDITION_DISABLE = "VEHICLE_CABIN_PRECONDITION_DISABLE"
156
+ VEHICLE_CABIN_PRECONDITION_ENABLE = "VEHICLE_CABIN_PRECONDITION_ENABLE"
157
+
158
+ # Closures
159
+ LOCK_ALL_CLOSURES_FEEDBACK = "LOCK_ALL_CLOSURES_FEEDBACK"
160
+ UNLOCK_ALL_CLOSURES = "UNLOCK_ALL_CLOSURES"
161
+ UNLOCK_DRIVER_DOOR = "UNLOCK_DRIVER_DOOR"
162
+ UNLOCK_PASSENGER_DOOR = "UNLOCK_PASSENGER_DOOR"
163
+
164
+ # Frunk
165
+ CLOSE_FRUNK = "CLOSE_FRUNK"
166
+ OPEN_FRUNK = "OPEN_FRUNK"
167
+
168
+ # Gear guard
169
+ ENABLE_GEAR_GUARD = "ENABLE_GEAR_GUARD"
170
+ ENABLE_GEAR_GUARD_VIDEO = "ENABLE_GEAR_GUARD_VIDEO"
171
+ DISABLE_GEAR_GUARD = "DISABLE_GEAR_GUARD"
172
+ DISABLE_GEAR_GUARD_VIDEO = "DISABLE_GEAR_GUARD_VIDEO"
173
+
174
+ # Liftgate (R1S only)
175
+ CLOSE_LIFTGATE = "CLOSE_LIFTGATE"
176
+
177
+ # Liftgate/tailgate
178
+ OPEN_LIFTGATE_UNLATCH_TAILGATE = "OPEN_LIFTGATE_UNLATCH_TAILGATE"
179
+
180
+ # OTA
181
+ OTA_INSTALL_NOW_ACKNOWLEDGE = "OTA_INSTALL_NOW_ACKNOWLEDGE"
182
+
183
+ # Panic
184
+ PANIC_OFF = "PANIC_OFF"
185
+ PANIC_ON = "PANIC_ON"
186
+
187
+ # Side bin (R1T only)
188
+ RELEASE_LEFT_SIDE_BIN = "RELEASE_LEFT_SIDE_BIN"
189
+ RELEASE_RIGHT_SIDE_BIN = "RELEASE_RIGHT_SIDE_BIN"
190
+
191
+ # Tonneau (Only for R1T with powered tonneau)
192
+ CLOSE_TONNEAU_COVER = "CLOSE_TONNEAU_COVER"
193
+ OPEN_TONNEAU_COVER = "OPEN_TONNEAU_COVER"
194
+
195
+ # Windows
196
+ CLOSE_ALL_WINDOWS = "CLOSE_ALL_WINDOWS"
197
+ OPEN_ALL_WINDOWS = "OPEN_ALL_WINDOWS"
198
+ UNLOCK_ALL_AND_OPEN_WINDOWS = "UNLOCK_ALL_AND_OPEN_WINDOWS"
@@ -31,3 +31,11 @@ class RivianTemporarilyLockedError(RivianApiException):
31
31
 
32
32
  class RivianApiRateLimitError(RivianApiException):
33
33
  """Rivian API is being rate limited."""
34
+
35
+
36
+ class RivianPhoneLimitReachedError(RivianApiException):
37
+ """Rivian phone limit has been reached."""
38
+
39
+
40
+ class RivianBadRequestError(RivianApiException):
41
+ """Rivian API bad request."""
File without changes
@@ -2,34 +2,35 @@
2
2
  from __future__ import annotations
3
3
 
4
4
  import asyncio
5
- import json
6
5
  import logging
7
6
  import socket
7
+ import time
8
8
  import uuid
9
9
  from collections.abc import Callable
10
10
  from typing import Any, Type
11
+ from warnings import warn
11
12
 
12
13
  import aiohttp
13
14
  import async_timeout
14
- from aiohttp import ClientRequest, ClientResponse, ClientWebSocketResponse
15
+ from aiohttp import ClientResponse, ClientWebSocketResponse
15
16
 
16
- from .const import LIVE_SESSION_PROPERTIES, VEHICLE_STATE_PROPERTIES
17
+ from .const import LIVE_SESSION_PROPERTIES, VEHICLE_STATE_PROPERTIES, VehicleCommand
17
18
  from .exceptions import (
18
19
  RivianApiException,
19
20
  RivianApiRateLimitError,
21
+ RivianBadRequestError,
20
22
  RivianDataError,
21
- RivianExpiredTokenError,
22
23
  RivianInvalidCredentials,
23
24
  RivianInvalidOTP,
25
+ RivianPhoneLimitReachedError,
24
26
  RivianTemporarilyLockedError,
25
27
  RivianUnauthenticated,
26
28
  )
29
+ from .utils import generate_vehicle_command_hmac
27
30
  from .ws_monitor import WebSocketMonitor
28
31
 
29
32
  _LOGGER = logging.getLogger(__name__)
30
33
 
31
- CESIUM_BASEPATH = "https://cesium.rivianservices.com/v2"
32
- AUTH_BASEPATH = "https://auth.rivianservices.com/auth/api/v1"
33
34
  GRAPHQL_BASEPATH = "https://rivian.com/api/gql"
34
35
  GRAPHQL_GATEWAY = GRAPHQL_BASEPATH + "/gateway/graphql"
35
36
  GRAPHQL_CHARGING = GRAPHQL_BASEPATH + "/chrg/user/graphql"
@@ -60,6 +61,7 @@ VALUE_RECORD_TEMPLATE = "{ __typename value updatedAt }"
60
61
 
61
62
  ERROR_CODE_CLASS_MAP: dict[str, Type[RivianApiException]] = {
62
63
  "BAD_CURRENT_PASSWORD": RivianInvalidCredentials,
64
+ "BAD_REQUEST_ERROR": RivianBadRequestError,
63
65
  "DATA_ERROR": RivianDataError,
64
66
  "INTERNAL_SERVER_ERROR": RivianApiException,
65
67
  "RATE_LIMIT": RivianApiRateLimitError,
@@ -68,27 +70,40 @@ ERROR_CODE_CLASS_MAP: dict[str, Type[RivianApiException]] = {
68
70
  }
69
71
 
70
72
 
73
+ def send_deprecation_warning(old_name: str, new_name: str) -> None: # pragma: no cover
74
+ """Send a deprecation warning."""
75
+ message = f"{old_name} has been deprecated in favor of {new_name}, the alias will be removed in the future"
76
+ warn(
77
+ message,
78
+ DeprecationWarning,
79
+ stacklevel=2,
80
+ )
81
+ _LOGGER.warning(message)
82
+
83
+
71
84
  class Rivian:
72
85
  """Main class for the Rivian API Client"""
73
86
 
74
87
  def __init__(
75
88
  self,
76
- client_id: str,
77
- client_secret: str,
78
89
  request_timeout: int = 10,
79
90
  session: aiohttp.client.ClientSession | None = None,
91
+ *,
92
+ access_token: str = "",
93
+ refresh_token: str = "",
94
+ csrf_token: str = "",
95
+ app_session_token: str = "",
96
+ user_session_token: str = "",
80
97
  ) -> None:
81
98
  self._session = session
82
99
  self._close_session = False
83
- self._session_token = ""
84
- self._access_token = ""
85
- self._refresh_token = ""
86
- self._csrf_token = ""
87
- self._app_session_token = ""
88
- self._user_session_token = ""
89
-
90
- self.client_id = client_id
91
- self.client_secret = client_secret
100
+
101
+ self._access_token = access_token
102
+ self._refresh_token = refresh_token
103
+ self._csrf_token = csrf_token
104
+ self._app_session_token = app_session_token
105
+ self._user_session_token = user_session_token
106
+
92
107
  self.request_timeout = request_timeout
93
108
 
94
109
  self._otp_needed = False
@@ -97,242 +112,7 @@ class Rivian:
97
112
  self._ws_monitor: WebSocketMonitor | None = None
98
113
  self._subscriptions: dict[str, str] = {}
99
114
 
100
- async def authenticate(
101
- self,
102
- username: str,
103
- password: str,
104
- ) -> ClientResponse | dict[str, str]:
105
- """Authenticate against the Rivian API with Username and Password"""
106
- url = AUTH_BASEPATH + "/token/auth"
107
-
108
- headers = {**BASE_HEADERS}
109
-
110
- json_data = {
111
- "grant_type": "password",
112
- "username": username,
113
- "client_id": self.client_id,
114
- "client_secret": self.client_secret,
115
- "pwd": password,
116
- }
117
-
118
- if self._session is None:
119
- self._session = aiohttp.ClientSession()
120
- self._close_session = True
121
-
122
- try:
123
- async with async_timeout.timeout(self.request_timeout):
124
- response = await self._session.request(
125
- "POST",
126
- url,
127
- json=json_data,
128
- headers=headers,
129
- )
130
- except asyncio.TimeoutError as exception:
131
- raise Exception(
132
- "Timeout occurred while connecting to Rivian API."
133
- ) from exception
134
- except (aiohttp.ClientError, socket.gaierror) as exception:
135
- raise Exception(
136
- "Error occurred while communicating with Rivian."
137
- ) from exception
138
-
139
- content_type = response.headers.get("Content-Type", "")
140
-
141
- if response.status == 401:
142
- self._otp_needed = True
143
- response_json = await response.json()
144
- self._session_token = response_json["session_token"]
145
- return response
146
-
147
- if response.status // 100 in [4, 5]:
148
- contents = await response.read()
149
- response.close()
150
-
151
- if content_type == "application/json":
152
- raise Exception(response.status, json.loads(contents.decode("utf8")))
153
- raise Exception(response.status, {"message": contents.decode("utf8")})
154
-
155
- if "application/json" in content_type:
156
- response_json = await response.json()
157
- self._access_token = response_json["access_token"]
158
- self._refresh_token = response_json["refresh_token"]
159
- return response
160
-
161
- text = await response.text()
162
- return {"message": text}
163
-
164
- async def validate_otp(
165
- self,
166
- username: str,
167
- otp: str,
168
- ) -> dict[str, Any]:
169
- """Validate the OTP"""
170
- url = AUTH_BASEPATH + "/token/auth"
171
-
172
- headers = BASE_HEADERS | {"Authorization": "Bearer " + self._session_token}
173
-
174
- json_data = {
175
- "grant_type": "password",
176
- "username": username,
177
- "client_id": self.client_id,
178
- "client_secret": self.client_secret,
179
- "otp_token": otp,
180
- }
181
-
182
- if self._session is None:
183
- self._session = aiohttp.ClientSession()
184
- self._close_session = True
185
-
186
- try:
187
- async with async_timeout.timeout(self.request_timeout):
188
- response = await self._session.request(
189
- "POST",
190
- url,
191
- json=json_data,
192
- headers=headers,
193
- )
194
- except asyncio.TimeoutError as exception:
195
- raise Exception(
196
- "Timeout occurred while connecting to Rivian API."
197
- ) from exception
198
- except (aiohttp.ClientError, socket.gaierror) as exception:
199
- raise Exception(
200
- "Error occurred while communicating with Rivian."
201
- ) from exception
202
-
203
- content_type = response.headers.get("Content-Type", "")
204
-
205
- if response.status // 100 in [4, 5]:
206
- contents = await response.read()
207
- response.close()
208
-
209
- if content_type == "application/json":
210
- raise Exception(response.status, json.loads(contents.decode("utf8")))
211
- raise Exception(response.status, {"message": contents.decode("utf8")})
212
-
213
- if "application/json" in content_type:
214
- response_json = await response.json()
215
- self._access_token = response_json["access_token"]
216
- self._refresh_token = response_json["refresh_token"]
217
- return response
218
-
219
- text = await response.text()
220
- return {"message": text}
221
-
222
- async def refresh_access_token(
223
- self,
224
- refresh_token: str,
225
- client_id: str,
226
- client_secret: str,
227
- ) -> ClientRequest:
228
- """Validate the OTP"""
229
- url = AUTH_BASEPATH + "/token/refresh"
230
-
231
- headers = {**BASE_HEADERS}
232
-
233
- json_data = {
234
- "token": refresh_token,
235
- "client_id": client_id,
236
- "client_secret": client_secret,
237
- }
238
-
239
- if self._session is None:
240
- self._session = aiohttp.ClientSession()
241
- self._close_session = True
242
-
243
- try:
244
- async with async_timeout.timeout(self.request_timeout):
245
- response = await self._session.request(
246
- "POST",
247
- url,
248
- json=json_data,
249
- headers=headers,
250
- )
251
- except asyncio.TimeoutError as exception:
252
- raise Exception(
253
- "Timeout occurred while connecting to Rivian API."
254
- ) from exception
255
- except (aiohttp.ClientError, socket.gaierror) as exception:
256
- raise Exception(
257
- "Error occurred while communicating with Rivian."
258
- ) from exception
259
-
260
- content_type = response.headers.get("Content-Type", "")
261
-
262
- if response.status // 100 in [4, 5]:
263
- contents = await response.read()
264
- response.close()
265
-
266
- if content_type == "application/json":
267
- raise Exception(
268
- response.status, json.loads(contents.decode("utf8")), json_data
269
- )
270
- raise Exception(response.status, {"message": contents.decode("utf8")})
271
-
272
- if "application/json" in content_type:
273
- response_json = await response.json()
274
- self._access_token = response_json["access_token"]
275
- return response
276
-
277
- return response
278
-
279
- async def get_vehicle_info(
280
- self, vin: str, access_token: str, properties: dict[str]
281
- ) -> dict[str, Any]:
282
- """get the vehicle info"""
283
- url = CESIUM_BASEPATH + "/vehicle/latest"
284
-
285
- headers = BASE_HEADERS | {"Authorization": "Bearer " + access_token}
286
-
287
- json_data = {
288
- "car": vin,
289
- "properties": properties,
290
- }
291
-
292
- if self._session is None:
293
- self._session = aiohttp.ClientSession()
294
- self._close_session = True
295
-
296
- try:
297
- async with async_timeout.timeout(self.request_timeout):
298
- response = await self._session.request(
299
- "POST",
300
- url,
301
- json=json_data,
302
- headers=headers,
303
- )
304
- except asyncio.TimeoutError as exception:
305
- raise Exception(
306
- "Timeout occurred while connecting to Rivian API."
307
- ) from exception
308
- except (aiohttp.ClientError, socket.gaierror) as exception:
309
- raise Exception(
310
- "Error occurred while communicating with Rivian."
311
- ) from exception
312
-
313
- if response.status // 100 in [4, 5]:
314
- contents = await response.read()
315
- response.close()
316
-
317
- response_json = await response.json()
318
- if response.status == 401 and response_json["error_code"] == -40:
319
- raise RivianExpiredTokenError(
320
- response.status,
321
- response_json,
322
- headers,
323
- json_data,
324
- )
325
-
326
- raise Exception(
327
- response.status,
328
- json.loads(contents.decode("utf8")),
329
- headers,
330
- json_data,
331
- )
332
-
333
- return response
334
-
335
- async def create_csrf_token(self) -> dict[str, Any]:
115
+ async def create_csrf_token(self) -> None:
336
116
  """Create cross-site-request-forgery (csrf) token."""
337
117
  url = GRAPHQL_GATEWAY
338
118
 
@@ -352,13 +132,7 @@ class Rivian:
352
132
  self._csrf_token = csrf_data["csrfToken"]
353
133
  self._app_session_token = csrf_data["appSessionToken"]
354
134
 
355
- return response
356
-
357
- async def authenticate_graphql(
358
- self,
359
- username: str,
360
- password: str,
361
- ) -> dict[str, Any]:
135
+ async def authenticate(self, username: str, password: str) -> None:
362
136
  """Authenticate against the Rivian GraphQL API with Username and Password"""
363
137
  url = GRAPHQL_GATEWAY
364
138
 
@@ -389,9 +163,17 @@ class Rivian:
389
163
  self._refresh_token = login_data["refreshToken"]
390
164
  self._user_session_token = login_data["userSessionToken"]
391
165
 
392
- return response
166
+ async def authenticate_graphql(
167
+ self, username: str, password: str
168
+ ) -> None: # pragma: no cover
169
+ """### DEPRECATED (use `authenticate` instead)
393
170
 
394
- async def validate_otp_graphql(self, username: str, otpCode: str) -> dict[str, Any]:
171
+ Authenticate against the Rivian GraphQL API with Username and Password.
172
+ """
173
+ send_deprecation_warning("authenticate_graphql", "authenticate")
174
+ return await self.authenticate(username=username, password=password)
175
+
176
+ async def validate_otp(self, username: str, otp_code: str) -> None:
395
177
  """Validates OTP against the Rivian GraphQL API with Username, OTP Code, and OTP Token"""
396
178
  url = GRAPHQL_GATEWAY
397
179
 
@@ -406,7 +188,7 @@ class Rivian:
406
188
  "query": "mutation LoginWithOTP($email: String!, $otpCode: String!, $otpToken: String!) {\n loginWithOTP(email: $email, otpCode: $otpCode, otpToken: $otpToken) {\n __typename\n ... on MobileLoginResponse {\n __typename\n accessToken\n refreshToken\n userSessionToken\n }\n }\n}",
407
189
  "variables": {
408
190
  "email": username,
409
- "otpCode": otpCode,
191
+ "otpCode": otp_code,
410
192
  "otpToken": self._otp_token,
411
193
  },
412
194
  }
@@ -421,10 +203,79 @@ class Rivian:
421
203
  self._refresh_token = login_data["refreshToken"]
422
204
  self._user_session_token = login_data["userSessionToken"]
423
205
 
424
- return response
206
+ async def validate_otp_graphql(
207
+ self, username: str, otpCode: str
208
+ ) -> None: # pragma: no cover
209
+ """### DEPRECATED (use `validate_otp` instead)
210
+
211
+ Validates OTP against the Rivian GraphQL API with Username, OTP Code, and OTP Token.
212
+ """
213
+ send_deprecation_warning("validate_otp_graphql", "validate_otp")
214
+ return await self.validate_otp(username=username, otp_code=otpCode)
215
+
216
+ async def disenroll_phone(self, identity_id: str) -> bool:
217
+ """Disenroll a phone."""
218
+ url = GRAPHQL_GATEWAY
219
+ headers = BASE_HEADERS | {
220
+ "Csrf-Token": self._csrf_token,
221
+ "A-Sess": self._app_session_token,
222
+ "U-Sess": self._user_session_token,
223
+ }
224
+ graphql_json = {
225
+ "operationName": "DisenrollPhone",
226
+ "variables": {"attrs": {"enrollmentId": identity_id}},
227
+ "query": "mutation DisenrollPhone($attrs: DisenrollPhoneAttributes!) { disenrollPhone(attrs: $attrs) { __typename success } }",
228
+ }
229
+
230
+ response = await self.__graphql_query(headers, url, graphql_json)
231
+ if response.status == 200:
232
+ data = await response.json()
233
+ return data.get("data", {}).get("disenrollPhone", {}).get("success")
234
+ return False
425
235
 
426
- async def get_user_information(self) -> ClientResponse:
427
- """get user information (user.id, vehicle vins)"""
236
+ async def enroll_phone(
237
+ self,
238
+ user_id: str,
239
+ vehicle_id: str,
240
+ device_type: str,
241
+ device_name: str,
242
+ public_key: str,
243
+ ) -> bool:
244
+ """Enable control of a vehicle by enrolling a phone.
245
+
246
+ To generate a public/private key for enrollment, use the `utils.generate_key_pair` function.
247
+ The private key will need to be retained to sign commands sent via the `send_vehicle_command` method.
248
+ """
249
+ url = GRAPHQL_GATEWAY
250
+ headers = BASE_HEADERS | {
251
+ "Csrf-Token": self._csrf_token,
252
+ "A-Sess": self._app_session_token,
253
+ "U-Sess": self._user_session_token,
254
+ }
255
+ graphql_json = {
256
+ "operationName": "EnrollPhone",
257
+ "variables": {
258
+ "attrs": {
259
+ "userId": user_id,
260
+ "vehicleId": vehicle_id,
261
+ "publicKey": public_key,
262
+ "type": device_type,
263
+ "name": device_name,
264
+ }
265
+ },
266
+ "query": "mutation EnrollPhone($attrs: EnrollPhoneAttributes!) { enrollPhone(attrs: $attrs) { __typename success } }",
267
+ }
268
+ response = await self.__graphql_query(headers, url, graphql_json)
269
+ if response.status == 200:
270
+ data = await response.json()
271
+ if data.get("data", {}).get("enrollPhone", {}).get("success"):
272
+ return True
273
+ return False
274
+
275
+ async def get_user_information(
276
+ self, include_phones: bool = False
277
+ ) -> ClientResponse:
278
+ """Get user information."""
428
279
  url = GRAPHQL_GATEWAY
429
280
 
430
281
  headers = BASE_HEADERS | {
@@ -432,16 +283,19 @@ class Rivian:
432
283
  "U-Sess": self._user_session_token,
433
284
  }
434
285
 
286
+ vehicles_fragment = "vehicles { id vin name vas { __typename vasVehicleId vehiclePublicKey } roles state createdAt updatedAt vehicle { __typename id vin modelYear make model expectedBuildDate plannedBuildDate expectedGeneralAssemblyStartDate actualGeneralAssemblyDate } }"
287
+ phones_fragment = "enrolledPhones { __typename vas { __typename vasPhoneId publicKey } enrolled { __typename deviceType deviceName vehicleId identityId shortName } }"
288
+
435
289
  graphql_json = {
436
290
  "operationName": "getUserInfo",
437
- "query": "query getUserInfo {\n currentUser {\n __typename\n id\n vehicles {\n id\n vin\n name\n vas {\n __typename\n vasVehicleId\n vehiclePublicKey\n }\n roles\n state\n createdAt\n updatedAt\n vehicle {\n __typename\n id\n vin\n modelYear\n make\n model\n expectedBuildDate\n plannedBuildDate\n expectedGeneralAssemblyStartDate\n actualGeneralAssemblyDate\n }\n }\n }\n}",
291
+ "query": f"query getUserInfo {{ currentUser {{ __typename id {vehicles_fragment} {phones_fragment if include_phones else ''} }} }}",
438
292
  "variables": None,
439
293
  }
440
294
 
441
295
  return await self.__graphql_query(headers, url, graphql_json)
442
296
 
443
297
  async def get_registered_wallboxes(self) -> ClientResponse:
444
- """get wallboxes (graphql)"""
298
+ """Get registered wallboxes."""
445
299
  url = GRAPHQL_CHARGING
446
300
 
447
301
  headers = BASE_HEADERS | {
@@ -458,8 +312,69 @@ class Rivian:
458
312
 
459
313
  return await self.__graphql_query(headers, url, graphql_json)
460
314
 
461
- async def get_vehicle_state(self, vin: str, properties: set[str]) -> ClientResponse:
462
- """get vehicle state (graphql)"""
315
+ async def get_vehicle_command_state(self, command_id: str) -> ClientResponse:
316
+ """Get vehicle command state."""
317
+ url = GRAPHQL_GATEWAY
318
+
319
+ headers = BASE_HEADERS | {
320
+ "A-Sess": self._app_session_token,
321
+ "U-Sess": self._user_session_token,
322
+ }
323
+
324
+ graphql_query = "query getVehicleCommand($id: String!) { getVehicleCommand(id: $id) { __typename id command createdAt state responseCode statusCode } }"
325
+
326
+ graphql_json = {
327
+ "operationName": "getVehicleCommand",
328
+ "query": graphql_query,
329
+ "variables": {"id": command_id},
330
+ }
331
+
332
+ return await self.__graphql_query(headers, url, graphql_json)
333
+
334
+ async def get_vehicle_images(
335
+ self,
336
+ *,
337
+ extension: str | None = None,
338
+ resolution: str | None = None,
339
+ vehicle_version: str | None = None,
340
+ preorder_version: str | None = None,
341
+ ) -> ClientResponse:
342
+ """Get vehicle images.
343
+
344
+ Known parameter values:
345
+ - extension: `png`, `webp`
346
+ - resolution: `@1x`, `@2x`, `@3x` (for png); `hdpi`, `xhdpi`, `xxhdpi`, `xxxhdpi` (for webp)
347
+ - vehicle_version/preorder_version: `1`, `2` (all other values return v1 images)
348
+ """
349
+ url = GRAPHQL_GATEWAY
350
+
351
+ headers = BASE_HEADERS | {
352
+ "A-Sess": self._app_session_token,
353
+ "U-Sess": self._user_session_token,
354
+ }
355
+
356
+ graphql_query = "query getVehicleImages( $extension: String $resolution: String $versionForVehicle: String $versionForPreOrder: String ) { getVehicleOrderMobileImages( resolution: $resolution extension: $extension version: $versionForPreOrder ) { ...image } getVehicleMobileImages( resolution: $resolution extension: $extension version: $versionForVehicle ) { ...image } } fragment image on VehicleMobileImage { orderId vehicleId url extension resolution size design placement }"
357
+
358
+ graphql_json = {
359
+ "operationName": "getVehicleImages",
360
+ "query": graphql_query,
361
+ "variables": {
362
+ "extension": extension,
363
+ "resolution": resolution,
364
+ "versionForVehicle": vehicle_version,
365
+ "versionForPreOrder": preorder_version,
366
+ },
367
+ }
368
+
369
+ return await self.__graphql_query(headers, url, graphql_json)
370
+
371
+ async def get_vehicle_state(
372
+ self, vin: str, properties: set[str] | None = None
373
+ ) -> ClientResponse:
374
+ """Get vehicle state."""
375
+ if not properties:
376
+ properties = VEHICLE_STATE_PROPERTIES
377
+
463
378
  url = GRAPHQL_GATEWAY
464
379
 
465
380
  headers = BASE_HEADERS | {
@@ -509,11 +424,112 @@ class Rivian:
509
424
 
510
425
  return await self.__graphql_query(headers, url, graphql_json)
511
426
 
427
+ def _validate_vehicle_command(
428
+ self, command: VehicleCommand | str, params: dict[str, Any] | None = None
429
+ ) -> None:
430
+ """Validate certian vehicle command/param combos."""
431
+ if command == VehicleCommand.CHARGING_LIMITS:
432
+ if not (
433
+ params
434
+ and isinstance((limit := params.get("SOC_limit")), int)
435
+ and 50 <= limit <= 100
436
+ ):
437
+ raise RivianBadRequestError(
438
+ "Charging limit must include parameter `SOC_limit` with a valid value between 50 and 100"
439
+ )
440
+ if command in (
441
+ VehicleCommand.CABIN_HVAC_DEFROST_DEFOG,
442
+ VehicleCommand.CABIN_HVAC_LEFT_SEAT_HEAT,
443
+ VehicleCommand.CABIN_HVAC_LEFT_SEAT_VENT,
444
+ VehicleCommand.CABIN_HVAC_REAR_LEFT_SEAT_HEAT,
445
+ VehicleCommand.CABIN_HVAC_REAR_RIGHT_SEAT_HEAT,
446
+ VehicleCommand.CABIN_HVAC_RIGHT_SEAT_HEAT,
447
+ VehicleCommand.CABIN_HVAC_RIGHT_SEAT_VENT,
448
+ VehicleCommand.CABIN_HVAC_STEERING_HEAT,
449
+ ):
450
+ if not (
451
+ params
452
+ and isinstance((level := params.get("level")), int)
453
+ and 0 <= level <= 4
454
+ ):
455
+ raise RivianBadRequestError(
456
+ "HVAC setting must include parameter `level` with a valid value between 0 and 4"
457
+ )
458
+ if command == VehicleCommand.CABIN_PRECONDITIONING_SET_TEMP:
459
+ if not (
460
+ params
461
+ and isinstance((temp := params.get("HVAC_set_temp")), (float, int))
462
+ and (16 <= temp <= 29 or temp in (0, 63.5))
463
+ ):
464
+ raise RivianBadRequestError(
465
+ "HVAC setting must include parameter `HVAC_set_temp` with a valid value between 16 and 29 or 0/63.5 for LO/HI, respectively"
466
+ )
467
+ params["HVAC_set_temp"] = str(params["HVAC_set_temp"])
468
+
469
+ async def send_vehicle_command(
470
+ self,
471
+ command: VehicleCommand | str,
472
+ vehicle_id: str,
473
+ phone_id: str,
474
+ identity_id: str,
475
+ vehicle_key: str,
476
+ private_key: str,
477
+ *,
478
+ params: dict[str, Any] | None = None,
479
+ ) -> str | None:
480
+ """Send a command to the vehicle.
481
+
482
+ To generate a public/private key for commands, use the `utils.generate_key_pair` function.
483
+ The public key will first need to be enrolled via the `enroll_phone` method, otherwise commands will fail.
484
+
485
+ Certain commands may require additional details via the `params` mapping.
486
+ Some known examples include:
487
+ - `CABIN_HVAC_*`: params = {"level": 0..4} where 0 is off, 1 is on, 2 is low/level_1, 3 is medium/level_2 and 4 is high/level_3
488
+ - `CABIN_PRECONDITIONING_SET_TEMP`: params = {"HVAC_set_temp": "deg_C"} where `deg_C` is a string value between 16 and 29 or 0/63.5 for LO/HI, respectively
489
+ - `CHARGING_LIMITS`: params = {"SOC_limit": 50..100}
490
+ """
491
+ self._validate_vehicle_command(command, params)
492
+
493
+ command = str(command)
494
+ timestamp = str(int(time.time()))
495
+ hmac = generate_vehicle_command_hmac(
496
+ command, timestamp, vehicle_key, private_key
497
+ )
498
+
499
+ url = GRAPHQL_GATEWAY
500
+ headers = BASE_HEADERS | {
501
+ "Csrf-Token": self._csrf_token,
502
+ "A-Sess": self._app_session_token,
503
+ "U-Sess": self._user_session_token,
504
+ }
505
+ graphql_json = {
506
+ "operationName": "sendVehicleCommand",
507
+ "variables": {
508
+ "attrs": {
509
+ "command": command,
510
+ "hmac": hmac,
511
+ "timestamp": str(timestamp),
512
+ "vasPhoneId": phone_id,
513
+ "deviceId": identity_id,
514
+ "vehicleId": vehicle_id,
515
+ }
516
+ | ({"params": params} if params else {})
517
+ },
518
+ "query": "mutation sendVehicleCommand($attrs: VehicleCommandAttributes!) { sendVehicleCommand(attrs: $attrs) { __typename id command state } }",
519
+ }
520
+
521
+ response = await self.__graphql_query(headers, url, graphql_json)
522
+ if response.status == 200:
523
+ data = await response.json()
524
+ if status := data.get("data", {}).get("sendVehicleCommand", {}):
525
+ return status.get("id")
526
+ return None
527
+
512
528
  async def subscribe_for_vehicle_updates(
513
529
  self,
514
530
  vehicle_id: str,
531
+ callback: Callable[[dict[str, Any]], None],
515
532
  properties: set[str] | None = None,
516
- callback: Callable = None,
517
533
  ) -> Callable | None:
518
534
  """Open a web socket connection to receive updates."""
519
535
  if not properties:
@@ -521,6 +537,7 @@ class Rivian:
521
537
 
522
538
  try:
523
539
  await self._ws_connect()
540
+ assert self._ws_monitor
524
541
  async with async_timeout.timeout(self.request_timeout):
525
542
  await self._ws_monitor.connection_ack.wait()
526
543
  payload = {
@@ -533,6 +550,7 @@ class Rivian:
533
550
  return unsubscribe
534
551
  except Exception as ex: # pylint: disable=broad-except
535
552
  _LOGGER.error(ex)
553
+ return None
536
554
 
537
555
  async def _ws_connect(self) -> ClientWebSocketResponse:
538
556
  """Initiate a websocket connection."""
@@ -557,12 +575,15 @@ class Rivian:
557
575
  ws_monitor = self._ws_monitor
558
576
  if ws_monitor.websocket is None or ws_monitor.websocket.closed:
559
577
  await ws_monitor.new_connection(True)
578
+ assert ws_monitor.websocket
560
579
  if ws_monitor.monitor is None or ws_monitor.monitor.done():
561
580
  await ws_monitor.start_monitor()
562
581
  return ws_monitor.websocket
563
582
 
564
- async def __graphql_query(self, headers: dict(str, str), url: str, body: str):
565
- """execute and return arbitrary graphql query"""
583
+ async def __graphql_query(
584
+ self, headers: dict[str, str], url: str, body: dict[str, Any]
585
+ ) -> ClientResponse:
586
+ """Execute and return arbitrary graphql query."""
566
587
  if self._session is None:
567
588
  self._session = aiohttp.ClientSession()
568
589
  self._close_session = True
@@ -590,14 +611,22 @@ class Rivian:
590
611
  for error in errors:
591
612
  if extensions := error.get("extensions"):
592
613
  code = extensions["code"]
593
- if err_cls := ERROR_CODE_CLASS_MAP.get(code):
594
- raise err_cls(response.status, response_json, headers, body)
595
- if code == "BAD_USER_INPUT" and (
596
- extensions["reason"] == "INVALID_OTP"
614
+ if (code, extensions.get("reason")) in (
615
+ ("BAD_USER_INPUT", "INVALID_OTP"),
616
+ ("UNAUTHENTICATED", "OTP_TOKEN_EXPIRED"),
597
617
  ):
598
618
  raise RivianInvalidOTP(
599
619
  response.status, response_json, headers, body
600
620
  )
621
+ if (code, extensions.get("reason")) == (
622
+ "CONFLICT",
623
+ "ENROLL_PHONE_LIMIT_REACHED",
624
+ ):
625
+ raise RivianPhoneLimitReachedError(
626
+ response.status, response_json, headers, body
627
+ )
628
+ if err_cls := ERROR_CODE_CLASS_MAP.get(code):
629
+ raise err_cls(response.status, response_json, headers, body)
601
630
  raise RivianApiException(
602
631
  "Error occurred while reading the graphql response from Rivian.",
603
632
  response.status,
@@ -0,0 +1,90 @@
1
+ """Utilities."""
2
+ from __future__ import annotations
3
+
4
+ import hashlib
5
+ import hmac
6
+ from base64 import b64decode, b64encode
7
+ from typing import cast
8
+
9
+ from cryptography.hazmat.primitives import hashes, serialization
10
+ from cryptography.hazmat.primitives.asymmetric import ec
11
+ from cryptography.hazmat.primitives.kdf.hkdf import HKDF
12
+
13
+
14
+ def base64_encode(data: bytes) -> str:
15
+ """Encode bytes to Base64 string"""
16
+ return b64encode(data).decode("utf-8")
17
+
18
+
19
+ def decode_private_key(private_key_str: str) -> ec.EllipticCurvePrivateKey:
20
+ """Decode an EC private key."""
21
+ key = serialization.load_pem_private_key(b64decode(private_key_str), password=None)
22
+ return cast(ec.EllipticCurvePrivateKey, key)
23
+
24
+
25
+ def decode_public_key(public_key_str) -> ec.EllipticCurvePublicKey:
26
+ """Decode an EC public key."""
27
+ return ec.EllipticCurvePublicKey.from_encoded_point(
28
+ ec.SECP256R1(), bytes.fromhex(public_key_str)
29
+ )
30
+
31
+
32
+ def encode_private_key(private_key: ec.EllipticCurvePrivateKey) -> str:
33
+ """Encode an EC public key."""
34
+ return base64_encode(
35
+ private_key.private_bytes(
36
+ encoding=serialization.Encoding.PEM,
37
+ format=serialization.PrivateFormat.PKCS8,
38
+ encryption_algorithm=serialization.NoEncryption(),
39
+ )
40
+ )
41
+
42
+
43
+ def encode_public_key(public_key: ec.EllipticCurvePublicKey) -> str:
44
+ """Encode an EC public key."""
45
+ return public_key.public_bytes(
46
+ encoding=serialization.Encoding.X962,
47
+ format=serialization.PublicFormat.UncompressedPoint,
48
+ ).hex()
49
+
50
+
51
+ def generate_key_pair() -> tuple[str, str]:
52
+ """Generate an ECDH public-private key pair.
53
+
54
+ Copied from https://rivian-api.kaedenb.org/app/controls/enroll-phone/
55
+ """
56
+ # Generate a private key
57
+ private_key = ec.generate_private_key(ec.SECP256R1())
58
+
59
+ # Get the corresponding public key
60
+ public_key = private_key.public_key()
61
+
62
+ # Serialize the keys in the standard format
63
+ private_key_str = encode_private_key(private_key)
64
+ public_key_str = encode_public_key(public_key)
65
+
66
+ # Return the public-private key pair as strings
67
+ return (public_key_str, private_key_str)
68
+
69
+
70
+ def generate_vehicle_command_hmac(
71
+ command: str, timestamp: str, vehicle_key: str, private_key: str
72
+ ):
73
+ """Generate vehicle command hmac."""
74
+ message = (command + timestamp).encode("utf-8")
75
+ secret_key = get_secret_key(private_key, vehicle_key)
76
+ return get_message_signature(secret_key, message)
77
+
78
+
79
+ def get_message_signature(secret_key: bytes, message: bytes) -> str:
80
+ """Get message signature."""
81
+ return hmac.new(secret_key, message, hashlib.sha256).hexdigest()
82
+
83
+
84
+ def get_secret_key(private_key_str: str, public_key_str: str) -> bytes:
85
+ """Get HKDF derived secrety key from private/public key pair."""
86
+ private_key = decode_private_key(private_key_str)
87
+ public_key = decode_public_key(public_key_str)
88
+ secret = private_key.exchange(ec.ECDH(), public_key)
89
+ hkdf = HKDF(algorithm=hashes.SHA256(), length=32, salt=None, info=b"")
90
+ return hkdf.derive(secret)
@@ -11,7 +11,7 @@ from typing import TYPE_CHECKING, Any
11
11
  from uuid import uuid4
12
12
 
13
13
  import async_timeout
14
- from aiohttp import ClientWebSocketResponse, WSMsgType
14
+ from aiohttp import ClientWebSocketResponse, WSMessage, WSMsgType
15
15
 
16
16
  if TYPE_CHECKING:
17
17
  from .rivian import Rivian
@@ -81,6 +81,7 @@ class WebSocketMonitor:
81
81
  await cancel_task(self._receiver_task)
82
82
  self._disconnect = False
83
83
  # pylint: disable=protected-access
84
+ assert self._account._session
84
85
  self._ws = await self._account._session.ws_connect(
85
86
  url=self._url, headers={"sec-websocket-protocol": "graphql-transport-ws"}
86
87
  )
@@ -94,7 +95,7 @@ class WebSocketMonitor:
94
95
  ) -> Callable[[], Awaitable[None]] | None:
95
96
  """Start a subscription."""
96
97
  if not self.connected:
97
- return
98
+ return None
98
99
  _id = str(uuid4())
99
100
  self._subscriptions[_id] = (callback, payload)
100
101
  await self._subscribe(_id, payload)
@@ -104,12 +105,14 @@ class WebSocketMonitor:
104
105
  if _id in self._subscriptions:
105
106
  del self._subscriptions[_id]
106
107
  if self.connected:
108
+ assert self._ws
107
109
  await self._ws.send_json({"id": _id, "type": "complete"})
108
110
 
109
111
  return unsubscribe
110
112
 
111
113
  async def _subscribe(self, _id: str, payload: dict[str, Any]) -> None:
112
114
  """Send a subscribe request."""
115
+ assert self._ws
113
116
  await self._ws.send_json({"id": _id, "payload": payload, "type": "subscribe"})
114
117
 
115
118
  async def _resubscribe_all(self) -> None:
@@ -142,8 +145,7 @@ class WebSocketMonitor:
142
145
  self._connection_ack.set()
143
146
  elif data_type == "next":
144
147
  if (_id := data.get("id")) in self._subscriptions:
145
- if callback := self._subscriptions[_id][0]:
146
- callback(data)
148
+ self._subscriptions[_id][0](data)
147
149
  else:
148
150
  self._log_message(msg)
149
151
  elif msg.type == WSMsgType.ERROR:
@@ -159,7 +161,8 @@ class WebSocketMonitor:
159
161
  attempt = 0
160
162
  while not self._disconnect:
161
163
  while self.connected:
162
- if self._receiver_task.done(): # Need to restart the receiver
164
+ if self._receiver_task and self._receiver_task.done():
165
+ # Need to restart the receiver
163
166
  self._receiver_task = asyncio.ensure_future(self._receiver())
164
167
  await asyncio.sleep(1)
165
168
  if not self._disconnect:
@@ -191,7 +194,9 @@ class WebSocketMonitor:
191
194
  await self._ws.close()
192
195
  await cancel_task(self._monitor_task, self._receiver_task)
193
196
 
194
- def _log_message(self, message: str | Exception, is_error: bool = False) -> None:
197
+ def _log_message(
198
+ self, message: str | Exception | WSMessage, is_error: bool = False
199
+ ) -> None:
195
200
  """Log a message."""
196
201
  log_method = _LOGGER.error if is_error else _LOGGER.debug
197
202
  log_method(message)