File size: 5,760 Bytes
e5034c3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 | use std::path::PathBuf;
use std::sync::Arc;
use anyhow::{Context as _, Result};
use forge_domain::{
Context, ContextMessage, DataGenerationParameters, ResultStreamExt, Template, ToolDefinition,
};
use futures::StreamExt;
use futures::stream::{self, BoxStream};
use schemars::Schema;
use tracing::{debug, info};
use crate::{AppConfigService, FsReadService, ProviderService, Services, TemplateEngine};
pub struct DataGenerationApp<A> {
services: Arc<A>,
}
type JsonSchema = String;
type SystemPrompt = String;
type UserPrompt = String;
type Input = Vec<serde_json::Value>;
impl<A: Services> DataGenerationApp<A> {
pub fn new(services: Arc<A>) -> Self {
Self { services }
}
/// Helper function to read a file from a path, resolving it relative to cwd
/// if necessary
async fn read_file(&self, path: PathBuf) -> Result<String> {
let resolved_path = if path.is_absolute() {
path
} else {
let cwd = self.services.get_environment().cwd;
cwd.join(path)
};
let content = self
.services
.read(resolved_path.display().to_string(), None, None)
.await?
.content
.file_content()
.to_owned();
Ok(content)
}
async fn read_file_opt(&self, path: Option<PathBuf>) -> Result<Option<String>> {
match path {
Some(path) => self.read_file(path).await.map(Some),
None => Ok(None),
}
}
async fn load_parameters(
&self,
params: DataGenerationParameters,
) -> Result<(JsonSchema, Option<SystemPrompt>, Option<UserPrompt>, Input)> {
debug!("Loading data generation parameters");
// Read all files in parallel
let (schema, system_prompt, user_prompt, input) = tokio::join!(
self.read_file(params.schema.clone()),
self.read_file_opt(params.system_prompt),
self.read_file_opt(params.user_prompt),
self.read_file(params.input)
);
let input: Vec<serde_json::Value> = input?
.lines()
.map(|text| {
serde_json::from_str(text).with_context(|| "Could not parse the input file")
})
.collect::<Result<Vec<_>>>()?;
debug!("Loaded {} input items", input.len());
Ok((schema?, system_prompt?, user_prompt?, input))
}
pub async fn execute(
&self,
params: DataGenerationParameters,
) -> Result<BoxStream<'static, Result<serde_json::Value>>> {
let concurrency = params.concurrency;
let (schema, system_prompt, user_prompt, input) = self.load_parameters(params).await?;
info!(
"Starting data generation with {} items (concurrency: {})",
input.len(),
concurrency
);
let model_config = self
.services
.get_session_config()
.await
.ok_or_else(|| forge_domain::Error::NoDefaultSession)?;
let provider = self.services.get_provider(model_config.provider).await?;
let model_id = model_config.model;
debug!("Using provider: {}, model: {}", provider.id, model_id);
let schema: Schema =
serde_json::from_str(&schema).with_context(|| "Could not parse the JSON schema")?;
let mut context =
Context::default().add_tool(ToolDefinition::new("output").input_schema(schema));
if let Some(content) = system_prompt {
context = context.add_message(ContextMessage::system(content))
}
let services = self.services.clone();
let json_stream = input.into_iter().map(move |input| {
let provider = provider.clone();
let context = context.clone();
let user_prompt = user_prompt.clone();
let model_id = model_id.clone();
let services = services.clone();
async move {
debug!("Processing data generation request");
let provider = provider.clone();
let mut context = context.clone();
let content = if let Some(ref content) = user_prompt {
TemplateEngine::default().render_template(Template::new(content), &input)?
} else {
serde_json::to_string(&input)?
};
context =
context.add_message(ContextMessage::user(content, Some(model_id.clone())));
let stream = services.chat(&model_id, context, provider.clone()).await?;
let response = stream.into_full(false).await?;
anyhow::Ok((input, response))
}
});
let json_stream = stream::iter(json_stream)
.buffer_unordered(concurrency)
.map(|result| {
result.and_then(|(input, response)| {
response
.tool_calls
.into_iter()
.map(|tool| {
let output = tool.arguments.parse()?;
let mut value = serde_json::Map::new();
value.insert("input".to_string(), input.clone());
value.insert("output".to_string(), output);
Ok(serde_json::Value::from(value))
})
.collect::<Result<Vec<_>>>()
})
})
.flat_map(|data| match data {
Ok(data) => stream::iter(data).map(Ok).boxed(),
Err(err) => stream::iter(Err(err)).boxed(),
})
.boxed();
Ok(json_stream)
}
}
|