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