use std::collections::HashMap;
use std::io;
use std::io::Read;
use std::sync::{Arc, Mutex};
use hyper;
use hyper::client::Client as HyperClient;
use hyper::header::{Authorization, Basic, ContentType, Headers};
use serde;
use serde_json;
use super::{Request, Response};
use util::HashableValue;
use error::Error;
pub struct Client {
url: String,
user: Option<String>,
pass: Option<String>,
client: HyperClient,
nonce: Arc<Mutex<u64>>,
}
impl Client {
pub fn new(url: String, user: Option<String>, pass: Option<String>) -> Client {
debug_assert!(pass.is_none() || user.is_some());
Client {
url: url,
user: user,
pass: pass,
client: HyperClient::new(),
nonce: Arc::new(Mutex::new(0)),
}
}
pub fn do_rpc<T: for<'a> serde::de::Deserialize<'a>>(
&self,
rpc_name: &str,
args: &[serde_json::value::Value],
) -> Result<T, Error> {
let request = self.build_request(rpc_name, args);
let response = self.send_request(&request)?;
Ok(response.into_result()?)
}
fn send_raw<B, R>(&self, body: &B) -> Result<R, Error>
where
B: serde::ser::Serialize,
R: for<'de> serde::de::Deserialize<'de>,
{
let request_raw = serde_json::to_vec(body)?;
let mut headers = Headers::new();
headers.set(ContentType::json());
if let Some(ref user) = self.user {
headers.set(Authorization(Basic {
username: user.clone(),
password: self.pass.clone(),
}));
}
let retry_headers = headers.clone();
let hyper_request = self.client.post(&self.url).headers(headers).body(&request_raw[..]);
let mut stream = match hyper_request.send() {
Ok(s) => s,
Err(hyper::error::Error::Io(e)) => {
if e.kind() == io::ErrorKind::BrokenPipe
|| e.kind() == io::ErrorKind::ConnectionAborted
{
try!(self
.client
.post(&self.url)
.headers(retry_headers)
.body(&request_raw[..])
.send()
.map_err(Error::Hyper))
} else {
return Err(Error::Hyper(hyper::error::Error::Io(e)));
}
}
Err(e) => {
return Err(Error::Hyper(e));
}
};
let response: R = serde_json::from_reader(&mut stream)?;
stream.bytes().count(); Ok(response)
}
pub fn send_request(&self, request: &Request) -> Result<Response, Error> {
let response: Response = self.send_raw(&request)?;
if response.jsonrpc != None && response.jsonrpc != Some(From::from("2.0")) {
return Err(Error::VersionMismatch);
}
if response.id != request.id {
return Err(Error::NonceMismatch);
}
Ok(response)
}
pub fn send_batch(&self, requests: &[Request]) -> Result<Vec<Option<Response>>, Error> {
if requests.len() < 1 {
return Err(Error::EmptyBatch);
}
let responses: Vec<Response> = self.send_raw(&requests)?;
if responses.len() > requests.len() {
return Err(Error::WrongBatchResponseSize);
}
let ids: Vec<serde_json::Value> = responses.iter().map(|r| r.id.clone()).collect();
let mut resp_by_id = HashMap::new();
for (id, resp) in ids.iter().zip(responses.into_iter()) {
if let Some(dup) = resp_by_id.insert(HashableValue(&id), resp) {
return Err(Error::BatchDuplicateResponseId(dup.id));
}
}
let results =
requests.into_iter().map(|r| resp_by_id.remove(&HashableValue(&r.id))).collect();
if let Some(incorrect) = resp_by_id.into_iter().nth(0) {
return Err(Error::WrongBatchResponseId(incorrect.1.id));
}
Ok(results)
}
pub fn build_request<'a, 'b>(
&self,
name: &'a str,
params: &'b [serde_json::Value],
) -> Request<'a, 'b> {
let mut nonce = self.nonce.lock().unwrap();
*nonce += 1;
Request {
method: name,
params: params,
id: From::from(*nonce),
jsonrpc: Some("2.0"),
}
}
pub fn last_nonce(&self) -> u64 {
*self.nonce.lock().unwrap()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sanity() {
let client = Client::new("localhost".to_owned(), None, None);
assert_eq!(client.last_nonce(), 0);
let req1 = client.build_request("test", &[]);
assert_eq!(client.last_nonce(), 1);
let req2 = client.build_request("test", &[]);
assert_eq!(client.last_nonce(), 2);
assert!(req1 != req2);
}
}