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