From e68ea1b8fc4f88e587121387ecac6858d04ebae2 Mon Sep 17 00:00:00 2001 From: Ruben Fiszel Date: Sat, 13 Aug 2022 17:25:16 +0200 Subject: [PATCH] feat: support union types (#398) --- backend/openapi.yaml | 13 +- backend/src/lib.rs | 2 + backend/src/parser.rs | 737 +------------------------------------- backend/src/parser_py.rs | 571 +++++++++++++++++++++++++++++ backend/src/parser_ts.rs | 215 +++++++++++ backend/src/scripts.rs | 8 +- backend/src/worker.rs | 9 +- frontend/src/lib/infer.ts | 15 +- 8 files changed, 822 insertions(+), 748 deletions(-) create mode 100644 backend/src/parser_py.rs create mode 100644 backend/src/parser_ts.rs diff --git a/backend/openapi.yaml b/backend/openapi.yaml index dfadda82c6..310d14c38a 100644 --- a/backend/openapi.yaml +++ b/backend/openapi.yaml @@ -3631,7 +3631,6 @@ components: - type: string enum: [ - "str", "float", "int", "bool", @@ -3646,13 +3645,25 @@ components: properties: resource: type: string + nullable: true + required: - resource + - type: object + properties: + str: + type: array + items: + type: string + nullable: true + required: + - str - type: object properties: list: type: string enum: ["str", "float", "int", "email"] + nullable: true required: - list has_default: diff --git a/backend/src/lib.rs b/backend/src/lib.rs index db8bc3d651..403aa11a4f 100644 --- a/backend/src/lib.rs +++ b/backend/src/lib.rs @@ -35,6 +35,8 @@ mod jobs; mod js_eval; mod oauth2; mod parser; +mod parser_py; +mod parser_ts; mod resources; mod schedule; mod scripts; diff --git a/backend/src/parser.rs b/backend/src/parser.rs index 8be373f34d..781cc5697a 100644 --- a/backend/src/parser.rs +++ b/backend/src/parser.rs @@ -6,20 +6,8 @@ * LICENSE-AGPL for a copy of the license. */ -use std::collections::HashMap; - -use itertools::Itertools; -use phf::phf_map; -use regex::Regex; use serde::Serialize; -use serde_json::json; -use crate::error; - -use rustpython_parser::{ - ast::{ExpressionType, Located, Number, StatementType, StringGroup, Varargs}, - parser, -}; #[derive(Serialize)] pub struct MainArgSignature { pub star_args: bool, @@ -40,7 +28,7 @@ pub enum InnerTyp { #[derive(Serialize, Clone)] #[serde(rename_all(serialize = "lowercase"))] pub enum Typ { - Str, + Str(Option>), Int, Float, Bool, @@ -61,726 +49,3 @@ pub struct Arg { pub default: Option, pub has_default: bool, } - -pub fn parse_python_signature(code: &str) -> error::Result { - let ast = parser::parse_program(code) - .map_err(|e| error::Error::ExecutionErr(format!("Error parsing code: {}", e.to_string())))? - .statements; - let param = ast.into_iter().find_map(|x| match x { - Located { - location: _, - node: - StatementType::FunctionDef { - is_async: _, - name, - args, - body: _, - decorator_list: _, - returns: _, - }, - } if &name == "main" => Some(*args), - _ => None, - }); - if let Some(params) = param { - //println!("{:?}", params); - let def_arg_start = params.args.len() - params.defaults.len(); - Ok(MainArgSignature { - star_args: params.vararg != Varargs::None, - star_kwargs: params.vararg != Varargs::None, - args: params - .args - .into_iter() - .enumerate() - .map(|(i, x)| { - let default = if i >= def_arg_start { - to_value(¶ms.defaults[i - def_arg_start].node) - } else { - None - }; - Arg { - name: x.arg, - typ: x.annotation.map_or(Typ::Unknown, |e| match *e { - Located { location: _, node: ExpressionType::Identifier { name } } => { - match name.as_ref() { - "str" => Typ::Str, - "float" => Typ::Float, - "int" => Typ::Int, - "bool" => Typ::Bool, - "dict" => Typ::Dict, - "list" => Typ::List(InnerTyp::Str), - "bytes" => Typ::Bytes, - "datetime" => Typ::Datetime, - "datetime.datetime" => Typ::Datetime, - _ => Typ::Unknown, - } - } - _ => Typ::Unknown, - }), - has_default: default.is_some(), - default, - } - }) - .collect(), - }) - } else { - Err(error::Error::ExecutionErr( - "main function was not findable".to_string(), - )) - } -} - -use swc_common::{sync::Lrc, FileName, SourceMap}; -use swc_ecma_ast::{ - AssignPat, BindingIdent, Decl, ExportDecl, FnDecl, Ident, ModuleDecl, ModuleItem, Pat, - TsArrayType, TsEntityName, TsKeywordTypeKind, TsType, TsTypeRef, -}; -use swc_ecma_parser::{lexer::Lexer, Parser, StringInput, Syntax, TsConfig}; - -pub fn parse_deno_signature(code: &str) -> error::Result { - let cm: Lrc = Default::default(); - let fm = cm.new_source_file(FileName::Custom("test.ts".into()), code.into()); - let lexer = Lexer::new( - // We want to parse ecmascript - Syntax::Typescript(TsConfig::default()), - // EsVersion defaults to es5 - Default::default(), - StringInput::from(&*fm), - None, - ); - - let mut parser = Parser::new_from(lexer); - - let mut err_s = "".to_string(); - for e in parser.take_errors() { - err_s += &e.into_kind().msg().to_string(); - } - - let ast = parser - .parse_module() - .map_err(|_| { - error::Error::ExecutionErr(format!( - "Error while parsing code, it is invalid typescript" - )) - })? - .body; - - // println!("{ast:?}"); - let params = - ast.into_iter().find_map(|x| match x { - ModuleItem::ModuleDecl(ModuleDecl::ExportDecl(ExportDecl { - decl: - Decl::Fn(FnDecl { - ident: Ident { span: _, sym, optional: _ }, - declare: _, - function, - }), - span: _, - })) if &sym.to_string() == "main" => Some(function.params), - _ => None, - }); - if let Some(params) = params { - Ok(MainArgSignature { - star_args: false, - star_kwargs: false, - args: params - .into_iter() - .map(|x| match x.pat { - Pat::Ident(ident) => { - let (name, typ) = binding_ident_to_arg(&ident)?; - Ok(Arg { name, typ, default: None, has_default: ident.id.optional }) - } - Pat::Assign(AssignPat { span: _, left, right, type_ann: _ }) => { - let (name, typ) = - left.as_ident().map(binding_ident_to_arg).ok_or_else(|| { - error::Error::ExecutionErr(format!( - "Arg {left:?} has unexpected syntax" - )) - })??; - Ok(Arg { - name, - typ, - default: serde_json::to_value(right) - .map_err(|e| error::Error::ExecutionErr(e.to_string()))? - .as_object() - .and_then(|x| x.get("value").to_owned()) - .cloned(), - - has_default: true, - }) - } - _ => Err(error::Error::ExecutionErr(format!( - "Arg {x:?} has unexpected syntax" - ))), - }) - .collect::, error::Error>>()?, - }) - } else { - Err(error::Error::ExecutionErr( - "main function was not findable (expected to find 'export main function(...)'" - .to_string(), - )) - } -} - -fn binding_ident_to_arg( - BindingIdent { id, type_ann }: &BindingIdent, -) -> anyhow::Result<(String, Typ)> { - Ok(( - id.sym.to_string(), - type_ann - .as_ref() - .map(|x| { - match &*x.type_ann { - TsType::TsKeywordType(t) => match t.kind { - TsKeywordTypeKind::TsObjectKeyword => Typ::Dict, - TsKeywordTypeKind::TsBooleanKeyword => Typ::Bool, - TsKeywordTypeKind::TsBigIntKeyword => Typ::Int, - TsKeywordTypeKind::TsNumberKeyword => Typ::Float, - TsKeywordTypeKind::TsStringKeyword => Typ::Str, - _ => Typ::Unknown, - }, - // TODO: we can do better here and extract the inner type of array - TsType::TsArrayType(TsArrayType { span: _, elem_type }) => { - match &**elem_type { - TsType::TsTypeRef(TsTypeRef { - span: _, - type_name: TsEntityName::Ident(Ident { span: _, sym, optional: _ }), - type_params: _, - }) => match sym.to_string().as_str() { - "Base64" => Typ::List(InnerTyp::Bytes), - "Email" => Typ::List(InnerTyp::Email), - "bigint" => Typ::List(InnerTyp::Int), - "number" => Typ::List(InnerTyp::Float), - _ => Typ::List(InnerTyp::Str), - }, - //TsType::TsKeywordType(()) - _ => Typ::List(InnerTyp::Str), - } - } - TsType::TsTypeRef(TsTypeRef { span: _, type_name, type_params }) => { - let sym = match type_name { - TsEntityName::Ident(Ident { span: _, sym, optional: _ }) => sym, - TsEntityName::TsQualifiedName(p) => &*p.right.sym, - }; - match sym.to_string().as_str() { - "Resource" => Typ::Resource( - type_params - .as_ref() - .and_then(|x| { - x.params.get(0).and_then(|y| { - y.as_ts_lit_type().and_then(|z| { - z.lit - .as_str() - .map(|a| a.to_owned().value.to_string()) - }) - }) - }) - .unwrap_or_else(|| "unknown".to_string()), - ), - "Base64" => Typ::Bytes, - "Email" => Typ::Email, - "Sql" => Typ::Sql, - _ => Typ::Unknown, - } - } - _ => Typ::Unknown, - } - }) - .unwrap_or(Typ::Unknown), - )) -} - -const STDIMPORTS: [&str; 301] = [ - "__future__", - "_abc", - "_aix_support", - "_ast", - "_asyncio", - "_bisect", - "_blake2", - "_bootsubprocess", - "_bz2", - "_codecs", - "_codecs_cn", - "_codecs_hk", - "_codecs_iso2022", - "_codecs_jp", - "_codecs_kr", - "_codecs_tw", - "_collections", - "_collections_abc", - "_compat_pickle", - "_compression", - "_contextvars", - "_crypt", - "_csv", - "_ctypes", - "_curses", - "_curses_panel", - "_datetime", - "_dbm", - "_decimal", - "_elementtree", - "_frozen_importlib", - "_frozen_importlib_external", - "_functools", - "_gdbm", - "_hashlib", - "_heapq", - "_imp", - "_io", - "_json", - "_locale", - "_lsprof", - "_lzma", - "_markupbase", - "_md5", - "_msi", - "_multibytecodec", - "_multiprocessing", - "_opcode", - "_operator", - "_osx_support", - "_overlapped", - "_pickle", - "_posixshmem", - "_posixsubprocess", - "_py_abc", - "_pydecimal", - "_pyio", - "_queue", - "_random", - "_sha1", - "_sha256", - "_sha3", - "_sha512", - "_signal", - "_sitebuiltins", - "_socket", - "_sqlite3", - "_sre", - "_ssl", - "_stat", - "_statistics", - "_string", - "_strptime", - "_struct", - "_symtable", - "_thread", - "_threading_local", - "_tkinter", - "_tracemalloc", - "_uuid", - "_warnings", - "_weakref", - "_weakrefset", - "_winapi", - "_zoneinfo", - "abc", - "aifc", - "antigravity", - "argparse", - "array", - "ast", - "asynchat", - "asyncio", - "asyncore", - "atexit", - "audioop", - "base64", - "bdb", - "binascii", - "binhex", - "bisect", - "builtins", - "bz2", - "cProfile", - "calendar", - "cgi", - "cgitb", - "chunk", - "cmath", - "cmd", - "code", - "codecs", - "codeop", - "collections", - "colorsys", - "compileall", - "concurrent", - "configparser", - "contextlib", - "contextvars", - "copy", - "copyreg", - "crypt", - "csv", - "ctypes", - "curses", - "dataclasses", - "datetime", - "dbm", - "decimal", - "difflib", - "dis", - "distutils", - "doctest", - "email", - "encodings", - "ensurepip", - "enum", - "errno", - "faulthandler", - "fcntl", - "filecmp", - "fileinput", - "fnmatch", - "fractions", - "ftplib", - "functools", - "gc", - "genericpath", - "getopt", - "getpass", - "gettext", - "glob", - "graphlib", - "grp", - "gzip", - "hashlib", - "heapq", - "hmac", - "html", - "http", - "idlelib", - "imaplib", - "imghdr", - "imp", - "importlib", - "inspect", - "io", - "ipaddress", - "itertools", - "json", - "keyword", - "lib2to3", - "linecache", - "locale", - "logging", - "lzma", - "mailbox", - "mailcap", - "marshal", - "math", - "mimetypes", - "mmap", - "modulefinder", - "msilib", - "msvcrt", - "multiprocessing", - "netrc", - "nis", - "nntplib", - "nt", - "ntpath", - "nturl2path", - "numbers", - "opcode", - "operator", - "optparse", - "os", - "ossaudiodev", - "pathlib", - "pdb", - "pickle", - "pickletools", - "pipes", - "pkgutil", - "platform", - "plistlib", - "poplib", - "posix", - "posixpath", - "pprint", - "profile", - "pstats", - "pty", - "pwd", - "py_compile", - "pyclbr", - "pydoc", - "pydoc_data", - "pyexpat", - "queue", - "quopri", - "random", - "re", - "readline", - "reprlib", - "resource", - "rlcompleter", - "runpy", - "sched", - "secrets", - "select", - "selectors", - "shelve", - "shlex", - "shutil", - "signal", - "site", - "smtpd", - "smtplib", - "sndhdr", - "socket", - "socketserver", - "spwd", - "sqlite3", - "sre_compile", - "sre_constants", - "sre_parse", - "ssl", - "stat", - "statistics", - "string", - "stringprep", - "struct", - "subprocess", - "sunau", - "symtable", - "sys", - "sysconfig", - "syslog", - "tabnanny", - "tarfile", - "telnetlib", - "tempfile", - "termios", - "textwrap", - "this", - "threading", - "time", - "timeit", - "tkinter", - "token", - "tokenize", - "trace", - "traceback", - "tracemalloc", - "tty", - "turtle", - "turtledemo", - "types", - "typing", - "unicodedata", - "unittest", - "urllib", - "uu", - "uuid", - "venv", - "warnings", - "wave", - "weakref", - "webbrowser", - "winreg", - "winsound", - "wsgiref", - "xdrlib", - "xml", - "xmlrpc", - "zipapp", - "zipfile", - "zipimport", - "", -]; - -fn to_value(et: &ExpressionType) -> Option { - match et { - ExpressionType::String { value: StringGroup::Constant { value } } => Some(json!(value)), - ExpressionType::Number { value } => match value { - Number::Integer { value } => Some(json!(value.to_string().parse::().unwrap())), - Number::Float { value } => Some(json!(value)), - _ => None, - }, - ExpressionType::True => Some(json!(true)), - ExpressionType::False => Some(json!(false)), - - ExpressionType::Dict { elements } => { - let v = elements - .into_iter() - .map(|(k, v)| { - let key = k - .as_ref() - .and_then(|x| to_value(&x.node)) - .and_then(|x| match x { - serde_json::Value::String(s) => Some(s), - _ => None, - }) - .unwrap_or_else(|| "no_key".to_string()); - (key, to_value(&v.node)) - }) - .collect::>(); - Some(json!(v)) - } - ExpressionType::List { elements } => { - let v = elements - .into_iter() - .map(|x| to_value(&x.node)) - .collect::>(); - Some(json!(v)) - } - ExpressionType::None => Some(json!(null)), - - ExpressionType::Call { function: _, args: _, keywords: _ } => { - Some(json!("")) - } - - _ => None, - } -} - -static PYTHON_IMPORTS_REPLACEMENT: phf::Map<&'static str, &'static str> = phf_map! { - "psycopg2" => "psycopg2-binary" -}; - -fn replace_import(x: String) -> String { - PYTHON_IMPORTS_REPLACEMENT - .get(&x) - .map(|x| x.to_owned()) - .unwrap_or(&x) - .to_string() -} - -pub fn parse_python_imports(code: &str) -> error::Result> { - let find_requirements = code - .lines() - .find_position(|x| x.starts_with("#requirements:")); - let re = Regex::new(r"^\#(\S+)$").unwrap(); - if let Some((pos, _)) = find_requirements { - let lines = code - .lines() - .skip(pos + 1) - .map_while(|x| { - re.captures(x) - .map(|x| x.get(1).unwrap().as_str().to_string()) - }) - .collect(); - Ok(lines) - } else { - let ast = parser::parse_program(code) - .map_err(|e| { - error::Error::ExecutionErr(format!("Error parsing code: {}", e.to_string())) - })? - .statements; - let imports = ast - .into_iter() - .filter_map(|x| match x { - Located { location: _, node } => match node { - StatementType::Import { names } => Some( - names - .into_iter() - .map(|x| x.symbol.split('.').next().unwrap_or("").to_string()) - .map(replace_import) - .collect::>(), - ), - StatementType::ImportFrom { level: _, module: Some(mod_), names: _ } => { - let imprt = mod_.split('.').next().unwrap_or("").replace("_", "-"); - - Some(vec![replace_import(imprt)]) - } - _ => None, - }, - }) - .flatten() - .filter(|x| !STDIMPORTS.contains(&x.as_str())) - .unique() - .collect(); - Ok(imports) - } -} - -#[cfg(test)] -mod tests { - - // Note this useful idiom: importing names from outer (for mod tests) scope. - use super::*; - - #[test] - fn test_parse_python_sig() -> anyhow::Result<()> { - //let code = "print(2 + 3, fd=sys.stderr)"; - let code = " - -import os - -def main(test1: str, name: datetime.datetime = datetime.now(), byte: bytes = bytes(1)): - - print(f\"Hello World and a warm welcome especially to {name}\") - print(\"The env variable at `all/pretty_secret`: \", os.environ.get(\"ALL_PRETTY_SECRET\")) - return {\"len\": len(name), \"splitted\": name.split() } - -"; - println!("{}", serde_json::to_string(&parse_python_signature(code)?)?); - - Ok(()) - } - - #[test] - fn test_parse_python_imports() -> anyhow::Result<()> { - //let code = "print(2 + 3, fd=sys.stderr)"; - let code = " - -import os -import wmill -from zanzibar.estonie import talin -import matplotlib.pyplot as plt - -def main(): - pass - -"; - let r = parse_python_imports(code)?; - println!("{}", serde_json::to_string(&r)?); - assert_eq!(r, vec!["wmill", "zanzibar", "matplotlib"]); - Ok(()) - } - - #[test] - fn test_parse_python_imports2() -> anyhow::Result<()> { - //let code = "print(2 + 3, fd=sys.stderr)"; - let code = " -#requirements: -#burkina=0.4 -#nigeria -# -#congo - -import os -import wmill -from zanzibar.estonie import talin - -def main(): - pass - -"; - let r = parse_python_imports(code)?; - println!("{}", serde_json::to_string(&r)?); - assert_eq!(r, vec!["burkina=0.4", "nigeria"]); - - Ok(()) - } - - #[test] - fn test_parse_deno_sig() -> anyhow::Result<()> { - let code = " - -export function main(test1?: string, test2: string = \"burkina\", - test3: wmill.Resource<'postgres'>, b64: Base64, ls: Base64[], email: Email) { - console.log(42) -} - -"; - println!("{}", serde_json::to_string(&parse_deno_signature(code)?)?); - - Ok(()) - } -} diff --git a/backend/src/parser_py.rs b/backend/src/parser_py.rs new file mode 100644 index 0000000000..d6beffade7 --- /dev/null +++ b/backend/src/parser_py.rs @@ -0,0 +1,571 @@ +/* + * Author: Ruben Fiszel + * Copyright: Windmill Labs, Inc 2022 + * This file and its contents are licensed under the AGPLv3 License. + * Please see the included NOTICE for copyright information and + * LICENSE-AGPL for a copy of the license. + */ + +use std::collections::HashMap; + +use itertools::Itertools; +use phf::phf_map; +use regex::Regex; +use serde_json::json; + +use crate::{ + error, + parser::{Arg, InnerTyp, MainArgSignature, Typ}, +}; + +use rustpython_parser::{ + ast::{ExpressionType, Located, Number, StatementType, StringGroup, Varargs}, + parser, +}; + +pub fn parse_python_signature(code: &str) -> error::Result { + let ast = parser::parse_program(code) + .map_err(|e| error::Error::ExecutionErr(format!("Error parsing code: {}", e.to_string())))? + .statements; + let param = ast.into_iter().find_map(|x| match x { + Located { + location: _, + node: + StatementType::FunctionDef { + is_async: _, + name, + args, + body: _, + decorator_list: _, + returns: _, + }, + } if &name == "main" => Some(*args), + _ => None, + }); + if let Some(params) = param { + //println!("{:?}", params); + let def_arg_start = params.args.len() - params.defaults.len(); + Ok(MainArgSignature { + star_args: params.vararg != Varargs::None, + star_kwargs: params.vararg != Varargs::None, + args: params + .args + .into_iter() + .enumerate() + .map(|(i, x)| { + let default = if i >= def_arg_start { + to_value(¶ms.defaults[i - def_arg_start].node) + } else { + None + }; + Arg { + name: x.arg, + typ: x.annotation.map_or(Typ::Unknown, |e| match *e { + Located { location: _, node: ExpressionType::Identifier { name } } => { + match name.as_ref() { + "str" => Typ::Str(None), + "float" => Typ::Float, + "int" => Typ::Int, + "bool" => Typ::Bool, + "dict" => Typ::Dict, + "list" => Typ::List(InnerTyp::Str), + "bytes" => Typ::Bytes, + "datetime" => Typ::Datetime, + "datetime.datetime" => Typ::Datetime, + _ => Typ::Unknown, + } + } + _ => Typ::Unknown, + }), + has_default: default.is_some(), + default, + } + }) + .collect(), + }) + } else { + Err(error::Error::ExecutionErr( + "main function was not findable".to_string(), + )) + } +} + +fn to_value(et: &ExpressionType) -> Option { + match et { + ExpressionType::String { value: StringGroup::Constant { value } } => Some(json!(value)), + ExpressionType::Number { value } => match value { + Number::Integer { value } => Some(json!(value.to_string().parse::().unwrap())), + Number::Float { value } => Some(json!(value)), + _ => None, + }, + ExpressionType::True => Some(json!(true)), + ExpressionType::False => Some(json!(false)), + + ExpressionType::Dict { elements } => { + let v = elements + .into_iter() + .map(|(k, v)| { + let key = k + .as_ref() + .and_then(|x| to_value(&x.node)) + .and_then(|x| match x { + serde_json::Value::String(s) => Some(s), + _ => None, + }) + .unwrap_or_else(|| "no_key".to_string()); + (key, to_value(&v.node)) + }) + .collect::>(); + Some(json!(v)) + } + ExpressionType::List { elements } => { + let v = elements + .into_iter() + .map(|x| to_value(&x.node)) + .collect::>(); + Some(json!(v)) + } + ExpressionType::None => Some(json!(null)), + + ExpressionType::Call { function: _, args: _, keywords: _ } => { + Some(json!("")) + } + + _ => None, + } +} + +static PYTHON_IMPORTS_REPLACEMENT: phf::Map<&'static str, &'static str> = phf_map! { + "psycopg2" => "psycopg2-binary" +}; + +fn replace_import(x: String) -> String { + PYTHON_IMPORTS_REPLACEMENT + .get(&x) + .map(|x| x.to_owned()) + .unwrap_or(&x) + .to_string() +} + +pub fn parse_python_imports(code: &str) -> error::Result> { + let find_requirements = code + .lines() + .find_position(|x| x.starts_with("#requirements:")); + let re = Regex::new(r"^\#(\S+)$").unwrap(); + if let Some((pos, _)) = find_requirements { + let lines = code + .lines() + .skip(pos + 1) + .map_while(|x| { + re.captures(x) + .map(|x| x.get(1).unwrap().as_str().to_string()) + }) + .collect(); + Ok(lines) + } else { + let ast = parser::parse_program(code) + .map_err(|e| { + error::Error::ExecutionErr(format!("Error parsing code: {}", e.to_string())) + })? + .statements; + let imports = ast + .into_iter() + .filter_map(|x| match x { + Located { location: _, node } => match node { + StatementType::Import { names } => Some( + names + .into_iter() + .map(|x| x.symbol.split('.').next().unwrap_or("").to_string()) + .map(replace_import) + .collect::>(), + ), + StatementType::ImportFrom { level: _, module: Some(mod_), names: _ } => { + let imprt = mod_.split('.').next().unwrap_or("").replace("_", "-"); + + Some(vec![replace_import(imprt)]) + } + _ => None, + }, + }) + .flatten() + .filter(|x| !STDIMPORTS.contains(&x.as_str())) + .unique() + .collect(); + Ok(imports) + } +} + +#[cfg(test)] +mod tests { + + // Note this useful idiom: importing names from outer (for mod tests) scope. + use super::*; + + #[test] + fn test_parse_python_sig() -> anyhow::Result<()> { + //let code = "print(2 + 3, fd=sys.stderr)"; + let code = " + +import os + +def main(test1: str, name: datetime.datetime = datetime.now(), byte: bytes = bytes(1)): + + print(f\"Hello World and a warm welcome especially to {name}\") + print(\"The env variable at `all/pretty_secret`: \", os.environ.get(\"ALL_PRETTY_SECRET\")) + return {\"len\": len(name), \"splitted\": name.split() } + +"; + println!("{}", serde_json::to_string(&parse_python_signature(code)?)?); + + Ok(()) + } + + #[test] + fn test_parse_python_imports() -> anyhow::Result<()> { + //let code = "print(2 + 3, fd=sys.stderr)"; + let code = " + +import os +import wmill +from zanzibar.estonie import talin +import matplotlib.pyplot as plt + +def main(): + pass + +"; + let r = parse_python_imports(code)?; + println!("{}", serde_json::to_string(&r)?); + assert_eq!(r, vec!["wmill", "zanzibar", "matplotlib"]); + Ok(()) + } + + #[test] + fn test_parse_python_imports2() -> anyhow::Result<()> { + //let code = "print(2 + 3, fd=sys.stderr)"; + let code = " +#requirements: +#burkina=0.4 +#nigeria +# +#congo + +import os +import wmill +from zanzibar.estonie import talin + +def main(): + pass + +"; + let r = parse_python_imports(code)?; + println!("{}", serde_json::to_string(&r)?); + assert_eq!(r, vec!["burkina=0.4", "nigeria"]); + + Ok(()) + } +} + +const STDIMPORTS: [&str; 301] = [ + "__future__", + "_abc", + "_aix_support", + "_ast", + "_asyncio", + "_bisect", + "_blake2", + "_bootsubprocess", + "_bz2", + "_codecs", + "_codecs_cn", + "_codecs_hk", + "_codecs_iso2022", + "_codecs_jp", + "_codecs_kr", + "_codecs_tw", + "_collections", + "_collections_abc", + "_compat_pickle", + "_compression", + "_contextvars", + "_crypt", + "_csv", + "_ctypes", + "_curses", + "_curses_panel", + "_datetime", + "_dbm", + "_decimal", + "_elementtree", + "_frozen_importlib", + "_frozen_importlib_external", + "_functools", + "_gdbm", + "_hashlib", + "_heapq", + "_imp", + "_io", + "_json", + "_locale", + "_lsprof", + "_lzma", + "_markupbase", + "_md5", + "_msi", + "_multibytecodec", + "_multiprocessing", + "_opcode", + "_operator", + "_osx_support", + "_overlapped", + "_pickle", + "_posixshmem", + "_posixsubprocess", + "_py_abc", + "_pydecimal", + "_pyio", + "_queue", + "_random", + "_sha1", + "_sha256", + "_sha3", + "_sha512", + "_signal", + "_sitebuiltins", + "_socket", + "_sqlite3", + "_sre", + "_ssl", + "_stat", + "_statistics", + "_string", + "_strptime", + "_struct", + "_symtable", + "_thread", + "_threading_local", + "_tkinter", + "_tracemalloc", + "_uuid", + "_warnings", + "_weakref", + "_weakrefset", + "_winapi", + "_zoneinfo", + "abc", + "aifc", + "antigravity", + "argparse", + "array", + "ast", + "asynchat", + "asyncio", + "asyncore", + "atexit", + "audioop", + "base64", + "bdb", + "binascii", + "binhex", + "bisect", + "builtins", + "bz2", + "cProfile", + "calendar", + "cgi", + "cgitb", + "chunk", + "cmath", + "cmd", + "code", + "codecs", + "codeop", + "collections", + "colorsys", + "compileall", + "concurrent", + "configparser", + "contextlib", + "contextvars", + "copy", + "copyreg", + "crypt", + "csv", + "ctypes", + "curses", + "dataclasses", + "datetime", + "dbm", + "decimal", + "difflib", + "dis", + "distutils", + "doctest", + "email", + "encodings", + "ensurepip", + "enum", + "errno", + "faulthandler", + "fcntl", + "filecmp", + "fileinput", + "fnmatch", + "fractions", + "ftplib", + "functools", + "gc", + "genericpath", + "getopt", + "getpass", + "gettext", + "glob", + "graphlib", + "grp", + "gzip", + "hashlib", + "heapq", + "hmac", + "html", + "http", + "idlelib", + "imaplib", + "imghdr", + "imp", + "importlib", + "inspect", + "io", + "ipaddress", + "itertools", + "json", + "keyword", + "lib2to3", + "linecache", + "locale", + "logging", + "lzma", + "mailbox", + "mailcap", + "marshal", + "math", + "mimetypes", + "mmap", + "modulefinder", + "msilib", + "msvcrt", + "multiprocessing", + "netrc", + "nis", + "nntplib", + "nt", + "ntpath", + "nturl2path", + "numbers", + "opcode", + "operator", + "optparse", + "os", + "ossaudiodev", + "pathlib", + "pdb", + "pickle", + "pickletools", + "pipes", + "pkgutil", + "platform", + "plistlib", + "poplib", + "posix", + "posixpath", + "pprint", + "profile", + "pstats", + "pty", + "pwd", + "py_compile", + "pyclbr", + "pydoc", + "pydoc_data", + "pyexpat", + "queue", + "quopri", + "random", + "re", + "readline", + "reprlib", + "resource", + "rlcompleter", + "runpy", + "sched", + "secrets", + "select", + "selectors", + "shelve", + "shlex", + "shutil", + "signal", + "site", + "smtpd", + "smtplib", + "sndhdr", + "socket", + "socketserver", + "spwd", + "sqlite3", + "sre_compile", + "sre_constants", + "sre_parse", + "ssl", + "stat", + "statistics", + "string", + "stringprep", + "struct", + "subprocess", + "sunau", + "symtable", + "sys", + "sysconfig", + "syslog", + "tabnanny", + "tarfile", + "telnetlib", + "tempfile", + "termios", + "textwrap", + "this", + "threading", + "time", + "timeit", + "tkinter", + "token", + "tokenize", + "trace", + "traceback", + "tracemalloc", + "tty", + "turtle", + "turtledemo", + "types", + "typing", + "unicodedata", + "unittest", + "urllib", + "uu", + "uuid", + "venv", + "warnings", + "wave", + "weakref", + "webbrowser", + "winreg", + "winsound", + "wsgiref", + "xdrlib", + "xml", + "xmlrpc", + "zipapp", + "zipfile", + "zipimport", + "", +]; diff --git a/backend/src/parser_ts.rs b/backend/src/parser_ts.rs new file mode 100644 index 0000000000..388df3083a --- /dev/null +++ b/backend/src/parser_ts.rs @@ -0,0 +1,215 @@ +/* + * Author: Ruben Fiszel + * Copyright: Windmill Labs, Inc 2022 + * This file and its contents are licensed under the AGPLv3 License. + * Please see the included NOTICE for copyright information and + * LICENSE-AGPL for a copy of the license. + */ + +use crate::{ + error, + parser::{Arg, InnerTyp, MainArgSignature, Typ}, +}; + +use swc_common::{sync::Lrc, FileName, SourceMap}; +use swc_ecma_ast::{ + AssignPat, BindingIdent, Decl, ExportDecl, FnDecl, Ident, ModuleDecl, ModuleItem, Pat, Str, + TsArrayType, TsEntityName, TsKeywordTypeKind, TsLit, TsLitType, TsType, TsTypeRef, + TsUnionOrIntersectionType, TsUnionType, +}; +use swc_ecma_parser::{lexer::Lexer, Parser, StringInput, Syntax, TsConfig}; + +pub fn parse_deno_signature(code: &str) -> error::Result { + let cm: Lrc = Default::default(); + let fm = cm.new_source_file(FileName::Custom("test.ts".into()), code.into()); + let lexer = Lexer::new( + // We want to parse ecmascript + Syntax::Typescript(TsConfig::default()), + // EsVersion defaults to es5 + Default::default(), + StringInput::from(&*fm), + None, + ); + + let mut parser = Parser::new_from(lexer); + + let mut err_s = "".to_string(); + for e in parser.take_errors() { + err_s += &e.into_kind().msg().to_string(); + } + + let ast = parser + .parse_module() + .map_err(|_| { + error::Error::ExecutionErr(format!( + "Error while parsing code, it is invalid typescript" + )) + })? + .body; + + // println!("{ast:?}"); + let params = ast.into_iter().find_map(|x| match x { + ModuleItem::ModuleDecl(ModuleDecl::ExportDecl(ExportDecl { + decl: Decl::Fn(FnDecl { ident: Ident { sym, .. }, function, .. }), + .. + })) if &sym.to_string() == "main" => Some(function.params), + _ => None, + }); + if let Some(params) = params { + Ok(MainArgSignature { + star_args: false, + star_kwargs: false, + args: params + .into_iter() + .map(|x| match x.pat { + Pat::Ident(ident) => { + let (name, typ) = binding_ident_to_arg(&ident)?; + Ok(Arg { name, typ, default: None, has_default: ident.id.optional }) + } + Pat::Assign(AssignPat { left, right, .. }) => { + let (name, typ) = + left.as_ident().map(binding_ident_to_arg).ok_or_else(|| { + error::Error::ExecutionErr(format!( + "Arg {left:?} has unexpected syntax" + )) + })??; + Ok(Arg { + name, + typ, + default: serde_json::to_value(right) + .map_err(|e| error::Error::ExecutionErr(e.to_string()))? + .as_object() + .and_then(|x| x.get("value").to_owned()) + .cloned(), + + has_default: true, + }) + } + _ => Err(error::Error::ExecutionErr(format!( + "Arg {x:?} has unexpected syntax" + ))), + }) + .collect::, error::Error>>()?, + }) + } else { + Err(error::Error::ExecutionErr( + "main function was not findable (expected to find 'export main function(...)'" + .to_string(), + )) + } +} + +fn binding_ident_to_arg( + BindingIdent { id, type_ann }: &BindingIdent, +) -> anyhow::Result<(String, Typ)> { + Ok(( + id.sym.to_string(), + type_ann + .as_ref() + .map(|x| { + println!(""); + println!("{:?}", id.sym.to_string()); + println!("{:?}", x); + match &*x.type_ann { + TsType::TsKeywordType(t) => match t.kind { + TsKeywordTypeKind::TsObjectKeyword => Typ::Dict, + TsKeywordTypeKind::TsBooleanKeyword => Typ::Bool, + TsKeywordTypeKind::TsBigIntKeyword => Typ::Int, + TsKeywordTypeKind::TsNumberKeyword => Typ::Float, + TsKeywordTypeKind::TsStringKeyword => Typ::Str(None), + _ => Typ::Unknown, + }, + // TODO: we can do better here and extract the inner type of array + TsType::TsArrayType(TsArrayType { elem_type, .. }) => { + match &**elem_type { + TsType::TsTypeRef(TsTypeRef { + type_name: TsEntityName::Ident(Ident { sym, .. }), + .. + }) => match sym.to_string().as_str() { + "Base64" => Typ::List(InnerTyp::Bytes), + "Email" => Typ::List(InnerTyp::Email), + "bigint" => Typ::List(InnerTyp::Int), + "number" => Typ::List(InnerTyp::Float), + _ => Typ::List(InnerTyp::Str), + }, + //TsType::TsKeywordType(()) + _ => Typ::List(InnerTyp::Str), + } + } + TsType::TsLitType(TsLitType { lit: TsLit::Str(Str { value, .. }), .. }) => { + Typ::Str(Some(vec![value.to_string()])) + } + TsType::TsUnionOrIntersectionType(TsUnionOrIntersectionType::TsUnionType( + TsUnionType { types, .. }, + )) => { + let literals = types + .into_iter() + .map(|x| match &**x { + TsType::TsLitType(TsLitType { + lit: TsLit::Str(Str { value, .. }), + .. + }) => Some(value.to_string()), + _ => None, + }) + .collect::>(); + if literals.iter().find(|x| x.is_none()).is_some() { + Typ::Unknown + } else { + Typ::Str(Some(literals.into_iter().filter_map(|x| x).collect())) + } + } + TsType::TsTypeRef(TsTypeRef { type_name, type_params, .. }) => { + let sym = match type_name { + TsEntityName::Ident(Ident { sym, .. }) => sym, + TsEntityName::TsQualifiedName(p) => &*p.right.sym, + }; + match sym.to_string().as_str() { + "Resource" => Typ::Resource( + type_params + .as_ref() + .and_then(|x| { + x.params.get(0).and_then(|y| { + y.as_ts_lit_type().and_then(|z| { + z.lit + .as_str() + .map(|a| a.to_owned().value.to_string()) + }) + }) + }) + .unwrap_or_else(|| "unknown".to_string()), + ), + "Base64" => Typ::Bytes, + "Email" => Typ::Email, + "Sql" => Typ::Sql, + _ => Typ::Unknown, + } + } + _ => Typ::Unknown, + } + }) + .unwrap_or(Typ::Unknown), + )) +} + +#[cfg(test)] +mod tests { + + // Note this useful idiom: importing names from outer (for mod tests) scope. + use super::*; + + #[test] + fn test_parse_deno_sig() -> anyhow::Result<()> { + let code = " + +export function main(test1?: string, test2: string = \"burkina\", + test3: wmill.Resource<'postgres'>, b64: Base64, ls: Base64[], + email: Email, literal: \"test\", literal_union: \"test\" | \"test2\") { + console.log(42) +} + +"; + println!("{}", serde_json::to_string(&parse_deno_signature(code)?)?); + + Ok(()) + } +} diff --git a/backend/src/scripts.rs b/backend/src/scripts.rs index 769d3b842b..471b4f1446 100644 --- a/backend/src/scripts.rs +++ b/backend/src/scripts.rs @@ -14,7 +14,7 @@ use crate::{ audit::{audit_log, ActionKind}, db::{UserDB, DB}, error::{to_anyhow, Error, JsonResult, Result}, - jobs, parser, + jobs, parser, parser_py, parser_ts, users::{owner_to_token_owner, truncate_token, Authed, Tokened}, utils::{http_get_from_hub, list_elems_from_hub, require_admin, Pagination, StripPath}, }; @@ -426,7 +426,7 @@ async fn create_script( .await?; let mut tx = if ns.lock.is_none() && ns.language == ScriptLang::Python3 { - let dependencies = parser::parse_python_imports(&ns.content)?; + let dependencies = parser_py::parse_python_imports(&ns.content)?; let (_, tx) = jobs::push( tx, &w_id, @@ -713,13 +713,13 @@ async fn delete_script_by_hash( async fn parse_python_code_to_jsonschema( Json(code): Json, ) -> JsonResult { - parser::parse_python_signature(&code).map(Json) + parser_py::parse_python_signature(&code).map(Json) } async fn parse_deno_code_to_jsonschema( Json(code): Json, ) -> JsonResult { - parser::parse_deno_signature(&code).map(Json) + parser_ts::parse_deno_signature(&code).map(Json) } pub fn to_i64(s: &str) -> Result { diff --git a/backend/src/worker.rs b/backend/src/worker.rs index 6a0c7e773f..6356e04fda 100644 --- a/backend/src/worker.rs +++ b/backend/src/worker.rs @@ -24,7 +24,8 @@ use crate::{ add_completed_job, add_completed_job_error, get_queued_job, postprocess_queued_job, pull, JobKind, QueuedJob, }, - parser::{self, Typ}, + parser::Typ, + parser_py, scripts::{ScriptHash, ScriptLang}, users::{create_token_for_owner, get_email_from_username}, variables, @@ -497,7 +498,7 @@ async fn handle_nondep_job( .map(|x| matches!(x, ScriptLang::Python3)) .unwrap_or(false) { - Some(parser::parse_python_imports(&code)?.join("\n")) + Some(parser_py::parse_python_imports(&code)?.join("\n")) } else { None }; @@ -588,7 +589,7 @@ async fn handle_nondep_job( let _ = write_file(job_dir, "inner.py", &inner_content).await?; - let sig = crate::parser::parse_python_signature(&inner_content)?; + let sig = crate::parser_py::parse_python_signature(&inner_content)?; let transforms = sig .args .into_iter() @@ -729,7 +730,7 @@ print(res_json) let _ = write_file(job_dir, "inner.ts", &inner_content).await?; - let sig = crate::parser::parse_deno_signature(&inner_content)?; + let sig = crate::parser_ts::parse_deno_signature(&inner_content)?; // let transforms = sig.args.clone().into_iter().map(|x| match x.typ { // Typ::Bytes => format!("if \"{}\" in kwargs and kwargs[\"{}\"] is not None:\n kwargs[\"{}\"] = base64.b64decode(kwargs[\"{}\"])\n", x.name, x.name, x.name, x.name), // Typ::Datetime => format!("if \"{}\" in kwargs and kwargs[\"{}\"] is not None:\n kwargs[\"{}\"] = datetime.strptime(kwargs[\"{}\"], '%Y-%m-%dT%H:%M')\n", x.name, x.name, x.name, x.name), diff --git a/frontend/src/lib/infer.ts b/frontend/src/lib/infer.ts index d1a839a834..cd7d40dc9b 100644 --- a/frontend/src/lib/infer.ts +++ b/frontend/src/lib/infer.ts @@ -39,17 +39,21 @@ export async function inferArgs( } function argSigToJsonSchemaType( - t: string | { resource: string } | { list: string }, + t: string | { resource: string | null } | { list: string | null } | { str: string[] | null }, s: SchemaProperty ): void { + for (const prop of Object.getOwnPropertyNames(s)) { + if (prop != "description") { + delete s[prop] + } + } + if (t === 'int') { s.type = 'integer' } else if (t === 'float') { s.type = 'number' } else if (t === 'bool') { s.type = 'boolean' - } else if (t === 'str') { - s.type = 'string' } else if (t === 'email') { s.type = 'string' s.format = 'email' @@ -64,6 +68,11 @@ function argSigToJsonSchemaType( } else if (t === 'datetime') { s.type = 'string' s.format = 'date-time' + } else if (typeof t !== 'string' && `str` in t) { + s.type = 'string' + if (t.str) { + s.enum = t.str + } } else if (typeof t !== 'string' && `resource` in t) { s.type = 'object' s.format = `resource-${t.resource}`