Skip to main content

openai_interface/files/
retrieve.rs

1//! This module provides functionality for retrieving files from the OpenAI API.
2//!
3//! It includes the `RetrieveRequest` struct which implements the `Get` and `GetNoStream` traits
4//! to build URLs and fetch file data asynchronously.
5//!
6//! # Example
7//!
8//! ```rust,no_run
9//! use openai_interface::files::retrieve::*;
10//! use openai_interface::{
11//!     files::{list::request::ListFilesRequest, retrieve::request::RetrieveRequest},
12//!     rest::{
13//!         RequestOptions,
14//!         default_client,
15//!         get::{Get, GetNoStream},
16//!     },
17//! };
18//! use anyhow::bail;
19//! use futures_util::future;
20//!
21//! const MODELSCOPE_BASE_URL: &str = "https://dashscope.aliyuncs.com/compatible-mode/v1/";
22//! const MODELSCOPE_KEY: &str = "YOUR_API_KEY";
23//!
24//! #[tokio::main]
25//! async fn main() -> Result<(), anyhow::Error> {
26//!     // Needs the `ferritls` cargo feature; drop this line if you install
27//!     // your own rustls crypto provider (see `openai_interface::rest`).
28//!     # #[cfg(feature = "ferritls")]
29//!     openai_interface::rest::install_crypto_provider().ok();
30//!
31//!     // first get all files
32//!     let list_request = ListFilesRequest {
33//!         limit: Some(5), // avoid rate limit
34//!         ..Default::default()
35//!     };
36//!
37//!     let list_response = list_request
38//!         .get_response(&default_client(), MODELSCOPE_BASE_URL, &RequestOptions::bearer(MODELSCOPE_KEY))
39//!         .await?;
40//!
41//!     let client = default_client();
42//!     let futures: Vec<_> = list_response
43//!         .data
44//!         .iter()
45//!         .map(|file_object| {
46//!             let client = client.clone();
47//!             let file_id = file_object.id.clone();
48//!             let base_url = MODELSCOPE_BASE_URL.to_string();
49//!             let key = MODELSCOPE_KEY.to_string();
50//!             async move {
51//!                 let retrieve_request = RetrieveRequest {
52//!                     file_id: &file_id,
53//!                     ..Default::default()
54//!                 };
55//!                 retrieve_request
56//!                     .get_response(&client, &base_url, &RequestOptions::bearer(&key))
57//!                     .await
58//!             }
59//!         })
60//!         .collect();
61//!
62//!     let results = future::join_all(futures).await;
63//!
64//!     for (i, result) in results.iter().enumerate() {
65//!         match result {
66//!             Ok(file_object) => {
67//!                 assert_eq!(&list_response.data[i].id, &file_object.id);
68//!                 assert_eq!(&list_response.data[i].filename, &file_object.filename);
69//!                 assert_eq!(&list_response.data[i].purpose, &file_object.purpose);
70//!             }
71//!             Err(e) => {
72//!                 bail!(
73//!                     "Failed to get response: {e:#}. The file is: index {i}, {:?}",
74//!                     list_response.data[i]
75//!                 )
76//!             }
77//!         }
78//!     }
79//!
80//!     Ok(())
81//! }
82//! ```
83
84pub mod request {
85    use std::collections::HashMap;
86
87    use url::Url;
88
89    use crate::{
90        errors::OapiError,
91        rest::get::{Get, GetNoStream},
92    };
93
94    /// Query parameters for retrieving a file.
95    #[derive(Debug, Clone, Default)]
96    pub struct RetrieveRequest<'a> {
97        pub file_id: &'a str,
98        pub extra_query: HashMap<&'a str, &'a str>,
99    }
100
101    impl<'a> Get for RetrieveRequest<'a> {
102        /// base_url should look like <https://api.openai.com/v1>
103        fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
104            let mut url =
105                Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
106            url.path_segments_mut()
107                .map_err(|_| OapiError::UrlError(url::ParseError::RelativeUrlWithoutBase))?
108                .push("files")
109                .push(self.file_id);
110
111            for (key, value) in &self.extra_query {
112                url.query_pairs_mut().append_pair(key, value);
113            }
114
115            Ok(url.to_string())
116        }
117    }
118
119    impl<'a> GetNoStream for RetrieveRequest<'a> {
120        type Response = crate::files::FileObject;
121    }
122}
123
124#[cfg(test)]
125mod tests {
126    use super::*;
127    use crate::{
128        files::list::request::ListFilesRequest,
129        rest::{
130            RequestOptions, default_client,
131            get::{Get, GetNoStream},
132        },
133    };
134    use anyhow::bail;
135    use futures_util::future::{self};
136
137    const MODELSCOPE_BASE_URL: &str = "https://dashscope.aliyuncs.com/compatible-mode/v1/";
138
139    fn modelstudio_key() -> Option<String> {
140        std::env::var("QWEN_API_KEY")
141            .ok()
142            .map(|key| key.trim().to_string())
143            .filter(|key| !key.is_empty())
144    }
145
146    #[test]
147    fn test_build_url() {
148        let request = request::RetrieveRequest {
149            file_id: "file_id",
150            ..Default::default()
151        };
152        let url = request.build_url("https://api.openai.com/v1/").unwrap();
153        assert_eq!(url, "https://api.openai.com/v1/files/file_id");
154    }
155
156    #[tokio::test]
157    async fn test_retrieve_file() -> Result<(), anyhow::Error> {
158        let Some(api_key) = modelstudio_key() else {
159            println!("Skipping: set QWEN_API_KEY to run this test");
160            return Ok(());
161        };
162
163        // first get all files
164        let list_request = ListFilesRequest {
165            limit: Some(5), // avoid rate limit
166            ..Default::default()
167        };
168
169        let list_response = list_request
170            .get_response(
171                &default_client(),
172                MODELSCOPE_BASE_URL,
173                &RequestOptions::bearer(&api_key),
174            )
175            .await?;
176
177        let client = default_client();
178        let futures: Vec<_> = list_response
179            .data
180            .iter()
181            .map(|file_object| {
182                let client = client.clone();
183                let file_id = file_object.id.clone();
184                let base_url = MODELSCOPE_BASE_URL.to_string();
185                let key = api_key.clone();
186                async move {
187                    let retrieve_request = request::RetrieveRequest {
188                        file_id: &file_id,
189                        ..Default::default()
190                    };
191                    retrieve_request
192                        .get_response(&client, &base_url, &RequestOptions::bearer(&key))
193                        .await
194                }
195            })
196            .collect();
197
198        let results = future::join_all(futures).await;
199
200        for (i, result) in results.iter().enumerate() {
201            match result {
202                Ok(file_object) => {
203                    assert_eq!(&list_response.data[i].id, &file_object.id);
204                    assert_eq!(&list_response.data[i].filename, &file_object.filename);
205                    assert_eq!(&list_response.data[i].purpose, &file_object.purpose);
206                    // assert_eq!(&list_response.data[i], file_object);
207                }
208                Err(e) => {
209                    bail!(
210                        "Failed to get response: {e:#}. The file is: index {i}, {:?}",
211                        list_response.data[i]
212                    )
213                }
214            }
215        }
216
217        Ok(())
218    }
219}