Skip to main content

qcs_api_client_common/configuration/
py.rs

1#![allow(unused_qualifications)]
2#![allow(non_local_definitions, reason = "necessary for pyo3::pymethods")]
3use std::collections::BTreeSet;
4
5use pyo3::{
6    exceptions::PyValueError,
7    prelude::*,
8    types::{PyAnyMethods, PyString},
9};
10use rigetti_pyo3::{create_init_submodule, impl_repr, py_function_sync_async, sync::Awaitable};
11use tokio_util::sync::CancellationToken;
12
13#[cfg(feature = "stubs")]
14use pyo3_stub_gen::derive::{gen_stub_pyfunction, gen_stub_pymethods};
15
16use crate::configuration::{
17    API_URL_VAR, ClientConfigurationBuilderError, DEFAULT_API_URL, DEFAULT_GRPC_API_URL,
18    DEFAULT_PROFILE_NAME, DEFAULT_QUILC_URL, DEFAULT_QVM_URL, GRPC_API_URL_VAR, PROFILE_NAME_VAR,
19    QUILC_URL_VAR, QVM_URL_VAR,
20    secrets::{DEFAULT_SECRETS_PATH, SECRETS_PATH_VAR},
21    settings::{DEFAULT_SETTINGS_PATH, SETTINGS_PATH_VAR},
22};
23use crate::errors;
24
25use super::{
26    ClientConfiguration, ClientConfigurationBuilder, LoadError, OAuthGrant, OAuthSession,
27    RefreshToken, TokenDispatcher,
28    error::TokenError,
29    secrets::{SecretAccessToken, SecretRefreshToken},
30    settings::AuthServer,
31    tokens::{AuthTokens, ClientCredentials, ClientSecret, ExternallyManaged},
32};
33
34create_init_submodule! {
35    classes: [
36        ClientConfiguration,
37        ClientConfigurationBuilder,
38        AuthServer,
39        OAuthSession,
40        RefreshToken,
41        ClientCredentials,
42        ClientSecret,
43        ExternallyManaged,
44        AuthTokens,
45        SecretAccessToken,
46        SecretRefreshToken,
47        TokenDispatcher
48    ],
49
50    consts: [
51        API_URL_VAR,
52        DEFAULT_API_URL,
53        DEFAULT_GRPC_API_URL,
54        DEFAULT_PROFILE_NAME,
55        DEFAULT_QUILC_URL,
56        DEFAULT_QVM_URL,
57        DEFAULT_SECRETS_PATH,
58        DEFAULT_SETTINGS_PATH,
59        GRPC_API_URL_VAR,
60        PROFILE_NAME_VAR,
61        QUILC_URL_VAR,
62        QVM_URL_VAR,
63        SECRETS_PATH_VAR,
64        SETTINGS_PATH_VAR
65    ],
66
67    errors: [
68        errors::ClientConfigurationBuilderError,
69        errors::ConfigurationError,
70        errors::LoadError,
71        errors::TokenError
72    ],
73
74    funcs: [
75        py_get_oauth_session,
76        py_get_oauth_session_async,
77        py_get_bearer_access_token,
78        py_get_bearer_access_token_async,
79        py_request_access_token,
80        py_request_access_token_async
81    ],
82
83}
84
85#[cfg(feature = "stubs")]
86#[derive(IntoPyObject)]
87struct Final<T>(T);
88
89#[cfg(feature = "stubs")]
90impl<T> pyo3_stub_gen::PyStubType for Final<T> {
91    fn type_output() -> pyo3_stub_gen::TypeInfo {
92        pyo3_stub_gen::TypeInfo::with_module("typing.Final", "typing".into())
93    }
94}
95
96/// Adds module-level `str` to the `qcs_api_client_common._qcs_api_client_common.configuration` stub file.
97macro_rules! stub_consts {
98    ( $($name:ident),* ) => {
99        $(
100            #[cfg(feature = "stubs")]
101            ::pyo3_stub_gen::module_variable!(
102                "qcs_api_client_common._qcs_api_client_common.configuration",
103                stringify!($name),
104                Final<&str>,
105                Final($name)
106            );
107        )*
108    };
109}
110
111stub_consts!(
112    API_URL_VAR,
113    DEFAULT_API_URL,
114    DEFAULT_GRPC_API_URL,
115    DEFAULT_PROFILE_NAME,
116    DEFAULT_QUILC_URL,
117    DEFAULT_QVM_URL,
118    DEFAULT_SECRETS_PATH,
119    DEFAULT_SETTINGS_PATH,
120    GRPC_API_URL_VAR,
121    PROFILE_NAME_VAR,
122    QUILC_URL_VAR,
123    QVM_URL_VAR,
124    SECRETS_PATH_VAR,
125    SETTINGS_PATH_VAR
126);
127
128/// Manual implementation to extract tokens from Python objects.
129///
130/// For Python functions that require a `SecretRefreshToken`,
131/// users can provide a Python `str`, a `RefreshToken`, or a `SecretRefreshToken`.
132impl FromPyObject<'_, '_> for SecretRefreshToken {
133    type Error = PyErr;
134
135    fn extract(obj: Borrowed<'_, '_, PyAny>) -> Result<Self, Self::Error> {
136        if let Ok(token) = obj.cast::<PyString>() {
137            Ok(Self::__new__(token.extract()?))
138        } else if let Ok(token) = obj.cast::<RefreshToken>() {
139            Ok(token.borrow().refresh_token.clone())
140        } else if let Ok(token) = obj.cast::<Self>() {
141            Ok(token.borrow().clone())
142        } else {
143            Err(PyValueError::new_err(
144                "expected str | SecretRefreshToken | RefreshToken",
145            ))
146        }
147    }
148}
149
150impl FromPyObject<'_, '_> for SecretAccessToken {
151    type Error = PyErr;
152
153    fn extract(obj: Borrowed<'_, '_, PyAny>) -> Result<Self, Self::Error> {
154        if let Ok(token) = obj.cast::<PyString>() {
155            Ok(Self::__new__(token.extract()?))
156        } else if let Ok(token) = obj.cast::<Self>() {
157            Ok(token.borrow().clone())
158        } else {
159            Err(PyValueError::new_err("expected str | SecretAccessToken"))
160        }
161    }
162}
163
164impl FromPyObject<'_, '_> for ClientSecret {
165    type Error = PyErr;
166
167    fn extract(obj: Borrowed<'_, '_, PyAny>) -> Result<Self, Self::Error> {
168        if let Ok(token) = obj.cast::<PyString>() {
169            Ok(Self::__new__(token.extract()?))
170        } else if let Ok(token) = obj.cast::<Self>() {
171            Ok(token.borrow().clone())
172        } else {
173            Err(PyValueError::new_err("expected str | ClientSecret"))
174        }
175    }
176}
177
178impl_repr!(RefreshToken);
179
180#[cfg_attr(feature = "stubs", gen_stub_pymethods)]
181#[pymethods]
182impl RefreshToken {
183    #[new]
184    const fn __new__(refresh_token: SecretRefreshToken) -> Self {
185        Self::new(refresh_token)
186    }
187}
188
189impl_repr!(ClientCredentials);
190
191#[cfg_attr(feature = "stubs", gen_stub_pymethods)]
192#[pymethods]
193impl ClientCredentials {
194    #[new]
195    fn __new__(client_id: String, client_secret: String) -> Self {
196        Self::new(client_id, ClientSecret::from(client_secret))
197    }
198}
199
200impl_repr!(ExternallyManaged);
201
202#[cfg_attr(not(feature = "stubs"), optipy::strip_pyo3(only_stubs))]
203#[cfg_attr(feature = "stubs", gen_stub_pymethods)]
204#[pymethods]
205impl ExternallyManaged {
206    #[new]
207    fn __new__(
208        #[gen_stub(
209            override_type(
210                type_repr="collections.abc.Callable[[AuthServer], str]",
211                imports=("collections.abc")
212            )
213        )]
214        refresh_function: Bound<'_, PyAny>,
215    ) -> PyResult<Self> {
216        if !refresh_function.is_callable() {
217            return Err(pyo3::exceptions::PyTypeError::new_err(
218                "refresh_function must be callable",
219            ));
220        }
221
222        let refresh_function = refresh_function.unbind();
223
224        #[allow(trivial_casts)] // Compilation fails without the cast.
225        // The provided refresh function will panic if there is an issue with the refresh function.
226        // This raises a `PanicException` within Python.
227        let refresh_closure = move |auth_server: AuthServer| {
228            let refresh_function = Python::attach(|py| refresh_function.clone_ref(py));
229            Box::pin(async move {
230                Python::attach(|py| {
231                    let result = refresh_function.call1(py, (auth_server,));
232                    match result {
233                        Ok(value) => value
234                            .extract::<String>(py)
235                            .map_or_else(|_| panic!("ExternallyManaged refresh function returned an unexpected type. Expected a string, got {value:?}"), Ok),
236                        Err(err) => Err(Box::<dyn std::error::Error + Send + Sync>::from(err))
237                    }
238                })
239            }) as super::tokens::RefreshResult
240        };
241
242        Ok(Self::new(refresh_closure))
243    }
244}
245
246impl_repr!(AuthTokens);
247
248#[cfg_attr(feature = "stubs", gen_stub_pymethods)]
249#[pymethods]
250impl AuthTokens {
251    #[new]
252    fn __new__(py: Python<'_>, auth_server: AuthServer) -> PyResult<Self> {
253        pyo3_async_runtimes::tokio::run(py, async move {
254            let cancel_token = cancel_token_with_ctrl_c();
255            Self::interactive_login(cancel_token, &auth_server)
256                .await
257                .map_err(|err| LoadError::from(err).into())
258        })
259    }
260}
261
262#[cfg(feature = "stubs")]
263pyo3_stub_gen::impl_stub_type!(
264    OAuthGrant = RefreshToken | ClientConfiguration | ExternallyManaged | AuthTokens
265);
266
267impl_repr!(OAuthSession);
268
269#[cfg_attr(feature = "stubs", gen_stub_pymethods)]
270#[pymethods]
271impl OAuthSession {
272    #[new]
273    #[pyo3(signature = (payload, auth_server, access_token = None))]
274    const fn __new__(
275        payload: OAuthGrant,
276        auth_server: AuthServer,
277        access_token: Option<SecretAccessToken>,
278    ) -> Self {
279        Self::new(payload, auth_server, access_token)
280    }
281
282    #[pyo3(name = "validate")]
283    fn py_validate(&self) -> Result<SecretAccessToken, TokenError> {
284        self.validate()
285    }
286
287    #[pyo3(name = "request_access_token")]
288    fn py_request_access_token(&self, py: Python<'_>) -> PyResult<SecretAccessToken> {
289        py_request_access_token(py, self.clone())
290    }
291
292    #[pyo3(name = "request_access_token_async")]
293    fn py_request_access_token_async<'py>(
294        &self,
295        py: Python<'py>,
296    ) -> PyResult<Awaitable<'py, SecretAccessToken>> {
297        py_request_access_token_async(py, self.clone())
298    }
299}
300
301py_function_sync_async! {
302    #[cfg_attr(feature = "stubs", gen_stub_pyfunction(module = "qcs_api_client_common._qcs_api_client_common.configuration"))]
303    #[pyfunction]
304    async fn get_oauth_session(tokens: Option<TokenDispatcher>) -> PyResult<OAuthSession> {
305        Ok(tokens.ok_or(TokenError::NoRefreshToken)?.tokens().await)
306    }
307}
308
309py_function_sync_async! {
310    /// Gets the `Bearer` access token, refreshing it if it is expired.
311    ///
312    /// # Errors
313    ///
314    /// Raises a `TokenError` if there's a problem providing the token.
315    #[cfg_attr(feature = "stubs", gen_stub_pyfunction(module = "qcs_api_client_common._qcs_api_client_common.configuration"))]
316    #[pyfunction]
317    async fn get_bearer_access_token(configuration: ClientConfiguration) -> PyResult<SecretAccessToken> {
318        configuration.get_bearer_access_token().await.map_err(PyErr::from)
319    }
320}
321
322py_function_sync_async! {
323    /// Request and return an updated access token using these credentials.
324    ///
325    /// # Errors
326    ///
327    /// Raises a `TokenError` if there's a problem providing the token.
328    #[cfg_attr(feature = "stubs", gen_stub_pyfunction(module = "qcs_api_client_common._qcs_api_client_common.configuration"))]
329    #[pyfunction]
330    async fn request_access_token(session: OAuthSession) -> PyResult<SecretAccessToken> {
331        session.clone().request_access_token().await.cloned().map_err(PyErr::from)
332    }
333}
334
335impl_repr!(ClientConfiguration);
336
337#[cfg_attr(feature = "stubs", gen_stub_pymethods)]
338#[pymethods]
339impl ClientConfiguration {
340    #[new]
341    #[pyo3(signature = (
342            api_url = None, grpc_api_url = None, quilc_url = None, qvm_url = None,
343            oauth_session = None,
344            ))]
345    fn __new__(
346        api_url: Option<String>,
347        grpc_api_url: Option<String>,
348        quilc_url: Option<String>,
349        qvm_url: Option<String>,
350        oauth_session: Option<OAuthSession>,
351    ) -> Self {
352        let mut builder = ClientConfigurationBuilder::default();
353
354        if let Some(api_url) = api_url {
355            builder.api_url(api_url);
356        }
357
358        if let Some(grpc_api_url) = grpc_api_url {
359            builder.grpc_api_url(grpc_api_url);
360        }
361
362        if let Some(quilc_url) = quilc_url {
363            builder.quilc_url(quilc_url);
364        }
365
366        if let Some(qvm_url) = qvm_url {
367            builder.qvm_url(qvm_url);
368        }
369
370        builder.oauth_session(oauth_session);
371
372        builder
373            .build()
374            .expect("our builder is valid regardless of which URLs are set")
375    }
376
377    #[staticmethod]
378    #[pyo3(name = "load_default")]
379    fn py_load_default(_py: Python<'_>) -> Result<Self, LoadError> {
380        Self::load_default()
381    }
382
383    #[staticmethod]
384    #[pyo3(name = "load_default_with_login")]
385    fn py_load_default_with_login(py: Python<'_>) -> PyResult<Self> {
386        pyo3_async_runtimes::tokio::run(py, async move {
387            let cancel_token = cancel_token_with_ctrl_c();
388            Self::load_with_login(cancel_token, None)
389                .await
390                .map_err(Into::into)
391        })
392    }
393
394    #[staticmethod]
395    #[pyo3(name = "builder")]
396    fn py_builder() -> ClientConfigurationBuilder {
397        ClientConfigurationBuilder::default()
398    }
399
400    #[staticmethod]
401    #[pyo3(name = "load_profile")]
402    fn py_load_profile(_py: Python<'_>, profile_name: String) -> Result<Self, LoadError> {
403        Self::load_profile(profile_name)
404    }
405
406    /// Gets the `Bearer` access token, refreshing it if it is expired.
407    ///
408    /// # Errors
409    ///
410    /// Raises a `TokenError` if there's a problem providing the token.
411    #[pyo3(name = "get_bearer_access_token")]
412    fn py_get_bearer_access_token(&self, py: Python<'_>) -> PyResult<SecretAccessToken> {
413        py_get_bearer_access_token(py, self.clone())
414    }
415
416    #[pyo3(name = "get_bearer_access_token_async")]
417    fn py_get_bearer_access_token_async<'py>(
418        &self,
419        py: Python<'py>,
420    ) -> PyResult<Awaitable<'py, SecretAccessToken>> {
421        py_get_bearer_access_token_async(py, self.clone())
422    }
423
424    /// Get the configured [`OAuthSession`].
425    ///
426    /// # Errors
427    ///
428    /// Raises a `TokenError` if there is a problem fetching the tokens.
429    pub fn get_oauth_session(&self, py: Python<'_>) -> PyResult<OAuthSession> {
430        py_get_oauth_session(py, self.oauth_session.clone())
431    }
432
433    fn get_oauth_session_async<'py>(
434        &self,
435        py: Python<'py>,
436    ) -> PyResult<Awaitable<'py, OAuthSession>> {
437        py_get_oauth_session_async(py, self.oauth_session.clone())
438    }
439}
440
441#[cfg_attr(feature = "stubs", gen_stub_pymethods)]
442#[pymethods]
443impl ClientConfigurationBuilder {
444    #[new]
445    fn __new__() -> Self {
446        Self::default()
447    }
448
449    /// The [`OAuthSession`] to use to authenticate with the QCS API.
450    ///
451    /// When set to [`None`], the configuration will not manage an OAuth Session, and access to the
452    /// QCS API will be limited to unauthenticated routes.
453    #[setter]
454    fn set_oauth_session(&mut self, oauth_session: Option<OAuthSession>) {
455        self.oauth_session = Some(oauth_session.map(Into::into));
456    }
457
458    #[pyo3(name = "build")]
459    fn py_build(&self) -> Result<ClientConfiguration, ClientConfigurationBuilderError> {
460        self.build()
461    }
462}
463
464impl_repr!(AuthServer);
465
466#[cfg_attr(feature = "stubs", gen_stub_pymethods)]
467#[pymethods]
468impl AuthServer {
469    #[new]
470    #[pyo3(signature = (client_id, issuer, scopes = None))]
471    const fn __new__(client_id: String, issuer: String, scopes: Option<BTreeSet<String>>) -> Self {
472        Self::new(client_id, issuer, scopes)
473    }
474
475    #[staticmethod]
476    #[pyo3(name = "default")]
477    fn py_default() -> Self {
478        Self::default()
479    }
480}
481
482fn cancel_token_with_ctrl_c() -> CancellationToken {
483    let cancel_token = CancellationToken::new();
484    let cancel_token_ctrl_c = cancel_token.clone();
485    tokio::spawn(cancel_token.clone().run_until_cancelled_owned(async move {
486        match tokio::signal::ctrl_c().await {
487            Ok(()) => cancel_token_ctrl_c.cancel(),
488            Err(error) => eprintln!("Failed to register signal handler: {error}"),
489        }
490    }));
491    cancel_token
492}