Compare commits
16
Commits
1d890395cc
...
develop
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cba96e4f46 | ||
|
|
d63b4ea470 | ||
|
|
b12785ed93 | ||
|
|
3c62970967 | ||
|
|
6c43ac9969 | ||
|
|
6d87da90d1 | ||
|
|
1fd7a4c698 | ||
|
|
71433b867e | ||
|
|
13e4b16b32 | ||
|
|
c281248bf8 | ||
|
|
a2f53d7879 | ||
|
|
16fda38039 | ||
|
|
5691b09fc8 | ||
|
|
33f1c218f2 | ||
|
|
81107c6f5c | ||
|
|
3b5c7e2d46 |
@@ -1,6 +1,7 @@
|
||||
# llm-worker-rs 開発instruction
|
||||
# llm-worker-rs Development Instructions
|
||||
|
||||
## パッケージ管理ルール
|
||||
## Package Management Rules
|
||||
|
||||
- クレートに依存関係を追加・更新する際は必ず
|
||||
`cargo`コマンドを使い、`Cargo.toml`を直接手で書き換えず、必ずコマンド経由で管理すること。
|
||||
- When adding or updating crate dependencies, always use the `cargo` command. Do
|
||||
not manually edit `Cargo.toml` directly; always manage dependencies via
|
||||
commands.
|
||||
|
||||
Generated
+355
-44
@@ -104,9 +104,9 @@ checksum = "b35204fbdc0b3f4446b89fc1ac2cf84a8a68971995d0bf2e925ec7cd960f9cb3"
|
||||
|
||||
[[package]]
|
||||
name = "cc"
|
||||
version = "1.2.51"
|
||||
version = "1.2.52"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7a0aeaff4ff1a90589618835a598e545176939b97874f7abc7851caa0618f203"
|
||||
checksum = "cd4932aefd12402b36c60956a4fe0035421f544799057659ff86f923657aada3"
|
||||
dependencies = [
|
||||
"find-msvc-tools",
|
||||
"shlex",
|
||||
@@ -203,6 +203,12 @@ version = "1.0.20"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d0881ea181b1df73ff77ffaaf9c7544ecc11e82fba9b5f27b262a3c73a332555"
|
||||
|
||||
[[package]]
|
||||
name = "equivalent"
|
||||
version = "1.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f"
|
||||
|
||||
[[package]]
|
||||
name = "errno"
|
||||
version = "0.3.14"
|
||||
@@ -232,9 +238,15 @@ checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be"
|
||||
|
||||
[[package]]
|
||||
name = "find-msvc-tools"
|
||||
version = "0.1.6"
|
||||
version = "0.1.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "645cbb3a84e60b7531617d5ae4e57f7e27308f6445f5abf653209ea76dec8dff"
|
||||
checksum = "f449e6c6c08c865631d4890cfacf252b3d396c9bcc83adb6623cdb02a8336c41"
|
||||
|
||||
[[package]]
|
||||
name = "fnv"
|
||||
version = "1.0.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1"
|
||||
|
||||
[[package]]
|
||||
name = "foreign-types"
|
||||
@@ -349,6 +361,17 @@ dependencies = [
|
||||
"slab",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "getrandom"
|
||||
version = "0.2.16"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "335ff9f135e4384c8150d6f27c6daed433577f86b4750418338c01a1a2528592"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"libc",
|
||||
"wasi",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "getrandom"
|
||||
version = "0.3.4"
|
||||
@@ -361,6 +384,37 @@ dependencies = [
|
||||
"wasip2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "glob"
|
||||
version = "0.3.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280"
|
||||
|
||||
[[package]]
|
||||
name = "h2"
|
||||
version = "0.4.13"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2f44da3a8150a6703ed5d34e164b875fd14c2cdab9af1252a9a1020bde2bdc54"
|
||||
dependencies = [
|
||||
"atomic-waker",
|
||||
"bytes",
|
||||
"fnv",
|
||||
"futures-core",
|
||||
"futures-sink",
|
||||
"http",
|
||||
"indexmap",
|
||||
"slab",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "hashbrown"
|
||||
version = "0.16.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100"
|
||||
|
||||
[[package]]
|
||||
name = "heck"
|
||||
version = "0.5.0"
|
||||
@@ -416,6 +470,7 @@ dependencies = [
|
||||
"bytes",
|
||||
"futures-channel",
|
||||
"futures-core",
|
||||
"h2",
|
||||
"http",
|
||||
"http-body",
|
||||
"httparse",
|
||||
@@ -427,6 +482,22 @@ dependencies = [
|
||||
"want",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "hyper-rustls"
|
||||
version = "0.27.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e3c93eb611681b207e1fe55d5a71ecf91572ec8a6705cdb6857f7d8d5242cf58"
|
||||
dependencies = [
|
||||
"http",
|
||||
"hyper",
|
||||
"hyper-util",
|
||||
"rustls",
|
||||
"rustls-pki-types",
|
||||
"tokio",
|
||||
"tokio-rustls",
|
||||
"tower-service",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "hyper-tls"
|
||||
version = "0.6.0"
|
||||
@@ -569,6 +640,16 @@ dependencies = [
|
||||
"icu_properties",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "indexmap"
|
||||
version = "2.13.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7714e70437a7dc3ac8eb7e6f8df75fd8eb422675fc7678aff7364301092b1017"
|
||||
dependencies = [
|
||||
"equivalent",
|
||||
"hashbrown",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ipnet"
|
||||
version = "2.11.0"
|
||||
@@ -631,6 +712,38 @@ version = "0.8.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6373607a59f0be73a39b6fe456b8192fcc3585f602af20751600e974dd455e77"
|
||||
|
||||
[[package]]
|
||||
name = "llm-worker"
|
||||
version = "0.2.1"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"clap",
|
||||
"dotenv",
|
||||
"eventsource-stream",
|
||||
"futures",
|
||||
"llm-worker-macros",
|
||||
"reqwest",
|
||||
"schemars",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tempfile",
|
||||
"thiserror",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
"trybuild",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "llm-worker-macros"
|
||||
version = "0.2.0"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "log"
|
||||
version = "0.4.29"
|
||||
@@ -865,10 +978,12 @@ dependencies = [
|
||||
"bytes",
|
||||
"futures-core",
|
||||
"futures-util",
|
||||
"h2",
|
||||
"http",
|
||||
"http-body",
|
||||
"http-body-util",
|
||||
"hyper",
|
||||
"hyper-rustls",
|
||||
"hyper-tls",
|
||||
"hyper-util",
|
||||
"js-sys",
|
||||
@@ -893,6 +1008,20 @@ dependencies = [
|
||||
"web-sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ring"
|
||||
version = "0.17.14"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7"
|
||||
dependencies = [
|
||||
"cc",
|
||||
"cfg-if",
|
||||
"getrandom 0.2.16",
|
||||
"libc",
|
||||
"untrusted",
|
||||
"windows-sys 0.52.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustix"
|
||||
version = "1.1.3"
|
||||
@@ -906,6 +1035,19 @@ dependencies = [
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustls"
|
||||
version = "0.23.36"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c665f33d38cea657d9614f766881e4d510e0eda4239891eea56b4cadcf01801b"
|
||||
dependencies = [
|
||||
"once_cell",
|
||||
"rustls-pki-types",
|
||||
"rustls-webpki",
|
||||
"subtle",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustls-pki-types"
|
||||
version = "1.13.2"
|
||||
@@ -915,6 +1057,17 @@ dependencies = [
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustls-webpki"
|
||||
version = "0.103.8"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2ffdfa2f5286e2247234e03f680868ac2815974dc39e00ea15adc445d0aafe52"
|
||||
dependencies = [
|
||||
"ring",
|
||||
"rustls-pki-types",
|
||||
"untrusted",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustversion"
|
||||
version = "1.0.22"
|
||||
@@ -1032,6 +1185,15 @@ dependencies = [
|
||||
"zmij",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "serde_spanned"
|
||||
version = "1.0.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f8bbf91e5a4d6315eee45e704372590b30e260ee83af6639d64557f51b067776"
|
||||
dependencies = [
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sharded-slab"
|
||||
version = "0.1.7"
|
||||
@@ -1081,6 +1243,12 @@ version = "0.11.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f"
|
||||
|
||||
[[package]]
|
||||
name = "subtle"
|
||||
version = "2.6.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292"
|
||||
|
||||
[[package]]
|
||||
name = "syn"
|
||||
version = "2.0.114"
|
||||
@@ -1112,6 +1280,12 @@ dependencies = [
|
||||
"syn",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "target-triple"
|
||||
version = "1.0.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "591ef38edfb78ca4771ee32cf494cb8771944bee237a9b91fc9c1424ac4b777b"
|
||||
|
||||
[[package]]
|
||||
name = "tempfile"
|
||||
version = "3.24.0"
|
||||
@@ -1119,12 +1293,21 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "655da9c7eb6305c55742045d5a8d2037996d61d8de95806335c7c86ce0f82e9c"
|
||||
dependencies = [
|
||||
"fastrand",
|
||||
"getrandom",
|
||||
"getrandom 0.3.4",
|
||||
"once_cell",
|
||||
"rustix",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "termcolor"
|
||||
version = "1.4.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "06794f8f6c5c898b3275aebefa6b8a1cb24cd2c6c79397ab15774837a0bc5755"
|
||||
dependencies = [
|
||||
"winapi-util",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "thiserror"
|
||||
version = "2.0.17"
|
||||
@@ -1200,6 +1383,16 @@ dependencies = [
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokio-rustls"
|
||||
version = "0.26.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61"
|
||||
dependencies = [
|
||||
"rustls",
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokio-util"
|
||||
version = "0.7.18"
|
||||
@@ -1213,6 +1406,45 @@ dependencies = [
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml"
|
||||
version = "1.0.3+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c7614eaf19ad818347db24addfa201729cf2a9b6fdfd9eb0ab870fcacc606c0c"
|
||||
dependencies = [
|
||||
"indexmap",
|
||||
"serde_core",
|
||||
"serde_spanned",
|
||||
"toml_datetime",
|
||||
"toml_parser",
|
||||
"toml_writer",
|
||||
"winnow",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml_datetime"
|
||||
version = "1.0.0+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "32c2555c699578a4f59f0cc68e5116c8d7cabbd45e1409b989d4be085b53f13e"
|
||||
dependencies = [
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml_parser"
|
||||
version = "1.0.9+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "702d4415e08923e7e1ef96cd5727c0dfed80b4d2fa25db9647fe5eb6f7c5a4c4"
|
||||
dependencies = [
|
||||
"winnow",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "toml_writer"
|
||||
version = "1.0.6+spec-1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ab16f14aed21ee8bfd8ec22513f7287cd4a91aa92e44edfe2c17ddd004e92607"
|
||||
|
||||
[[package]]
|
||||
name = "tower"
|
||||
version = "0.5.2"
|
||||
@@ -1325,12 +1557,33 @@ version = "0.2.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b"
|
||||
|
||||
[[package]]
|
||||
name = "trybuild"
|
||||
version = "1.0.116"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "47c635f0191bd3a2941013e5062667100969f8c4e9cd787c14f977265d73616e"
|
||||
dependencies = [
|
||||
"glob",
|
||||
"serde",
|
||||
"serde_derive",
|
||||
"serde_json",
|
||||
"target-triple",
|
||||
"termcolor",
|
||||
"toml",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "unicode-ident"
|
||||
version = "1.0.22"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9312f7c4f6ff9069b165498234ce8be658059c6728633667c526e27dc2cf1df5"
|
||||
|
||||
[[package]]
|
||||
name = "untrusted"
|
||||
version = "0.9.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1"
|
||||
|
||||
[[package]]
|
||||
name = "url"
|
||||
version = "2.5.8"
|
||||
@@ -1472,19 +1725,37 @@ dependencies = [
|
||||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "winapi-util"
|
||||
version = "0.1.11"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22"
|
||||
dependencies = [
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "windows-link"
|
||||
version = "0.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
|
||||
|
||||
[[package]]
|
||||
name = "windows-sys"
|
||||
version = "0.52.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d"
|
||||
dependencies = [
|
||||
"windows-targets 0.52.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "windows-sys"
|
||||
version = "0.60.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb"
|
||||
dependencies = [
|
||||
"windows-targets",
|
||||
"windows-targets 0.53.5",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -1496,6 +1767,22 @@ dependencies = [
|
||||
"windows-link",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "windows-targets"
|
||||
version = "0.52.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973"
|
||||
dependencies = [
|
||||
"windows_aarch64_gnullvm 0.52.6",
|
||||
"windows_aarch64_msvc 0.52.6",
|
||||
"windows_i686_gnu 0.52.6",
|
||||
"windows_i686_gnullvm 0.52.6",
|
||||
"windows_i686_msvc 0.52.6",
|
||||
"windows_x86_64_gnu 0.52.6",
|
||||
"windows_x86_64_gnullvm 0.52.6",
|
||||
"windows_x86_64_msvc 0.52.6",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "windows-targets"
|
||||
version = "0.53.5"
|
||||
@@ -1503,100 +1790,124 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3"
|
||||
dependencies = [
|
||||
"windows-link",
|
||||
"windows_aarch64_gnullvm",
|
||||
"windows_aarch64_msvc",
|
||||
"windows_i686_gnu",
|
||||
"windows_i686_gnullvm",
|
||||
"windows_i686_msvc",
|
||||
"windows_x86_64_gnu",
|
||||
"windows_x86_64_gnullvm",
|
||||
"windows_x86_64_msvc",
|
||||
"windows_aarch64_gnullvm 0.53.1",
|
||||
"windows_aarch64_msvc 0.53.1",
|
||||
"windows_i686_gnu 0.53.1",
|
||||
"windows_i686_gnullvm 0.53.1",
|
||||
"windows_i686_msvc 0.53.1",
|
||||
"windows_x86_64_gnu 0.53.1",
|
||||
"windows_x86_64_gnullvm 0.53.1",
|
||||
"windows_x86_64_msvc 0.53.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "windows_aarch64_gnullvm"
|
||||
version = "0.52.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3"
|
||||
|
||||
[[package]]
|
||||
name = "windows_aarch64_gnullvm"
|
||||
version = "0.53.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a9d8416fa8b42f5c947f8482c43e7d89e73a173cead56d044f6a56104a6d1b53"
|
||||
|
||||
[[package]]
|
||||
name = "windows_aarch64_msvc"
|
||||
version = "0.52.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469"
|
||||
|
||||
[[package]]
|
||||
name = "windows_aarch64_msvc"
|
||||
version = "0.53.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006"
|
||||
|
||||
[[package]]
|
||||
name = "windows_i686_gnu"
|
||||
version = "0.52.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b"
|
||||
|
||||
[[package]]
|
||||
name = "windows_i686_gnu"
|
||||
version = "0.53.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "960e6da069d81e09becb0ca57a65220ddff016ff2d6af6a223cf372a506593a3"
|
||||
|
||||
[[package]]
|
||||
name = "windows_i686_gnullvm"
|
||||
version = "0.52.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66"
|
||||
|
||||
[[package]]
|
||||
name = "windows_i686_gnullvm"
|
||||
version = "0.53.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c"
|
||||
|
||||
[[package]]
|
||||
name = "windows_i686_msvc"
|
||||
version = "0.52.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66"
|
||||
|
||||
[[package]]
|
||||
name = "windows_i686_msvc"
|
||||
version = "0.53.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1e7ac75179f18232fe9c285163565a57ef8d3c89254a30685b57d83a38d326c2"
|
||||
|
||||
[[package]]
|
||||
name = "windows_x86_64_gnu"
|
||||
version = "0.52.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78"
|
||||
|
||||
[[package]]
|
||||
name = "windows_x86_64_gnu"
|
||||
version = "0.53.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499"
|
||||
|
||||
[[package]]
|
||||
name = "windows_x86_64_gnullvm"
|
||||
version = "0.52.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d"
|
||||
|
||||
[[package]]
|
||||
name = "windows_x86_64_gnullvm"
|
||||
version = "0.53.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0ffa179e2d07eee8ad8f57493436566c7cc30ac536a3379fdf008f47f6bb7ae1"
|
||||
|
||||
[[package]]
|
||||
name = "windows_x86_64_msvc"
|
||||
version = "0.52.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec"
|
||||
|
||||
[[package]]
|
||||
name = "windows_x86_64_msvc"
|
||||
version = "0.53.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650"
|
||||
|
||||
[[package]]
|
||||
name = "winnow"
|
||||
version = "0.7.14"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5a5364e9d77fcdeeaa6062ced926ee3381faa2ee02d3eb83a5c27a8825540829"
|
||||
|
||||
[[package]]
|
||||
name = "wit-bindgen"
|
||||
version = "0.46.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f17a85883d4e6d00e8a97c586de764dabcc06133f7f1d55dce5cdc070ad7fe59"
|
||||
|
||||
[[package]]
|
||||
name = "worker"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"clap",
|
||||
"dotenv",
|
||||
"eventsource-stream",
|
||||
"futures",
|
||||
"reqwest",
|
||||
"schemars",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tempfile",
|
||||
"thiserror",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
"worker-macros",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "worker-macros"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "writeable"
|
||||
version = "0.6.2"
|
||||
|
||||
+2
-2
@@ -1,8 +1,8 @@
|
||||
[workspace]
|
||||
resolver = "2"
|
||||
members = [
|
||||
"worker",
|
||||
"worker-macros",
|
||||
"llm-worker",
|
||||
"llm-worker-macros",
|
||||
]
|
||||
|
||||
[workspace.package]
|
||||
|
||||
@@ -1,5 +1,35 @@
|
||||
# llm-worker-rs
|
||||
# llm-worker
|
||||
|
||||
Rusty, Efficient, and Agentic LLM Client Library
|
||||
|
||||
`llm-worker-rs` is a Rust library designed for building robust and efficient LLM applications. It unifies interactions with multiple LLM providers (Anthropic, Gemini, OpenAI, Ollama) under a single abstraction, with type-safe state management, efficient context caching, and a powerful event-driven architecture.
|
||||
`llm-worker` is a Rust library for building autonomous LLM-powered systems. Define tools, register hooks, and let the Worker handle the agentic loop — tool calls are executed automatically until the task completes.
|
||||
|
||||
## Features
|
||||
|
||||
- Autonomous Execution: The `Worker` manages the full request-response-tool cycle. You provide tools and input; it loops until done.
|
||||
- Multi-Provider Support: Unified interface for Anthropic, Gemini, OpenAI, and Ollama.
|
||||
- Tool System: Define tools as async functions. The Worker automatically parses LLM tool calls, executes them in parallel, and feeds results back.
|
||||
- Hook System: Intercept execution flow with `before_tool_call`, `after_tool_call`, and `on_turn_end` hooks for validation, logging, or self-correction.
|
||||
- Event-Driven Streaming: Subscribe to real-time events (text deltas, tool calls, usage) for responsive UIs.
|
||||
- Cache-Aware State Management: Type-state pattern (`Mutable` → `CacheLocked`) ensures KV cache efficiency by protecting the conversation prefix.
|
||||
|
||||
## Quick Start
|
||||
|
||||
```rust
|
||||
use llm_worker::{Worker, Message};
|
||||
|
||||
// Create a Worker with your LLM client
|
||||
let mut worker = Worker::new(client)
|
||||
.system_prompt("You are a helpful assistant.");
|
||||
|
||||
// Register tools (optional)
|
||||
worker.register_tool(SearchTool::new());
|
||||
worker.register_tool(CalculatorTool::new());
|
||||
|
||||
// Run — the Worker handles tool calls automatically
|
||||
let history = worker.run("What is 2+2?").await?;
|
||||
```
|
||||
|
||||
## License
|
||||
|
||||
MIT
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
# Worker API/DSL 実装計画
|
||||
|
||||
## 目的
|
||||
|
||||
- [Open Responses](https://www.openresponses.org)(以後"OR")に準拠した正規化を前提に、
|
||||
Item/Part の2段スコープを扱える Worker API を設計する。
|
||||
- APIの煩雑化を防ぐため、worker.on_xxx として公開するのを避けつつ、
|
||||
Text/Thinking/Tool など型の違いを静的に扱える DSL を提供する。
|
||||
|
||||
## 方針
|
||||
|
||||
- 内部は Timeline が Event を正規化し、Item/Part/Meta
|
||||
を単一ストリームとして扱う。
|
||||
- API では Item/Part 型ごとに ctx を持てるようにし、DSL
|
||||
で記述の冗長さを削減する。
|
||||
- まず macro_rules! 版を作り、必要なら proc-macro に拡張する。
|
||||
- Item/Part の型パラメータはクレートが公開する Kind 型を使う。
|
||||
|
||||
## 仕様の前提
|
||||
|
||||
- Item は OR の item (message, function_call, reasoning など) に対応する。
|
||||
- Part は OR の content part (output_text, reasoning_text など) に対応する。
|
||||
- Item は必ず start/stop を持つ。Part は Item 内で複数発生し得る。
|
||||
- Item/Part の型指定は `Item<Message>` / `Part<ReasoningText>` のように書く。
|
||||
|
||||
## 設計ステップ
|
||||
|
||||
### 1. 内部イベントモデルの整理
|
||||
|
||||
- Event を Item/Part/Meta の3層に整理する。
|
||||
- ItemEvent / PartEvent は型パラメータで区別する。
|
||||
- 例: ItemEvent<Message>, PartEvent<Message, OutputText>
|
||||
|
||||
### 2. スコープの二段化
|
||||
|
||||
- Item ctx: Item 型ごとに1つ
|
||||
- Part ctx: Part 型ごとに1つ
|
||||
- Part のイベントでは常に item ctx と part ctx の両方を渡す。
|
||||
|
||||
### 3. Handler trait の再定義
|
||||
|
||||
- Item/Part を型で指定できる trait を導入する。
|
||||
- 例:
|
||||
- trait ItemHandler<I>
|
||||
- trait PartHandler<I, P>
|
||||
- PartHandler には ItemHandler の ItemCtx を必須で渡す。
|
||||
- Part の ctx 型は `PartKind::Ctx` 方式 or enum 方式で切り替える。
|
||||
|
||||
### 4. Timeline との結合
|
||||
|
||||
- Timeline は ItemStart で ItemCtx を生成
|
||||
- PartStart で PartCtx を生成
|
||||
- Delta/Stop は対応 ctx に流す
|
||||
- ItemStop で ItemCtx を破棄
|
||||
|
||||
### 5. DSL (macro_rules!) の導入
|
||||
|
||||
- まず宣言的 DSL を提供する。
|
||||
- 例:
|
||||
- handler! { Item<Message> { type ItemCtx = ...; Part<OutputText> { type
|
||||
PartCtx = ...; } } }
|
||||
- DSL は ItemHandler / PartHandler 実装を生成する。
|
||||
- Item/Part の Kind 型はクレートが公開する型を参照する。
|
||||
|
||||
### 6. 拡張ポイント
|
||||
|
||||
- 追加 Part (output_image など) を DSL に追加しやすい形にする。
|
||||
- 必要なら proc-macro に移行して構文自由度を上げる。
|
||||
|
||||
## 実装順序
|
||||
|
||||
1. Event/Item/Part の型定義の整理
|
||||
2. Item/Part ctx を持つ Timeline 実装
|
||||
3. Handler trait の定義・既存コードの移行
|
||||
4. macro_rules! DSL の実装
|
||||
5. 既存ユースケースの移植
|
||||
|
||||
## TODO
|
||||
|
||||
- Item と Part の型対応表を整理する
|
||||
- OR と既存 llm_client の差分を再確認する
|
||||
- Tool args の delta を OR 拡張として扱うか検討する
|
||||
- macro_rules! で表現可能な DSL の最小文法を確定する
|
||||
@@ -0,0 +1,80 @@
|
||||
# Open Responses mapping (llm_client -> Open Responses)
|
||||
|
||||
This document maps the current `llm_client` event model to Open Responses items
|
||||
and streaming events. It focuses on output streaming; input items are noted
|
||||
where they are the closest semantic match.
|
||||
|
||||
## Legend
|
||||
|
||||
- **OR item**: Open Responses item types used in `response.output`.
|
||||
- **OR event**: Open Responses streaming events (`response.*`).
|
||||
- **Note**: Gaps or required adaptation decisions.
|
||||
|
||||
## Response lifecycle / meta events
|
||||
|
||||
| llm_client | Open Responses | Note |
|
||||
| ------------------------ | ------------------------------------------------------------- | ---------------------------------------------------------------------------------------------- |
|
||||
| `StatusEvent::Started` | `response.created`, `response.queued`, `response.in_progress` | OR has finer-grained lifecycle states; pick a subset or map Started -> `response.in_progress`. |
|
||||
| `StatusEvent::Completed` | `response.completed` | |
|
||||
| `StatusEvent::Failed` | `response.failed` | |
|
||||
| `StatusEvent::Cancelled` | (no direct event) | Could map to `response.incomplete` or `response.failed` depending on semantics. |
|
||||
| `UsageEvent` | `response.completed` payload usage | OR reports usage on the response object, not as a dedicated streaming event. |
|
||||
| `ErrorEvent` | `error` event | OR has a dedicated error streaming event. |
|
||||
| `PingEvent` | (no direct event) | OR does not define a heartbeat event. |
|
||||
|
||||
## Output block lifecycle
|
||||
|
||||
### Text block
|
||||
|
||||
| llm_client | Open Responses | Note |
|
||||
| ------------------------------------------------- | ---------------------------------------------------------------------------------------- | ----------------------------------------------------------------------------------- |
|
||||
| `BlockStart { block_type: Text, metadata: Text }` | `response.output_item.added` with item type `message` (assistant) | OR output items are message/function_call/reasoning. This creates the message item. |
|
||||
| `BlockDelta { delta: Text(..) }` | `response.output_text.delta` | Text deltas map 1:1 to output text deltas. |
|
||||
| `BlockStop { block_type: Text }` | `response.output_text.done` + `response.content_part.done` + `response.output_item.done` | OR emits separate done events for content parts and items. |
|
||||
|
||||
### Tool use (function call)
|
||||
|
||||
| llm_client | Open Responses | Note |
|
||||
| -------------------------------------------------------------------- | --------------------------------------------------------------------- | ----------------------------------------------------------------------------------------------------- |
|
||||
| `BlockStart { block_type: ToolUse, metadata: ToolUse { id, name } }` | `response.output_item.added` with item type `function_call` | OR uses `call_id` + `name` + `arguments` string. Map `id` -> `call_id`. |
|
||||
| `BlockDelta { delta: InputJson(..) }` | `response.function_call_arguments.delta` | OR spec does not explicitly require argument deltas; treat as OpenAI-compatible extension if adopted. |
|
||||
| `BlockStop { block_type: ToolUse }` | `response.function_call_arguments.done` + `response.output_item.done` | Item status can be set to `completed` or `incomplete`. |
|
||||
|
||||
### Tool result (function call output)
|
||||
|
||||
| llm_client | Open Responses | Note |
|
||||
| ----------------------------------------------------------------------------- | ------------------------------------- | ---------------------------------------------------------------------------------------- |
|
||||
| `BlockStart { block_type: ToolResult, metadata: ToolResult { tool_use_id } }` | **Input item** `function_call_output` | OR treats tool results as input items, not output items. This is a request-side mapping. |
|
||||
| `BlockDelta` | (no direct output event) | OR does not stream tool output deltas as response events. |
|
||||
| `BlockStop` | (no direct output event) | Tool output lives on the next request as an input item. |
|
||||
|
||||
### Thinking / reasoning
|
||||
|
||||
| llm_client | Open Responses | Note |
|
||||
| --------------------------------------------------------- | ------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------- |
|
||||
| `BlockStart { block_type: Thinking, metadata: Thinking }` | `response.output_item.added` with item type `reasoning` | OR models reasoning as a separate item type. |
|
||||
| `BlockDelta { delta: Thinking(..) }` | `response.reasoning.delta` | OR has dedicated reasoning delta events. |
|
||||
| `BlockStop { block_type: Thinking }` | `response.reasoning.done` | OR separates reasoning summary events (`response.reasoning_summary_*`) from reasoning deltas. Decide whether Thinking maps to full reasoning or summary only. |
|
||||
|
||||
## Stop reasons
|
||||
|
||||
| llm_client `StopReason` | Open Responses | Note |
|
||||
| ----------------------- | ------------------------------------------------------------------------------ | ---------------------------------------------- |
|
||||
| `EndTurn` | `response.completed` + item status `completed` | |
|
||||
| `MaxTokens` | `response.incomplete` + item status `incomplete` | |
|
||||
| `StopSequence` | `response.completed` | |
|
||||
| `ToolUse` | `response.completed` for message item, followed by `function_call` output item | OR models tool call as a separate output item. |
|
||||
|
||||
## Gaps / open decisions
|
||||
|
||||
- `PingEvent` has no OR equivalent. If needed, keep as internal only.
|
||||
- `Cancelled` status needs a policy: map to `response.incomplete` or
|
||||
`response.failed`.
|
||||
- OR has `response.refusal.delta` / `response.refusal.done`. `llm_client` has no
|
||||
refusal delta type; consider adding a new block or delta variant if needed.
|
||||
- OR splits _item_ and _content part_ lifecycles. `llm_client` currently has a
|
||||
single block lifecycle, so mapping should decide whether to synthesize
|
||||
`content_part.*` events or ignore them.
|
||||
- The OR specification does not state how `function_call.arguments` stream
|
||||
deltas; `response.function_call_arguments.*` should be treated as a compatible
|
||||
extension if required.
|
||||
+1
-1
@@ -17,7 +17,7 @@ LLMを用いたワーカーを作成する小型のSDK・ライブラリ。
|
||||
|
||||
module構成概念図
|
||||
|
||||
```
|
||||
```plaintext
|
||||
worker
|
||||
├── context
|
||||
├── llm_client
|
||||
|
||||
@@ -27,7 +27,7 @@ RustのType-stateパターンを利用し、Workerの状態によって利用可
|
||||
* 自由な編集が可能な状態。
|
||||
* システムプロンプトの設定・変更が可能。
|
||||
* メッセージ履歴の初期構築(ロード、編集)が可能。
|
||||
* **`Locked` (キャッシュ保護状態)**
|
||||
* **`CacheLocked` (キャッシュ保護状態)**
|
||||
* キャッシュの有効活用を目的とした、前方不変状態。
|
||||
* **システムプロンプトの変更不可**。
|
||||
* **既存メッセージ履歴の変更不可**(追記のみ許可)。
|
||||
@@ -47,7 +47,7 @@ worker.history_mut().push(initial_message);
|
||||
|
||||
// 3. ロックしてLocked状態へ遷移
|
||||
// これにより、ここまでのコンテキストが "Fixed Prefix" として扱われる
|
||||
let mut locked_worker: Worker<Locked> = worker.lock();
|
||||
let mut locked_worker: Worker<CacheLocked> = worker.lock();
|
||||
|
||||
// 4. 利用 (Locked状態)
|
||||
// 実行は可能。新しいメッセージは履歴の末尾に追記される。
|
||||
@@ -65,4 +65,4 @@ locked_worker.run(new_user_input).await?;
|
||||
|
||||
* **状態パラメータの導入**: `Worker<S: WorkerState>` の導入。
|
||||
* **コンテキスト所有権の委譲**: `run` メソッドの引数でコンテキストを受け取るのではなく、`Worker` 内部に `history: Vec<Message>` を保持し管理する形へ移行する。
|
||||
* **APIの分離**: `Mutable` 特有のメソッド(setter等)と、`Locked` でも使えるメソッド(実行、参照等)をトレイト境界で分離する。
|
||||
* **APIの分離**: `Mutable` 特有のメソッド(setter等)と、`CacheLocked` でも使えるメソッド(実行、参照等)をトレイト境界で分離する。
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
# 非同期キャンセル設計
|
||||
|
||||
Workerの非同期キャンセル機構についての設計ドキュメント。
|
||||
|
||||
## 概要
|
||||
|
||||
`tokio::sync::mpsc`の通知チャネルを用いて、別タスクからWorkerの実行を安全にキャンセルできる。
|
||||
|
||||
```rust
|
||||
let worker = Arc::new(Mutex::new(Worker::new(client)));
|
||||
|
||||
// 実行タスク
|
||||
let w = worker.clone();
|
||||
let handle = tokio::spawn(async move {
|
||||
w.lock().await.run("prompt").await
|
||||
});
|
||||
|
||||
// キャンセル
|
||||
worker.lock().await.cancel();
|
||||
```
|
||||
|
||||
## キャンセル時の処理フロー
|
||||
|
||||
```
|
||||
キャンセル検知
|
||||
↓
|
||||
timeline.abort_current_block() // 進行中ブロックの終端処理
|
||||
↓
|
||||
run_on_abort_hooks("Cancelled") // on_abort フック呼び出し
|
||||
↓
|
||||
Err(WorkerError::Cancelled) // エラー返却
|
||||
```
|
||||
|
||||
## API
|
||||
|
||||
| メソッド | 説明 |
|
||||
| ----------------- | ------------------------------ |
|
||||
| `cancel()` | キャンセルをトリガー |
|
||||
| `cancel_sender()` | キャンセル通知用のSenderを取得 |
|
||||
|
||||
## on_abort フック
|
||||
|
||||
`Hook::on_abort(&self, reason: &str)`がキャンセル時に呼ばれる。
|
||||
クリーンアップ処理やログ記録に使用できる。
|
||||
|
||||
```rust
|
||||
async fn on_abort(&self, reason: &str) -> Result<(), HookError> {
|
||||
log::info!("Aborted: {}", reason);
|
||||
Ok(())
|
||||
}
|
||||
```
|
||||
|
||||
呼び出しタイミング:
|
||||
|
||||
- `WorkerError::Cancelled` — reason: `"Cancelled"`
|
||||
- `ControlFlow::Abort(reason)` — reason: フックが指定した理由
|
||||
|
||||
---
|
||||
|
||||
## 既知の問題
|
||||
|
||||
### on_abort の発火基準
|
||||
|
||||
`on_abort` は **interrupt(中断)** された場合に必ず発火する。
|
||||
|
||||
interrupt の例:
|
||||
|
||||
- `WorkerError::Cancelled`(キャンセル)
|
||||
- `WorkerError::Aborted`(フックによるAbort)
|
||||
- ストリーム/ツール/クライアント/Hook の各種エラーで処理が中断された場合
|
||||
+219
-139
@@ -3,7 +3,8 @@
|
||||
## 概要
|
||||
|
||||
HookはWorker層でのターン制御に介入するためのメカニズムです。
|
||||
Claude CodeのHooks機能に着想を得ており、メッセージ送信・ツール実行・ターン終了の各ポイントで処理を差し込むことができます。
|
||||
|
||||
メッセージ送信・ツール実行・ターン終了等の各ポイントで処理を差し込むことができます。
|
||||
|
||||
## コンセプト
|
||||
|
||||
@@ -11,120 +12,184 @@ Claude CodeのHooks機能に着想を得ており、メッセージ送信・ツ
|
||||
- **Contextへのアクセス**: メッセージ履歴を読み書き可能
|
||||
- **非破壊的チェーン**: 複数のHookを登録順に実行、後続Hookへの影響を制御
|
||||
|
||||
## Hook一覧
|
||||
|
||||
| Hook | タイミング | 主な用途 | 戻り値 |
|
||||
| ------------------ | -------------------------- | -------------------------- | ---------------------- |
|
||||
| `on_prompt_submit` | `run()` 呼び出し時 | ユーザーメッセージの前処理 | `OnPromptSubmitResult` |
|
||||
| `pre_llm_request` | 各ターンのLLM送信前 | コンテキスト改変/検証 | `PreLlmRequestResult` |
|
||||
| `pre_tool_call` | ツール実行前 | 実行許可/引数改変 | `PreToolCallResult` |
|
||||
| `post_tool_call` | ツール実行後 | 結果加工/マスキング | `PostToolCallResult` |
|
||||
| `on_turn_end` | ツールなしでターン終了直前 | 検証/リトライ指示 | `OnTurnEndResult` |
|
||||
| `on_abort` | 中断時 | クリーンアップ/通知 | `()` |
|
||||
|
||||
## Hook Trait
|
||||
|
||||
```rust
|
||||
#[async_trait]
|
||||
pub trait WorkerHook: Send + Sync {
|
||||
/// メッセージ送信前
|
||||
/// リクエストに含まれるメッセージリストを改変できる
|
||||
async fn on_message_send(
|
||||
&self,
|
||||
context: &mut Vec<Message>,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// ツール実行前
|
||||
/// 実行をキャンセルしたり、引数を書き換えることができる
|
||||
async fn before_tool_call(
|
||||
&self,
|
||||
tool_call: &mut ToolCall,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// ツール実行後
|
||||
/// 結果を書き換えたり、隠蔽したりできる
|
||||
async fn after_tool_call(
|
||||
&self,
|
||||
tool_result: &mut ToolResult,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// ターン終了時
|
||||
/// 生成されたメッセージを検査し、必要ならリトライを指示できる
|
||||
async fn on_turn_end(
|
||||
&self,
|
||||
messages: &[Message],
|
||||
) -> Result<TurnResult, HookError> {
|
||||
Ok(TurnResult::Finish)
|
||||
}
|
||||
pub trait Hook<E: HookEventKind>: Send + Sync {
|
||||
async fn call(&self, input: &mut E::Input) -> Result<E::Output, HookError>;
|
||||
}
|
||||
```
|
||||
|
||||
## 制御フロー型
|
||||
|
||||
### ControlFlow
|
||||
### HookEventKind / Result
|
||||
|
||||
Hook処理の継続/中断を制御する列挙型。
|
||||
Hookイベントごとに入力/出力型を分離し、意味のない制御フローを排除する。
|
||||
|
||||
```rust
|
||||
pub enum ControlFlow {
|
||||
/// 処理を続行(後続Hookも実行)
|
||||
pub trait HookEventKind {
|
||||
type Input;
|
||||
type Output;
|
||||
}
|
||||
|
||||
pub struct OnPromptSubmit;
|
||||
pub struct PreLlmRequest;
|
||||
pub struct PreToolCall;
|
||||
pub struct PostToolCall;
|
||||
pub struct OnTurnEnd;
|
||||
pub struct OnAbort;
|
||||
|
||||
pub enum OnPromptSubmitResult {
|
||||
Continue,
|
||||
Cancel(String),
|
||||
}
|
||||
|
||||
pub enum PreLlmRequestResult {
|
||||
Continue,
|
||||
Cancel(String),
|
||||
}
|
||||
|
||||
pub enum PreToolCallResult {
|
||||
Continue,
|
||||
/// 現在の処理をスキップ(ツール実行をスキップ等)
|
||||
Skip,
|
||||
/// 処理全体を中断(エラーとして扱う)
|
||||
Abort(String),
|
||||
Pause,
|
||||
}
|
||||
|
||||
pub enum PostToolCallResult {
|
||||
Continue,
|
||||
Abort(String),
|
||||
}
|
||||
|
||||
pub enum OnTurnEndResult {
|
||||
Finish,
|
||||
ContinueWithMessages(Vec<Message>),
|
||||
Paused,
|
||||
}
|
||||
```
|
||||
|
||||
### TurnResult
|
||||
### Tool Call Context
|
||||
|
||||
ターン終了時の判定結果を表す列挙型。
|
||||
`pre_tool_call` / `post_tool_call` は、ツール実行の文脈を含む入力を受け取る。
|
||||
|
||||
```rust
|
||||
pub enum TurnResult {
|
||||
/// ターンを正常終了
|
||||
Finish,
|
||||
/// メッセージを追加してターン継続(自己修正など)
|
||||
ContinueWithMessages(Vec<Message>),
|
||||
pub struct ToolCallContext {
|
||||
pub call: ToolCall,
|
||||
pub meta: ToolMeta, // 不変メタデータ
|
||||
pub tool: Arc<dyn Tool>, // 状態アクセス用
|
||||
}
|
||||
|
||||
pub struct PostToolCallContext {
|
||||
pub call: ToolCall,
|
||||
pub result: ToolResult,
|
||||
pub meta: ToolMeta,
|
||||
pub tool: Arc<dyn Tool>,
|
||||
}
|
||||
```
|
||||
|
||||
## 呼び出しタイミング
|
||||
|
||||
```
|
||||
Worker::run() ループ
|
||||
Worker::run(user_input)
|
||||
│
|
||||
├─▶ on_message_send ──────────────────────────────┐
|
||||
│ コンテキストの改変、バリデーション、 │
|
||||
│ システムプロンプト注入などが可能 │
|
||||
├─▶ on_prompt_submit ───────────────────────────┐
|
||||
│ ユーザーメッセージの前処理・検証 │
|
||||
│ (最初の1回のみ) │
|
||||
│ │
|
||||
├─▶ LLMリクエスト送信 & ストリーム処理 │
|
||||
│ │
|
||||
├─▶ ツール呼び出しがある場合: │
|
||||
│ │ │
|
||||
│ ├─▶ before_tool_call (各ツールごと・逐次) │
|
||||
│ │ 実行可否の判定、引数の改変 │
|
||||
│ │ │
|
||||
│ ├─▶ ツール並列実行 (join_all) │
|
||||
│ │ │
|
||||
│ └─▶ after_tool_call (各結果ごと・逐次) │
|
||||
│ 結果の確認、加工、ログ出力 │
|
||||
│ │
|
||||
├─▶ ツール結果をコンテキストに追加 → ループ先頭へ │
|
||||
│ │
|
||||
└─▶ ツールなしの場合: │
|
||||
└─▶ loop {
|
||||
│
|
||||
├─▶ pre_llm_request ──────────────────────│
|
||||
│ コンテキストの改変、バリデーション、 │
|
||||
│ システムプロンプト注入などが可能 │
|
||||
│ (毎ターン実行) │
|
||||
│ │
|
||||
└─▶ on_turn_end ─────────────────────────────┘
|
||||
├─▶ LLMリクエスト送信 & ストリーム処理 │
|
||||
│ │
|
||||
├─▶ ツール呼び出しがある場合: │
|
||||
│ │ │
|
||||
│ ├─▶ pre_tool_call (各ツールごと・逐次) │
|
||||
│ │ 実行可否の判定、引数の改変 │
|
||||
│ │ │
|
||||
│ ├─▶ ツール並列実行 (join_all) │
|
||||
│ │ │
|
||||
│ └─▶ post_tool_call (各結果ごと・逐次) │
|
||||
│ 結果の確認、加工、ログ出力 │
|
||||
│ │
|
||||
├─▶ ツール結果をコンテキストに追加 │
|
||||
│ → ループ先頭へ │
|
||||
│ │
|
||||
└─▶ ツールなしの場合: │
|
||||
│ │
|
||||
└─▶ on_turn_end ───────────────────┘
|
||||
最終応答のチェック(Lint/Fmt等)
|
||||
エラーがあればContinueWithMessagesでリトライ
|
||||
}
|
||||
|
||||
※ 中断時は on_abort が呼ばれる
|
||||
```
|
||||
|
||||
## 各Hookの詳細
|
||||
|
||||
### on_message_send
|
||||
### on_prompt_submit
|
||||
|
||||
**呼び出しタイミング**: LLMへリクエスト送信前(ターンループの冒頭)
|
||||
**呼び出しタイミング**: `run()`
|
||||
でユーザーメッセージを受け取った直後(最初の1回のみ)
|
||||
|
||||
**用途**:
|
||||
|
||||
- ユーザー入力のバリデーション
|
||||
- 入力のサニタイズ・フィルタリング
|
||||
- ログ出力
|
||||
- `OnPromptSubmitResult::Cancel` による実行キャンセル
|
||||
|
||||
**入力**: `&mut Message` - ユーザーメッセージ(改変可能)
|
||||
|
||||
**例**: 入力のバリデーション
|
||||
|
||||
```rust
|
||||
struct InputValidator;
|
||||
|
||||
#[async_trait]
|
||||
impl Hook<OnPromptSubmit> for InputValidator {
|
||||
async fn call(
|
||||
&self,
|
||||
message: &mut Message,
|
||||
) -> Result<OnPromptSubmitResult, HookError> {
|
||||
if let MessageContent::Text(text) = &message.content {
|
||||
if text.trim().is_empty() {
|
||||
return Ok(OnPromptSubmitResult::Cancel("Empty input".to_string()));
|
||||
}
|
||||
}
|
||||
Ok(OnPromptSubmitResult::Continue)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### pre_llm_request
|
||||
|
||||
**呼び出しタイミング**: 各ターンのLLMリクエスト送信前(ループの毎回)
|
||||
|
||||
**用途**:
|
||||
|
||||
- コンテキストへのシステムメッセージ注入
|
||||
- メッセージのバリデーション
|
||||
- 機密情報のフィルタリング
|
||||
- リクエスト内容のログ出力
|
||||
- `PreLlmRequestResult::Cancel` による送信キャンセル
|
||||
|
||||
**入力**: `&mut Vec<Message>` - コンテキスト全体(改変可能)
|
||||
|
||||
**例**: メッセージにタイムスタンプを追加
|
||||
|
||||
@@ -132,27 +197,33 @@ Worker::run() ループ
|
||||
struct TimestampHook;
|
||||
|
||||
#[async_trait]
|
||||
impl WorkerHook for TimestampHook {
|
||||
async fn on_message_send(
|
||||
impl Hook<PreLlmRequest> for TimestampHook {
|
||||
async fn call(
|
||||
&self,
|
||||
context: &mut Vec<Message>,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
) -> Result<PreLlmRequestResult, HookError> {
|
||||
let timestamp = chrono::Local::now().to_rfc3339();
|
||||
context.insert(0, Message::user(format!("[{}]", timestamp)));
|
||||
Ok(ControlFlow::Continue)
|
||||
Ok(PreLlmRequestResult::Continue)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### before_tool_call
|
||||
### pre_tool_call
|
||||
|
||||
**呼び出しタイミング**: 各ツール実行前(並列実行フェーズの前)
|
||||
|
||||
**用途**:
|
||||
|
||||
- 危険なツールのブロック
|
||||
- 引数のサニタイズ
|
||||
- 確認プロンプトの表示(UIとの連携)
|
||||
- 実行ログの記録
|
||||
- `PreToolCallResult::Pause` による一時停止
|
||||
|
||||
**入力**:
|
||||
|
||||
- `ToolCallContext`(`ToolCall` + `ToolMeta` + `Arc<dyn Tool>`)
|
||||
|
||||
**例**: 特定ツールをブロック
|
||||
|
||||
@@ -162,46 +233,52 @@ struct ToolBlocker {
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl WorkerHook for ToolBlocker {
|
||||
async fn before_tool_call(
|
||||
impl Hook<PreToolCall> for ToolBlocker {
|
||||
async fn call(
|
||||
&self,
|
||||
tool_call: &mut ToolCall,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
if self.blocked_tools.contains(&tool_call.name) {
|
||||
println!("Blocked tool: {}", tool_call.name);
|
||||
Ok(ControlFlow::Skip)
|
||||
ctx: &mut ToolCallContext,
|
||||
) -> Result<PreToolCallResult, HookError> {
|
||||
if self.blocked_tools.contains(&ctx.call.name) {
|
||||
println!("Blocked tool: {}", ctx.call.name);
|
||||
Ok(PreToolCallResult::Skip)
|
||||
} else {
|
||||
Ok(ControlFlow::Continue)
|
||||
Ok(PreToolCallResult::Continue)
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### after_tool_call
|
||||
### post_tool_call
|
||||
|
||||
**呼び出しタイミング**: 各ツール実行後(並列実行フェーズの後)
|
||||
|
||||
**用途**:
|
||||
|
||||
- 結果の加工・フォーマット
|
||||
- 機密情報のマスキング
|
||||
- 結果のキャッシュ
|
||||
- 実行結果のログ出力
|
||||
|
||||
**入力**:
|
||||
|
||||
- `PostToolCallContext`(`ToolCall` + `ToolResult` + `ToolMeta` +
|
||||
`Arc<dyn Tool>`)
|
||||
|
||||
**例**: 結果にプレフィックスを追加
|
||||
|
||||
```rust
|
||||
struct ResultFormatter;
|
||||
|
||||
#[async_trait]
|
||||
impl WorkerHook for ResultFormatter {
|
||||
async fn after_tool_call(
|
||||
impl Hook<PostToolCall> for ResultFormatter {
|
||||
async fn call(
|
||||
&self,
|
||||
tool_result: &mut ToolResult,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
if !tool_result.is_error {
|
||||
tool_result.content = format!("[OK] {}", tool_result.content);
|
||||
ctx: &mut PostToolCallContext,
|
||||
) -> Result<PostToolCallResult, HookError> {
|
||||
if !ctx.result.is_error {
|
||||
ctx.result.content = format!("[OK] {}", ctx.result.content);
|
||||
}
|
||||
Ok(ControlFlow::Continue)
|
||||
Ok(PostToolCallResult::Continue)
|
||||
}
|
||||
}
|
||||
```
|
||||
@@ -211,10 +288,22 @@ impl WorkerHook for ResultFormatter {
|
||||
**呼び出しタイミング**: ツール呼び出しなしでターンが終了する直前
|
||||
|
||||
**用途**:
|
||||
|
||||
- 生成されたコードのLint/Fmt
|
||||
- 出力形式のバリデーション
|
||||
- 自己修正のためのリトライ指示
|
||||
- 最終結果のログ出力
|
||||
- `OnTurnEndResult::Paused` による一時停止
|
||||
|
||||
### on_abort
|
||||
|
||||
**呼び出しタイミング**: キャンセル/エラー/AbortなどでWorkerが中断された時
|
||||
|
||||
**用途**:
|
||||
|
||||
- クリーンアップ処理
|
||||
- 中断理由のログ出力
|
||||
- 外部システムへの通知
|
||||
|
||||
**例**: JSON形式のバリデーション
|
||||
|
||||
@@ -222,11 +311,11 @@ impl WorkerHook for ResultFormatter {
|
||||
struct JsonValidator;
|
||||
|
||||
#[async_trait]
|
||||
impl WorkerHook for JsonValidator {
|
||||
async fn on_turn_end(
|
||||
impl Hook<OnTurnEnd> for JsonValidator {
|
||||
async fn call(
|
||||
&self,
|
||||
messages: &[Message],
|
||||
) -> Result<TurnResult, HookError> {
|
||||
messages: &mut Vec<Message>,
|
||||
) -> Result<OnTurnEndResult, HookError> {
|
||||
// 最後のアシスタントメッセージを取得
|
||||
let last = messages.iter().rev()
|
||||
.find(|m| m.role == Role::Assistant);
|
||||
@@ -236,25 +325,25 @@ impl WorkerHook for JsonValidator {
|
||||
// JSONとしてパースを試みる
|
||||
if serde_json::from_str::<serde_json::Value>(text).is_err() {
|
||||
// 失敗したらリトライ指示
|
||||
return Ok(TurnResult::ContinueWithMessages(vec![
|
||||
return Ok(OnTurnEndResult::ContinueWithMessages(vec![
|
||||
Message::user("Invalid JSON. Please fix and try again.")
|
||||
]));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(TurnResult::Finish)
|
||||
Ok(OnTurnEndResult::Finish)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## 複数Hookの実行順序
|
||||
|
||||
Hookは**登録順**に実行されます。
|
||||
Hookは**イベントごとに登録順**に実行されます。
|
||||
|
||||
```rust
|
||||
worker.add_hook(HookA); // 1番目に実行
|
||||
worker.add_hook(HookB); // 2番目に実行
|
||||
worker.add_hook(HookC); // 3番目に実行
|
||||
worker.add_pre_tool_call_hook(HookA); // 1番目に実行
|
||||
worker.add_pre_tool_call_hook(HookB); // 2番目に実行
|
||||
worker.add_pre_tool_call_hook(HookC); // 3番目に実行
|
||||
```
|
||||
|
||||
### 制御フローの伝播
|
||||
@@ -262,6 +351,7 @@ worker.add_hook(HookC); // 3番目に実行
|
||||
- `Continue`: 後続Hookも実行
|
||||
- `Skip`: 現在の処理をスキップし、後続Hookは実行しない
|
||||
- `Abort`: 即座にエラーを返し、処理全体を中断
|
||||
- `Pause`: Workerを一時停止(再開は`resume`)
|
||||
|
||||
```
|
||||
Hook A: Continue → Hook B: Skip → (Hook Cは実行されない)
|
||||
@@ -271,52 +361,39 @@ Hook A: Continue → Hook B: Skip → (Hook Cは実行されない)
|
||||
Hook A: Continue → Hook B: Abort("reason")
|
||||
↓
|
||||
WorkerError::Aborted
|
||||
|
||||
Hook A: Continue → Hook B: Pause
|
||||
↓
|
||||
WorkerResult::Paused
|
||||
```
|
||||
|
||||
## 設計上のポイント
|
||||
|
||||
### 1. デフォルト実装
|
||||
### 1. イベントごとの実装
|
||||
|
||||
全メソッドにデフォルト実装があるため、必要なメソッドだけオーバーライドすれば良い。
|
||||
|
||||
```rust
|
||||
struct SimpleLogger;
|
||||
|
||||
#[async_trait]
|
||||
impl WorkerHook for SimpleLogger {
|
||||
// on_message_send だけ実装
|
||||
async fn on_message_send(
|
||||
&self,
|
||||
context: &mut Vec<Message>,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
println!("Sending {} messages", context.len());
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
// 他のメソッドはデフォルト(Continue/Finish)
|
||||
}
|
||||
```
|
||||
必要なイベントのみ `Hook<Event>` を実装する。
|
||||
|
||||
### 2. 可変参照による改変
|
||||
|
||||
`&mut`で引数を受け取るため、直接改変が可能。
|
||||
|
||||
```rust
|
||||
async fn before_tool_call(&self, tool_call: &mut ToolCall) -> ... {
|
||||
async fn call(&self, ctx: &mut ToolCallContext) -> ... {
|
||||
// 引数を直接書き換え
|
||||
tool_call.input["sanitized"] = json!(true);
|
||||
Ok(ControlFlow::Continue)
|
||||
ctx.call.input["sanitized"] = json!(true);
|
||||
Ok(PreToolCallResult::Continue)
|
||||
}
|
||||
```
|
||||
|
||||
### 3. 並列実行との統合
|
||||
|
||||
- `before_tool_call`: 並列実行**前**に逐次実行(許可判定のため)
|
||||
- `pre_tool_call`: 並列実行**前**に逐次実行(許可判定のため)
|
||||
- ツール実行: `join_all`で**並列**実行
|
||||
- `after_tool_call`: 並列実行**後**に逐次実行(結果加工のため)
|
||||
- `post_tool_call`: 並列実行**後**に逐次実行(結果加工のため)
|
||||
|
||||
### 4. Send + Sync 要件
|
||||
|
||||
`WorkerHook`は`Send + Sync`を要求するため、スレッドセーフな実装が必要。
|
||||
`Hook`は`Send + Sync`を要求するため、スレッドセーフな実装が必要。
|
||||
状態を持つ場合は`Arc<Mutex<T>>`や`AtomicUsize`などを使用する。
|
||||
|
||||
```rust
|
||||
@@ -325,10 +402,10 @@ struct CountingHook {
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl WorkerHook for CountingHook {
|
||||
async fn before_tool_call(&self, _: &mut ToolCall) -> Result<ControlFlow, HookError> {
|
||||
impl Hook<PreToolCall> for CountingHook {
|
||||
async fn call(&self, _: &mut ToolCallContext) -> Result<PreToolCallResult, HookError> {
|
||||
self.count.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(ControlFlow::Continue)
|
||||
Ok(PreToolCallResult::Continue)
|
||||
}
|
||||
}
|
||||
```
|
||||
@@ -336,13 +413,13 @@ impl WorkerHook for CountingHook {
|
||||
## 典型的なユースケース
|
||||
|
||||
| ユースケース | 使用Hook | 処理内容 |
|
||||
|-------------|----------|----------|
|
||||
| ツール許可制御 | `before_tool_call` | 危険なツールをSkip |
|
||||
| 実行ログ | `before/after_tool_call` | 呼び出しと結果を記録 |
|
||||
| ------------------ | -------------------- | -------------------------- |
|
||||
| ツール許可制御 | `pre_tool_call` | 危険なツールをSkip |
|
||||
| 実行ログ | `pre/post_tool_call` | 呼び出しと結果を記録 |
|
||||
| 出力バリデーション | `on_turn_end` | 形式チェック、リトライ指示 |
|
||||
| コンテキスト注入 | `on_message_send` | システムメッセージ追加 |
|
||||
| 結果のサニタイズ | `after_tool_call` | 機密情報のマスキング |
|
||||
| レート制限 | `before_tool_call` | 呼び出し頻度の制御 |
|
||||
| 結果のサニタイズ | `post_tool_call` | 機密情報のマスキング |
|
||||
| レート制限 | `pre_tool_call` | 呼び出し頻度の制御 |
|
||||
|
||||
## TODO
|
||||
|
||||
@@ -350,11 +427,14 @@ impl WorkerHook for CountingHook {
|
||||
|
||||
現在のHooks実装は基本的なユースケースをカバーしているが、以下の点について将来的に厳密な仕様を定義する必要がある:
|
||||
|
||||
- **エラーハンドリングの明確化**: `HookError`発生時のリカバリー戦略、部分的な失敗の扱い
|
||||
- **エラーハンドリングの明確化**:
|
||||
`HookError`発生時のリカバリー戦略、部分的な失敗の扱い
|
||||
- **Hook間の依存関係**: 複数Hookの実行順序が結果に影響する場合のセマンティクス
|
||||
- **非同期キャンセル**: Hook実行中のキャンセル(タイムアウト等)の振る舞い
|
||||
- **状態の一貫性**: `on_message_send`で改変されたコンテキストが後続処理で期待通りに反映される保証
|
||||
- **リトライ制限**: `on_turn_end`での`ContinueWithMessages`による無限ループ防止策
|
||||
- **状態の一貫性**:
|
||||
`on_message_send`で改変されたコンテキストが後続処理で期待通りに反映される保証
|
||||
- **リトライ制限**:
|
||||
`on_turn_end`での`ContinueWithMessages`による無限ループ防止策
|
||||
- **Hook優先度**: 登録順以外の優先度指定メカニズムの必要性
|
||||
- **条件付きHook**: 特定条件でのみ有効化されるHookパターン
|
||||
- **テスト容易性**: Hookのモック/スタブ作成のためのユーティリティ
|
||||
|
||||
@@ -0,0 +1,191 @@
|
||||
# Tool 設計
|
||||
|
||||
## 概要
|
||||
|
||||
`llm-worker`のツールシステムは、LLMが外部リソースにアクセスしたり計算を実行するための仕組みを提供する。
|
||||
メタ情報の不変性とセッションスコープの状態管理を両立させる設計となっている。
|
||||
|
||||
## 主要な型
|
||||
|
||||
```
|
||||
type ToolDefinition
|
||||
Fn() -> (ToolMeta, Arc<dyn Tool>)
|
||||
|
||||
worker.register_tool() で呼び出し
|
||||
|
||||
▼
|
||||
|
||||
- struct ToolMeta (name, desc, schema)
|
||||
不変・登録時固定
|
||||
- trait Tool (executer)
|
||||
登録時生成・セッション中再利用
|
||||
```
|
||||
|
||||
### ToolMeta
|
||||
|
||||
ツールのメタ情報を保持する不変構造体。登録時に固定され、Worker内で変更されない。
|
||||
|
||||
```rust
|
||||
pub struct ToolMeta {
|
||||
pub name: String,
|
||||
pub description: String,
|
||||
pub input_schema: Value,
|
||||
}
|
||||
```
|
||||
|
||||
**目的:**
|
||||
|
||||
- LLM へのツール定義として送信
|
||||
- Hook からの参照(読み取り専用)
|
||||
- 登録後の不変性を保証
|
||||
|
||||
### Tool trait
|
||||
|
||||
ツールの実行ロジックのみを定義するトレイト。
|
||||
|
||||
```rust
|
||||
#[async_trait]
|
||||
pub trait Tool: Send + Sync {
|
||||
async fn execute(&self, input_json: &str) -> Result<String, ToolError>;
|
||||
}
|
||||
```
|
||||
|
||||
**設計方針:**
|
||||
|
||||
- メタ情報(name, description, schema)は含まない
|
||||
- 状態を持つことが可能(セッション中のカウンターなど)
|
||||
- `Send + Sync` で並列実行に対応
|
||||
|
||||
**インスタンスのライフサイクル:**
|
||||
|
||||
1. `register_tool()` 呼び出し時にファクトリが実行され、インスタンスが生成される
|
||||
2. LLM がツールを呼び出すと、既存インスタンスの `execute()` が実行される
|
||||
3. 同じセッション中は同一インスタンスが再利用される
|
||||
|
||||
※ 「最初に呼ばれたとき」の遅延初期化ではなく、**登録時の即時初期化**である。
|
||||
|
||||
### ToolDefinition
|
||||
|
||||
メタ情報とツールインスタンスを生成するファクトリ。
|
||||
|
||||
```rust
|
||||
pub type ToolDefinition = Arc<dyn Fn() -> (ToolMeta, Arc<dyn Tool>) + Send + Sync>;
|
||||
```
|
||||
|
||||
**なぜファクトリか:**
|
||||
|
||||
- Worker への登録時に一度だけ呼び出される
|
||||
- メタ情報とインスタンスを同時に生成し、整合性を保証
|
||||
- クロージャでコンテキスト(`self.clone()`)をキャプチャ可能
|
||||
|
||||
## Worker でのツール管理
|
||||
|
||||
```rust
|
||||
// Worker 内部
|
||||
tools: HashMap<String, (ToolMeta, Arc<dyn Tool>)>
|
||||
|
||||
// 登録 API
|
||||
pub fn register_tool(&mut self, factory: ToolDefinition) -> Result<(), ToolRegistryError>
|
||||
```
|
||||
|
||||
登録時の処理:
|
||||
|
||||
1. ファクトリを呼び出し `(meta, instance)` を取得
|
||||
2. 同名ツールが既に登録されていればエラー
|
||||
3. HashMap に `(meta, instance)` を保存
|
||||
|
||||
## マクロによる自動生成
|
||||
|
||||
`#[tool_registry]` マクロは `{method}_definition()` メソッドを生成する。
|
||||
|
||||
```rust
|
||||
#[tool_registry]
|
||||
impl MyApp {
|
||||
/// 検索を実行する
|
||||
#[tool]
|
||||
async fn search(&self, query: String) -> String {
|
||||
// 実装
|
||||
}
|
||||
}
|
||||
|
||||
// 生成されるコード:
|
||||
impl MyApp {
|
||||
pub fn search_definition(&self) -> ToolDefinition {
|
||||
let ctx = self.clone();
|
||||
Arc::new(move || {
|
||||
let meta = ToolMeta::new("search")
|
||||
.description("検索を実行する")
|
||||
.input_schema(/* schemars で生成 */);
|
||||
let tool = Arc::new(ToolSearch { ctx: ctx.clone() });
|
||||
(meta, tool)
|
||||
})
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Hook との連携
|
||||
|
||||
Hook は `ToolCallContext` / `AfterToolCallContext`
|
||||
を通じてメタ情報とインスタンスにアクセスできる。
|
||||
|
||||
```rust
|
||||
pub struct ToolCallContext {
|
||||
pub call: ToolCall, // 呼び出し情報(改変可能)
|
||||
pub meta: ToolMeta, // メタ情報(読み取り専用)
|
||||
pub tool: Arc<dyn Tool>, // インスタンス(状態アクセス用)
|
||||
}
|
||||
```
|
||||
|
||||
**用途:**
|
||||
|
||||
- `meta` で名前やスキーマを確認
|
||||
- `tool` でツールの内部状態を読み取り(ダウンキャスト必要)
|
||||
- `call` の引数を改変してツールに渡す
|
||||
|
||||
## 使用例
|
||||
|
||||
### 手動実装
|
||||
|
||||
```rust
|
||||
struct Counter { count: AtomicUsize }
|
||||
|
||||
impl Tool for Counter {
|
||||
async fn execute(&self, _: &str) -> Result<String, ToolError> {
|
||||
let n = self.count.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(format!("count: {}", n))
|
||||
}
|
||||
}
|
||||
|
||||
let def: ToolDefinition = Arc::new(|| {
|
||||
let meta = ToolMeta::new("counter")
|
||||
.description("カウンターを増加")
|
||||
.input_schema(json!({"type": "object"}));
|
||||
(meta, Arc::new(Counter { count: AtomicUsize::new(0) }))
|
||||
});
|
||||
|
||||
worker.register_tool(def)?;
|
||||
```
|
||||
|
||||
### マクロ使用(推奨)
|
||||
|
||||
```rust
|
||||
#[tool_registry]
|
||||
impl App {
|
||||
#[tool]
|
||||
async fn greet(&self, name: String) -> String {
|
||||
format!("Hello, {}!", name)
|
||||
}
|
||||
}
|
||||
|
||||
let app = App;
|
||||
worker.register_tool(app.greet_definition())?;
|
||||
```
|
||||
|
||||
## 設計上の決定
|
||||
|
||||
| 問題 | 決定 | 理由 |
|
||||
| -------------------- | ------------------------------ | ---------------------------------------------- |
|
||||
| メタ情報の変更可能性 | ToolMeta を分離・不変化 | 登録後の整合性を保証 |
|
||||
| 状態管理 | 登録時にインスタンス生成 | セッション中の状態保持、同一インスタンス再利用 |
|
||||
| Factory vs Instance | Factory + 登録時即時呼び出し | コンテキストキャプチャと登録時検証 |
|
||||
| Hook からのアクセス | Context に meta と tool を含む | 柔軟な介入を可能に |
|
||||
+98
-62
@@ -33,39 +33,46 @@ Workerは以下のループ(ターン)を実行します。
|
||||
|
||||
1. **Start Turn**: `Worker::run(messages)` 呼び出し
|
||||
2. **Hook: OnMessageSend**:
|
||||
* ユーザーメッセージの改変、バリデーション、キャンセルが可能。
|
||||
* コンテキストへのシステムプロンプト注入などもここで行う想定。
|
||||
- ユーザーメッセージの改変、バリデーション、キャンセルが可能。
|
||||
- コンテキストへのシステムプロンプト注入などもここで行う想定。
|
||||
3. **Request & Stream**:
|
||||
* LLMへリクエスト送信。イベントストリーム開始。
|
||||
* `Timeline`によるイベント処理。
|
||||
- LLMへリクエスト送信。イベントストリーム開始。
|
||||
- `Timeline`によるイベント処理。
|
||||
4. **Tool Handling (Parallel)**:
|
||||
* レスポンス内に含まれる全てのTool Callを収集。
|
||||
* 各Toolに対して **Hook: BeforeToolCall** を実行(実行可否、引数改変)。
|
||||
* 許可されたToolを**並列実行 (`join_all`)**。
|
||||
* 各Tool実行後に **Hook: AfterToolCall** を実行(結果の確認、加工)。
|
||||
- レスポンス内に含まれる全てのTool Callを収集。
|
||||
- 各Toolに対して **Hook: BeforeToolCall** を実行(実行可否、引数改変)。
|
||||
- 許可されたToolを**並列実行 (`join_all`)**。
|
||||
- 各Tool実行後に **Hook: AfterToolCall** を実行(結果の確認、加工)。
|
||||
5. **Next Request Decision**:
|
||||
* Tool実行結果がある場合 -> 結果をMessageとしてContextに追加し、**Step 3へ戻る** (自動ループ)。
|
||||
* Tool実行がない場合 -> Step 6へ。
|
||||
- Tool実行結果がある場合 -> 結果をMessageとしてContextに追加し、**Step
|
||||
3へ戻る** (自動ループ)。
|
||||
- Tool実行がない場合 -> Step 6へ。
|
||||
6. **Hook: OnTurnEnd**:
|
||||
* 最終的な応答に対するチェック(Lint/Fmt)。
|
||||
* エラーがある場合、エラーメッセージをContextに追加して **Step 3へ戻る** ことで自己修正を促せる。
|
||||
* 問題なければターン終了。
|
||||
- 最終的な応答に対するチェック(Lint/Fmt)。
|
||||
- エラーがある場合、エラーメッセージをContextに追加して **Step 3へ戻る**
|
||||
ことで自己修正を促せる。
|
||||
- 問題なければターン終了。
|
||||
|
||||
## Tool 設計
|
||||
|
||||
### アーキテクチャ概要
|
||||
|
||||
Rustの静的型付けシステムとLLMの動的なツール呼び出し(文字列による指定)を、**Trait Object** と **動的ディスパッチ** を用いて接続します。
|
||||
Rustの静的型付けシステムとLLMの動的なツール呼び出し(文字列による指定)を、**Trait
|
||||
Object** と **動的ディスパッチ** を用いて接続します。
|
||||
|
||||
1. **共通インターフェース (`Tool` Trait)**: 全てのツールが実装すべき共通の振る舞い(メタデータ取得と実行)を定義します。
|
||||
2. **ラッパー生成 (`#[tool]` Macro)**: ユーザー定義のメソッドをラップし、`Tool` Traitを実装した構造体を自動生成します。
|
||||
3. **レジストリ (`HashMap`)**: Workerは動的ディスパッチ用に `HashMap<String, Box<dyn Tool>>` でツールを管理します。
|
||||
1. **共通インターフェース (`Tool` Trait)**:
|
||||
全てのツールが実装すべき共通の振る舞い(メタデータ取得と実行)を定義します。
|
||||
2. **ラッパー生成 (`#[tool]` Macro)**: ユーザー定義のメソッドをラップし、`Tool`
|
||||
Traitを実装した構造体を自動生成します。
|
||||
3. **レジストリ (`HashMap`)**: Workerは動的ディスパッチ用に
|
||||
`HashMap<String, Box<dyn Tool>>` でツールを管理します。
|
||||
|
||||
この仕組みにより、「名前からツールを探し、JSON引数を型変換して関数を実行する」フローを安全に実現します。
|
||||
|
||||
### 1. Tool Trait 定義
|
||||
|
||||
ツールが最低限持つべきインターフェースです。`Send + Sync` を必須とし、マルチスレッド(並列実行)に対応します。
|
||||
ツールが最低限持つべきインターフェースです。`Send + Sync`
|
||||
を必須とし、マルチスレッド(並列実行)に対応します。
|
||||
|
||||
```rust
|
||||
#[async_trait]
|
||||
@@ -113,7 +120,8 @@ impl MyApp {
|
||||
|
||||
**マクロ展開後のイメージ (擬似コード):**
|
||||
|
||||
マクロは、元のメソッドに対応する**ラッパー構造体**を生成します。このラッパーが `Tool` Trait を実装します。
|
||||
マクロは、元のメソッドに対応する**ラッパー構造体**を生成します。このラッパーが
|
||||
`Tool` Trait を実装します。
|
||||
|
||||
```rust
|
||||
// 1. 引数をデシリアライズ用の中間構造体に変換
|
||||
@@ -155,15 +163,18 @@ impl Tool for GetUserTool {
|
||||
|
||||
### 3. Workerによる実行フロー
|
||||
|
||||
Workerは生成されたラッパー構造体を `Box<dyn Tool>` として保持し、以下のフローで実行します。
|
||||
Workerは生成されたラッパー構造体を `Box<dyn Tool>`
|
||||
として保持し、以下のフローで実行します。
|
||||
|
||||
1. **登録**: アプリケーション開始時、コンテキスト(`MyApp`)から各ツールのラッパー(`GetUserTool`)を生成し、WorkerのMapに登録。
|
||||
2. **解決**: LLMからのレスポンスに含まれる `ToolUse { name: "get_user", ... }` を受け取る。
|
||||
1. **登録**:
|
||||
アプリケーション開始時、コンテキスト(`MyApp`)から各ツールのラッパー(`GetUserTool`)を生成し、WorkerのMapに登録。
|
||||
2. **解決**: LLMからのレスポンスに含まれる `ToolUse { name: "get_user", ... }`
|
||||
を受け取る。
|
||||
3. **検索**: `name` をキーに Map から `Box<dyn Tool>` を取得。
|
||||
4. **実行**:
|
||||
* `tool.execute(json)` を呼び出す。
|
||||
* 内部で `serde_json` による型変換とメソッド実行が行われる。
|
||||
* 結果が返る。
|
||||
- `tool.execute(json)` を呼び出す。
|
||||
- 内部で `serde_json` による型変換とメソッド実行が行われる。
|
||||
- 結果が返る。
|
||||
|
||||
これにより、型安全性を保ちつつ、動的なツール実行が可能になります。
|
||||
|
||||
@@ -171,64 +182,88 @@ Workerは生成されたラッパー構造体を `Box<dyn Tool>` として保持
|
||||
|
||||
### コンセプト
|
||||
|
||||
* **制御の介入**: ターンの進行、メッセージの内容、ツールの実行に対して介入します。
|
||||
* **Contextへのアクセス**: メッセージ履歴(Context)を読み書きできます。
|
||||
- **制御の介入**:
|
||||
ターンの進行、メッセージの内容、ツールの実行に対して介入します。
|
||||
- **Contextへのアクセス**: メッセージ履歴(Context)を読み書きできます。
|
||||
|
||||
### Hook Trait
|
||||
|
||||
```rust
|
||||
#[async_trait]
|
||||
pub trait WorkerHook: Send + Sync {
|
||||
/// メッセージ送信前。
|
||||
/// リクエストに含まれるメッセージリストを改変できる。
|
||||
async fn on_message_send(&self, context: &mut Vec<Message>) -> Result<ControlFlow, Error> {
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// ツール実行前。
|
||||
/// 実行をキャンセルしたり、引数を書き換えることができる。
|
||||
async fn before_tool_call(&self, tool_call: &mut ToolCall) -> Result<ControlFlow, Error> {
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// ツール実行後。
|
||||
/// 結果を書き換えたり、隠蔽したりできる。
|
||||
async fn after_tool_call(&self, tool_result: &mut ToolResult) -> Result<ControlFlow, Error> {
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// ターン終了時。
|
||||
/// 生成されたメッセージを検査し、必要ならリトライ(ContinueWithMessages)を指示できる。
|
||||
async fn on_turn_end(&self, messages: &[Message]) -> Result<TurnResult, Error> {
|
||||
Ok(TurnResult::Finish)
|
||||
}
|
||||
pub trait Hook<E: HookEventKind>: Send + Sync {
|
||||
async fn call(&self, input: &mut E::Input) -> Result<E::Output, Error>;
|
||||
}
|
||||
|
||||
pub enum ControlFlow {
|
||||
pub trait HookEventKind {
|
||||
type Input;
|
||||
type Output;
|
||||
}
|
||||
|
||||
pub struct OnMessageSend;
|
||||
pub struct BeforeToolCall;
|
||||
pub struct AfterToolCall;
|
||||
pub struct OnTurnEnd;
|
||||
pub struct OnAbort;
|
||||
|
||||
pub enum OnMessageSendResult {
|
||||
Continue,
|
||||
Cancel(String),
|
||||
}
|
||||
|
||||
pub enum BeforeToolCallResult {
|
||||
Continue,
|
||||
Skip, // Tool実行などをスキップ
|
||||
Abort(String), // 処理中断
|
||||
Pause,
|
||||
}
|
||||
|
||||
pub enum TurnResult {
|
||||
pub enum AfterToolCallResult {
|
||||
Continue,
|
||||
Abort(String),
|
||||
}
|
||||
|
||||
pub enum OnTurnEndResult {
|
||||
Finish,
|
||||
ContinueWithMessages(Vec<Message>), // メッセージを追加してターン継続(自己修正など)
|
||||
Paused,
|
||||
}
|
||||
```
|
||||
|
||||
### Tool Call Context
|
||||
|
||||
`before_tool_call` / `after_tool_call`
|
||||
は、ツール実行の文脈を含む入力を受け取る。
|
||||
|
||||
```rust
|
||||
pub struct ToolCallContext {
|
||||
pub call: ToolCall,
|
||||
pub meta: ToolMeta, // 不変メタデータ
|
||||
pub tool: Arc<dyn Tool>, // 状態アクセス用
|
||||
}
|
||||
|
||||
pub struct ToolResultContext {
|
||||
pub result: ToolResult,
|
||||
pub meta: ToolMeta,
|
||||
pub tool: Arc<dyn Tool>,
|
||||
}
|
||||
```
|
||||
|
||||
## 実装方針
|
||||
|
||||
1. **Worker Struct**:
|
||||
* `Timeline`を所有。
|
||||
* `Handler`として「ToolCallCollector」をTimelineに登録。
|
||||
* `stream`終了後に収集したToolCallを処理するロジックを持つ。
|
||||
- `Timeline`を所有。
|
||||
- `Handler`として「ToolCallCollector」をTimelineに登録。
|
||||
- `stream`終了後に収集したToolCallを処理するロジックを持つ。
|
||||
- **履歴管理**: `set_history`, `with_messages`, `history_mut`
|
||||
等を通じて、会話履歴の注入や編集を可能にする。
|
||||
|
||||
2. **Tool Executor Handler**:
|
||||
* Timeline上ではツール実行を行わず、あくまで「ToolCallブロックの収集」に徹する(Toolの実行は非同期かつ並列で、ストリーム終了後あるいはブロック確定後に行うため)。
|
||||
* ただし、リアルタイム性を重視する場合(ストリーミング中にToolを実行開始等)は将来的な拡張とするが、現状は「結果が揃うのを待って」という要件に従い、収集フェーズと実行フェーズを分ける。
|
||||
- Timeline上ではツール実行を行わず、あくまで「ToolCallブロックの収集」に徹する(Toolの実行は非同期かつ並列で、ストリーム終了後あるいはブロック確定後に行うため)。
|
||||
- ただし、リアルタイム性を重視する場合(ストリーミング中にToolを実行開始等)は将来的な拡張とするが、現状は「結果が揃うのを待って」という要件に従い、収集フェーズと実行フェーズを分ける。
|
||||
|
||||
3. **worker-macros**:
|
||||
* `syn`, `quote` を用いて、関数定義から `Tool` トレイト実装と `InputSchema` (schemars利用) を生成。
|
||||
- `syn`, `quote` を用いて、関数定義から `Tool` トレイト実装と `InputSchema`
|
||||
(schemars利用) を生成。
|
||||
|
||||
## Worker Event API 設計
|
||||
|
||||
@@ -237,6 +272,7 @@ pub enum TurnResult {
|
||||
Workerは内部でイベントを処理し結果を返しますが、UIへのストリーミング表示やリアルタイムフィードバックには、イベントを外部に公開する仕組みが必要です。
|
||||
|
||||
**要件**:
|
||||
|
||||
1. テキストデルタをリアルタイムでUIに表示
|
||||
2. ツール呼び出しの進行状況を表示
|
||||
3. ブロック完了時に累積結果を受け取る
|
||||
@@ -246,7 +282,7 @@ Workerは内部でイベントを処理し結果を返しますが、UIへのス
|
||||
Worker APIは **Timeline層のHandler機構の薄いラッパー** として設計します。
|
||||
|
||||
| 層 | 目的 | 提供するもの |
|
||||
|---|------|-------------|
|
||||
| ------------------------ | ------------------ | ---------------------------------- |
|
||||
| **Handler (Timeline層)** | 内部実装、役割分離 | スコープ管理 + Deltaイベント |
|
||||
| **Worker Event API** | ユーザー向け利便性 | Handler露出 + Completeイベント追加 |
|
||||
|
||||
@@ -429,8 +465,8 @@ impl<C: LlmClient> Worker<C> {
|
||||
### 設計上のポイント
|
||||
|
||||
1. **Handlerの再利用**: 既存のHandler traitをそのまま活用
|
||||
2. **スコープ管理の維持**: ブロックイベントはStart→Delta→Endのライフサイクルを保持
|
||||
2. **スコープ管理の維持**:
|
||||
ブロックイベントはStart→Delta→Endのライフサイクルを保持
|
||||
3. **選択的購読**: on_*で必要なイベントだけ、またはSubscriberで一括
|
||||
4. **累積イベントの追加**: Worker層でComplete系イベントを追加提供
|
||||
5. **後方互換性**: 従来の`run()`も引き続き使用可能
|
||||
|
||||
|
||||
Generated
+3
-3
@@ -35,11 +35,11 @@
|
||||
},
|
||||
"nixpkgs": {
|
||||
"locked": {
|
||||
"lastModified": 1767116409,
|
||||
"narHash": "sha256-5vKw92l1GyTnjoLzEagJy5V5mDFck72LiQWZSOnSicw=",
|
||||
"lastModified": 1771369470,
|
||||
"narHash": "sha256-0NBlEBKkN3lufyvFegY4TYv5mCNHbi5OmBDrzihbBMQ=",
|
||||
"owner": "nixos",
|
||||
"repo": "nixpkgs",
|
||||
"rev": "cad22e7d996aea55ecab064e84834289143e44a0",
|
||||
"rev": "0182a361324364ae3f436a63005877674cf45efb",
|
||||
"type": "github"
|
||||
},
|
||||
"original": {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
[package]
|
||||
name = "worker-macros"
|
||||
version = "0.1.0"
|
||||
name = "llm-worker-macros"
|
||||
description = "llm-worker's proc macros"
|
||||
version = "0.2.0"
|
||||
publish.workspace = true
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
@@ -1,7 +1,7 @@
|
||||
//! worker-macros - Tool生成用手続きマクロ
|
||||
//! llm-worker-macros - Procedural macros for Tool generation
|
||||
//!
|
||||
//! `#[tool_registry]` と `#[tool]` マクロを提供し、
|
||||
//! ユーザー定義のメソッドから `Tool` トレイト実装を自動生成する。
|
||||
//! Provides `#[tool_registry]` and `#[tool]` macros to
|
||||
//! automatically generate `Tool` trait implementations from user-defined methods.
|
||||
|
||||
use proc_macro::TokenStream;
|
||||
use quote::{format_ident, quote};
|
||||
@@ -9,22 +9,22 @@ use syn::{
|
||||
Attribute, FnArg, ImplItem, ItemImpl, Lit, Meta, Pat, ReturnType, Type, parse_macro_input,
|
||||
};
|
||||
|
||||
/// `impl` ブロックに付与し、内部の `#[tool]` 属性がついたメソッドからツールを生成するマクロ。
|
||||
/// Macro applied to an `impl` block that generates tools from methods marked with `#[tool]`.
|
||||
///
|
||||
/// # Example
|
||||
/// ```ignore
|
||||
/// #[tool_registry]
|
||||
/// impl MyApp {
|
||||
/// /// ユーザー情報を取得する
|
||||
/// /// 指定されたIDのユーザーをDBから検索します。
|
||||
/// /// Get user information
|
||||
/// /// Retrieves a user from the database by their ID.
|
||||
/// #[tool]
|
||||
/// async fn get_user(&self, user_id: String) -> Result<User, Error> { ... }
|
||||
/// }
|
||||
/// ```
|
||||
///
|
||||
/// これにより以下が生成されます:
|
||||
/// - `GetUserArgs` 構造体(引数用)
|
||||
/// - `Tool_get_user` 構造体(Toolラッパー)
|
||||
/// This generates:
|
||||
/// - `GetUserArgs` struct (for arguments)
|
||||
/// - `Tool_get_user` struct (Tool wrapper)
|
||||
/// - `impl Tool for Tool_get_user`
|
||||
/// - `impl MyApp { fn get_user_tool(&self) -> Tool_get_user }`
|
||||
#[proc_macro_attribute]
|
||||
@@ -36,14 +36,14 @@ pub fn tool_registry(_attr: TokenStream, item: TokenStream) -> TokenStream {
|
||||
|
||||
for item in &mut impl_block.items {
|
||||
if let ImplItem::Fn(method) = item {
|
||||
// #[tool] 属性を探す
|
||||
// Look for #[tool] attribute
|
||||
let mut is_tool = false;
|
||||
|
||||
// 属性を走査してtoolがあるか確認し、削除する
|
||||
// Iterate through attributes to check for tool and remove it
|
||||
method.attrs.retain(|attr| {
|
||||
if attr.path().is_ident("tool") {
|
||||
is_tool = true;
|
||||
false // 属性を削除
|
||||
false // Remove the attribute
|
||||
} else {
|
||||
true
|
||||
}
|
||||
@@ -65,7 +65,7 @@ pub fn tool_registry(_attr: TokenStream, item: TokenStream) -> TokenStream {
|
||||
TokenStream::from(expanded)
|
||||
}
|
||||
|
||||
/// ドキュメントコメントから説明文を抽出
|
||||
/// Extract description from doc comments
|
||||
fn extract_doc_comment(attrs: &[Attribute]) -> String {
|
||||
let mut lines = Vec::new();
|
||||
|
||||
@@ -75,7 +75,7 @@ fn extract_doc_comment(attrs: &[Attribute]) -> String {
|
||||
if let syn::Expr::Lit(expr_lit) = &meta.value {
|
||||
if let Lit::Str(lit_str) = &expr_lit.lit {
|
||||
let line = lit_str.value();
|
||||
// 先頭の空白を1つだけ除去(/// の後のスペース)
|
||||
// Remove only the leading space (after ///)
|
||||
let trimmed = line.strip_prefix(' ').unwrap_or(&line);
|
||||
lines.push(trimmed.to_string());
|
||||
}
|
||||
@@ -87,7 +87,7 @@ fn extract_doc_comment(attrs: &[Attribute]) -> String {
|
||||
lines.join("\n")
|
||||
}
|
||||
|
||||
/// #[description = "..."] 属性から説明を抽出
|
||||
/// Extract description from #[description = "..."] attribute
|
||||
fn extract_description_attr(attrs: &[syn::Attribute]) -> Option<String> {
|
||||
for attr in attrs {
|
||||
if attr.path().is_ident("description") {
|
||||
@@ -103,19 +103,19 @@ fn extract_description_attr(attrs: &[syn::Attribute]) -> Option<String> {
|
||||
None
|
||||
}
|
||||
|
||||
/// メソッドからTool実装を生成
|
||||
/// Generate Tool implementation from a method
|
||||
fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::TokenStream {
|
||||
let sig = &method.sig;
|
||||
let method_name = &sig.ident;
|
||||
let tool_name = method_name.to_string();
|
||||
|
||||
// 構造体名を生成(PascalCase変換)
|
||||
// Generate struct names (convert to PascalCase)
|
||||
let pascal_name = to_pascal_case(&method_name.to_string());
|
||||
let tool_struct_name = format_ident!("Tool{}", pascal_name);
|
||||
let args_struct_name = format_ident!("{}Args", pascal_name);
|
||||
let factory_name = format_ident!("{}_tool", method_name);
|
||||
let definition_name = format_ident!("{}_definition", method_name);
|
||||
|
||||
// ドキュメントコメントから説明を取得
|
||||
// Get description from doc comments
|
||||
let description = extract_doc_comment(&method.attrs);
|
||||
let description = if description.is_empty() {
|
||||
format!("Tool: {}", tool_name)
|
||||
@@ -123,7 +123,7 @@ fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::
|
||||
description
|
||||
};
|
||||
|
||||
// 引数を解析(selfを除く)
|
||||
// Parse arguments (excluding self)
|
||||
let args: Vec<_> = sig
|
||||
.inputs
|
||||
.iter()
|
||||
@@ -131,12 +131,12 @@ fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::
|
||||
if let FnArg::Typed(pat_type) = arg {
|
||||
Some(pat_type)
|
||||
} else {
|
||||
None // selfを除外
|
||||
None // Exclude self
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
// 引数構造体のフィールドを生成
|
||||
// Generate argument struct fields
|
||||
let arg_fields: Vec<_> = args
|
||||
.iter()
|
||||
.map(|pat_type| {
|
||||
@@ -144,14 +144,14 @@ fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::
|
||||
let ty = &pat_type.ty;
|
||||
let desc = extract_description_attr(&pat_type.attrs);
|
||||
|
||||
// パターンから識別子を抽出
|
||||
// Extract identifier from pattern
|
||||
let field_name = if let Pat::Ident(pat_ident) = pat.as_ref() {
|
||||
&pat_ident.ident
|
||||
} else {
|
||||
panic!("Only simple identifiers are supported for tool arguments");
|
||||
};
|
||||
|
||||
// #[description] があればschemarsのdocに変換
|
||||
// Convert #[description] to schemars doc if present
|
||||
if let Some(desc_str) = desc {
|
||||
quote! {
|
||||
#[schemars(description = #desc_str)]
|
||||
@@ -165,7 +165,7 @@ fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::
|
||||
})
|
||||
.collect();
|
||||
|
||||
// execute内で引数を展開するコード
|
||||
// Code to expand arguments in execute
|
||||
let arg_names: Vec<_> = args
|
||||
.iter()
|
||||
.map(|pat_type| {
|
||||
@@ -178,22 +178,22 @@ fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::
|
||||
})
|
||||
.collect();
|
||||
|
||||
// メソッドが非同期かどうか
|
||||
// Check if method is async
|
||||
let is_async = sig.asyncness.is_some();
|
||||
|
||||
// 戻り値の型を解析してResult判定
|
||||
// Parse return type and determine if Result
|
||||
let awaiter = if is_async {
|
||||
quote! { .await }
|
||||
} else {
|
||||
quote! {}
|
||||
};
|
||||
|
||||
// 戻り値がResultかどうかを判定
|
||||
// Determine if return type is Result
|
||||
let result_handling = if is_result_type(&sig.output) {
|
||||
quote! {
|
||||
match result {
|
||||
Ok(val) => Ok(format!("{:?}", val)),
|
||||
Err(e) => Err(::worker::tool::ToolError::ExecutionFailed(format!("{}", e))),
|
||||
Err(e) => Err(::llm_worker::tool::ToolError::ExecutionFailed(format!("{}", e))),
|
||||
}
|
||||
}
|
||||
} else {
|
||||
@@ -202,7 +202,7 @@ fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::
|
||||
}
|
||||
};
|
||||
|
||||
// 引数がない場合は空のArgs構造体を作成
|
||||
// Create empty Args struct if no arguments
|
||||
let args_struct_def = if arg_fields.is_empty() {
|
||||
quote! {
|
||||
#[derive(serde::Deserialize, schemars::JsonSchema)]
|
||||
@@ -217,10 +217,10 @@ fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::
|
||||
}
|
||||
};
|
||||
|
||||
// 引数がない場合のexecute処理
|
||||
// Execute body handling for no arguments case
|
||||
let execute_body = if args.is_empty() {
|
||||
quote! {
|
||||
// 引数なしでも空のJSONオブジェクトを許容
|
||||
// Allow empty JSON object even with no arguments
|
||||
let _: #args_struct_name = serde_json::from_str(input_json)
|
||||
.unwrap_or(#args_struct_name {});
|
||||
|
||||
@@ -230,7 +230,7 @@ fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::
|
||||
} else {
|
||||
quote! {
|
||||
let args: #args_struct_name = serde_json::from_str(input_json)
|
||||
.map_err(|e| ::worker::tool::ToolError::InvalidArgument(e.to_string()))?;
|
||||
.map_err(|e| ::llm_worker::tool::ToolError::InvalidArgument(e.to_string()))?;
|
||||
|
||||
let result = self.ctx.#method_name(#(#arg_names),*)#awaiter;
|
||||
#result_handling
|
||||
@@ -246,41 +246,36 @@ fn generate_tool_impl(self_ty: &Type, method: &syn::ImplItemFn) -> proc_macro2::
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl ::worker::tool::Tool for #tool_struct_name {
|
||||
fn name(&self) -> &str {
|
||||
#tool_name
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
#description
|
||||
}
|
||||
|
||||
fn input_schema(&self) -> serde_json::Value {
|
||||
let schema = schemars::schema_for!(#args_struct_name);
|
||||
serde_json::to_value(schema).unwrap_or(serde_json::json!({}))
|
||||
}
|
||||
|
||||
async fn execute(&self, input_json: &str) -> Result<String, ::worker::tool::ToolError> {
|
||||
impl ::llm_worker::tool::Tool for #tool_struct_name {
|
||||
async fn execute(&self, input_json: &str) -> Result<String, ::llm_worker::tool::ToolError> {
|
||||
#execute_body
|
||||
}
|
||||
}
|
||||
|
||||
impl #self_ty {
|
||||
pub fn #factory_name(&self) -> #tool_struct_name {
|
||||
#tool_struct_name {
|
||||
ctx: self.clone()
|
||||
}
|
||||
/// Get ToolDefinition (for registering with Worker)
|
||||
pub fn #definition_name(&self) -> ::llm_worker::tool::ToolDefinition {
|
||||
let ctx = self.clone();
|
||||
::std::sync::Arc::new(move || {
|
||||
let schema = schemars::schema_for!(#args_struct_name);
|
||||
let meta = ::llm_worker::tool::ToolMeta::new(#tool_name)
|
||||
.description(#description)
|
||||
.input_schema(serde_json::to_value(schema).unwrap_or(serde_json::json!({})));
|
||||
let tool: ::std::sync::Arc<dyn ::llm_worker::tool::Tool> =
|
||||
::std::sync::Arc::new(#tool_struct_name { ctx: ctx.clone() });
|
||||
(meta, tool)
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 戻り値の型がResultかどうかを判定
|
||||
/// Determine if return type is Result
|
||||
fn is_result_type(return_type: &ReturnType) -> bool {
|
||||
match return_type {
|
||||
ReturnType::Default => false,
|
||||
ReturnType::Type(_, ty) => {
|
||||
// Type::Pathの場合、最後のセグメントが"Result"かチェック
|
||||
// For Type::Path, check if last segment is "Result"
|
||||
if let Type::Path(type_path) = ty.as_ref() {
|
||||
if let Some(segment) = type_path.path.segments.last() {
|
||||
return segment.ident == "Result";
|
||||
@@ -291,7 +286,7 @@ fn is_result_type(return_type: &ReturnType) -> bool {
|
||||
}
|
||||
}
|
||||
|
||||
/// snake_case を PascalCase に変換
|
||||
/// Convert snake_case to PascalCase
|
||||
fn to_pascal_case(s: &str) -> String {
|
||||
s.split('_')
|
||||
.map(|part| {
|
||||
@@ -304,20 +299,20 @@ fn to_pascal_case(s: &str) -> String {
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// マーカー属性。`tool_registry` によって処理されるため、ここでは何もしない。
|
||||
/// Marker attribute. Does nothing here as it's processed by `tool_registry`.
|
||||
#[proc_macro_attribute]
|
||||
pub fn tool(_attr: TokenStream, item: TokenStream) -> TokenStream {
|
||||
item
|
||||
}
|
||||
|
||||
/// 引数属性用のマーカー。パース時に`tool_registry`で解釈される。
|
||||
/// Marker for argument attributes. Interpreted by `tool_registry` during parsing.
|
||||
///
|
||||
/// # Example
|
||||
/// ```ignore
|
||||
/// #[tool]
|
||||
/// async fn get_user(
|
||||
/// &self,
|
||||
/// #[description = "取得したいユーザーのID"] user_id: String
|
||||
/// #[description = "The ID of the user to retrieve"] user_id: String
|
||||
/// ) -> Result<User, Error> { ... }
|
||||
/// ```
|
||||
#[proc_macro_attribute]
|
||||
@@ -1,7 +1,7 @@
|
||||
[package]
|
||||
name = "worker"
|
||||
description = ""
|
||||
version = "0.1.0"
|
||||
name = "llm-worker"
|
||||
description = "A library for building autonomous LLM-powered systems"
|
||||
version = "0.2.1"
|
||||
publish.workspace = true
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
@@ -15,9 +15,10 @@ tracing = "0.1"
|
||||
async-trait = "0.1"
|
||||
futures = "0.3"
|
||||
tokio = { version = "1.49", features = ["macros", "rt-multi-thread"] }
|
||||
reqwest = { version = "0.13.1", default-features = false, features = ["stream", "json", "native-tls"] }
|
||||
tokio-util = "0.7"
|
||||
reqwest = { version = "0.13.1", default-features = false, features = ["stream", "json", "native-tls", "http2"] }
|
||||
eventsource-stream = "0.2"
|
||||
worker-macros = { path = "../worker-macros", version = "0.1" }
|
||||
llm-worker-macros = { path = "../llm-worker-macros", version = "0.2" }
|
||||
|
||||
[dev-dependencies]
|
||||
clap = { version = "4.5", features = ["derive", "env"] }
|
||||
@@ -25,3 +26,4 @@ schemars = "1.2"
|
||||
tempfile = "3.24"
|
||||
dotenv = "0.15"
|
||||
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
|
||||
trybuild = "1.0.116"
|
||||
+12
-12
@@ -1,18 +1,18 @@
|
||||
//! テストフィクスチャ記録ツール
|
||||
//! Test fixture recording tool
|
||||
//!
|
||||
//! 定義されたシナリオのAPIレスポンスを記録する。
|
||||
//! Records API responses for defined scenarios.
|
||||
//!
|
||||
//! ## 使用方法
|
||||
//! ## Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # 利用可能なシナリオを表示
|
||||
//! # Show available scenarios
|
||||
//! cargo run --example record_test_fixtures
|
||||
//!
|
||||
//! # 特定のシナリオを記録
|
||||
//! # Record specific scenario
|
||||
//! ANTHROPIC_API_KEY=your-key cargo run --example record_test_fixtures -- simple_text
|
||||
//! ANTHROPIC_API_KEY=your-key cargo run --example record_test_fixtures -- tool_call
|
||||
//!
|
||||
//! # 全シナリオを記録
|
||||
//! # Record all scenarios
|
||||
//! ANTHROPIC_API_KEY=your-key cargo run --example record_test_fixtures -- --all
|
||||
//! ```
|
||||
|
||||
@@ -20,9 +20,9 @@ mod recorder;
|
||||
mod scenarios;
|
||||
|
||||
use clap::{Parser, ValueEnum};
|
||||
use worker::llm_client::providers::anthropic::AnthropicClient;
|
||||
use worker::llm_client::providers::gemini::GeminiClient;
|
||||
use worker::llm_client::providers::openai::OpenAIClient;
|
||||
use llm_worker::llm_client::providers::anthropic::AnthropicClient;
|
||||
use llm_worker::llm_client::providers::gemini::GeminiClient;
|
||||
use llm_worker::llm_client::providers::openai::OpenAIClient;
|
||||
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(author, version, about, long_about = None)]
|
||||
@@ -101,7 +101,7 @@ async fn run_scenario_with_ollama(
|
||||
subdir: &str,
|
||||
model: Option<String>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
use worker::llm_client::providers::ollama::OllamaClient;
|
||||
use llm_worker::llm_client::providers::ollama::OllamaClient;
|
||||
// Ollama typically runs local, no key needed or placeholder
|
||||
let model = model.as_deref().unwrap_or("llama3"); // default example
|
||||
let client = OllamaClient::new(model); // base_url placeholder, handled by client default
|
||||
@@ -193,8 +193,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
ClientType::Ollama => "ollama",
|
||||
};
|
||||
|
||||
// シナリオのフィルタリングは main.rs のロジックで実行済み
|
||||
// ここでは単純なループで実行
|
||||
// Scenario filtering is already done in main.rs logic
|
||||
// Here we just execute in a simple loop
|
||||
for scenario in scenarios_to_run {
|
||||
match args.client {
|
||||
ClientType::Anthropic => {
|
||||
+8
-8
@@ -1,6 +1,6 @@
|
||||
//! テストフィクスチャ記録機構
|
||||
//! Test fixture recording mechanism
|
||||
//!
|
||||
//! イベントをJSONLフォーマットでファイルに保存する
|
||||
//! Saves events to files in JSONL format
|
||||
|
||||
use std::fs::{self, File};
|
||||
use std::io::{BufWriter, Write};
|
||||
@@ -8,9 +8,9 @@ use std::path::Path;
|
||||
use std::time::{Instant, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use futures::StreamExt;
|
||||
use worker::llm_client::{LlmClient, Request};
|
||||
use llm_worker::llm_client::{LlmClient, Request};
|
||||
|
||||
/// 記録されたイベント
|
||||
/// Recorded event
|
||||
#[derive(Debug, serde::Serialize, serde::Deserialize)]
|
||||
pub struct RecordedEvent {
|
||||
pub elapsed_ms: u64,
|
||||
@@ -18,7 +18,7 @@ pub struct RecordedEvent {
|
||||
pub data: String,
|
||||
}
|
||||
|
||||
/// セッションメタデータ
|
||||
/// Session metadata
|
||||
#[derive(Debug, serde::Serialize, serde::Deserialize)]
|
||||
pub struct SessionMetadata {
|
||||
pub timestamp: u64,
|
||||
@@ -26,7 +26,7 @@ pub struct SessionMetadata {
|
||||
pub description: String,
|
||||
}
|
||||
|
||||
/// イベントシーケンスをファイルに保存
|
||||
/// Save event sequence to file
|
||||
pub fn save_fixture(
|
||||
path: impl AsRef<Path>,
|
||||
metadata: &SessionMetadata,
|
||||
@@ -43,7 +43,7 @@ pub fn save_fixture(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// リクエストを送信してイベントを記録
|
||||
/// Send request and record events
|
||||
pub async fn record_request<C: LlmClient>(
|
||||
client: &C,
|
||||
request: Request,
|
||||
@@ -78,7 +78,7 @@ pub async fn record_request<C: LlmClient>(
|
||||
}
|
||||
}
|
||||
|
||||
// 保存
|
||||
// Save
|
||||
let fixtures_dir = Path::new("worker/tests/fixtures").join(subdir);
|
||||
fs::create_dir_all(&fixtures_dir)?;
|
||||
|
||||
+11
-11
@@ -1,20 +1,20 @@
|
||||
//! テストフィクスチャ用リクエスト定義
|
||||
//! Test fixture request definitions
|
||||
//!
|
||||
//! 各シナリオのリクエストと出力ファイル名を定義
|
||||
//! Defines requests and output file names for each scenario
|
||||
|
||||
use worker::llm_client::{Request, ToolDefinition};
|
||||
use llm_worker::llm_client::{Request, ToolDefinition};
|
||||
|
||||
/// テストシナリオ
|
||||
/// Test scenario
|
||||
pub struct TestScenario {
|
||||
/// シナリオ名(説明)
|
||||
/// Scenario name (description)
|
||||
pub name: &'static str,
|
||||
/// 出力ファイル名(拡張子なし)
|
||||
/// Output file name (without extension)
|
||||
pub output_name: &'static str,
|
||||
/// リクエスト
|
||||
/// Request
|
||||
pub request: Request,
|
||||
}
|
||||
|
||||
/// 全てのテストシナリオを取得
|
||||
/// Get all test scenarios
|
||||
pub fn scenarios() -> Vec<TestScenario> {
|
||||
vec![
|
||||
simple_text_scenario(),
|
||||
@@ -23,7 +23,7 @@ pub fn scenarios() -> Vec<TestScenario> {
|
||||
]
|
||||
}
|
||||
|
||||
/// シンプルなテキストレスポンス
|
||||
/// Simple text response
|
||||
fn simple_text_scenario() -> TestScenario {
|
||||
TestScenario {
|
||||
name: "Simple text response",
|
||||
@@ -35,7 +35,7 @@ fn simple_text_scenario() -> TestScenario {
|
||||
}
|
||||
}
|
||||
|
||||
/// ツール呼び出しを含むレスポンス
|
||||
/// Response with tool call
|
||||
fn tool_call_scenario() -> TestScenario {
|
||||
let get_weather_tool = ToolDefinition::new("get_weather")
|
||||
.description("Get the current weather for a city")
|
||||
@@ -61,7 +61,7 @@ fn tool_call_scenario() -> TestScenario {
|
||||
}
|
||||
}
|
||||
|
||||
/// 長文生成シナリオ
|
||||
/// Long text generation scenario
|
||||
fn long_text_scenario() -> TestScenario {
|
||||
TestScenario {
|
||||
name: "Long text response",
|
||||
@@ -0,0 +1,71 @@
|
||||
//! Worker cancellation demo
|
||||
//!
|
||||
//! Example of cancelling from another thread during streaming
|
||||
|
||||
use llm_worker::llm_client::providers::anthropic::AnthropicClient;
|
||||
use llm_worker::{Worker, WorkerResult};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
// Load .env file
|
||||
dotenv::dotenv().ok();
|
||||
|
||||
// Initialize logging
|
||||
tracing_subscriber::fmt()
|
||||
.with_env_filter(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env()
|
||||
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")),
|
||||
)
|
||||
.init();
|
||||
|
||||
let api_key =
|
||||
std::env::var("ANTHROPIC_API_KEY").expect("ANTHROPIC_API_KEY environment variable not set");
|
||||
|
||||
let client = AnthropicClient::new(&api_key, "claude-sonnet-4-20250514");
|
||||
let worker = Arc::new(Mutex::new(Worker::new(client)));
|
||||
|
||||
println!("🚀 Starting Worker...");
|
||||
println!("💡 Will cancel after 2 seconds\n");
|
||||
|
||||
// Get cancel sender first (without holding lock)
|
||||
let cancel_tx = {
|
||||
let w = worker.lock().await;
|
||||
w.cancel_sender()
|
||||
};
|
||||
|
||||
// Task 1: Run Worker
|
||||
let worker_clone = worker.clone();
|
||||
let task = tokio::spawn(async move {
|
||||
let mut w = worker_clone.lock().await;
|
||||
println!("📡 Sending request to LLM...");
|
||||
|
||||
match w.run("Tell me a very long story about a brave knight. Make it as detailed as possible with many paragraphs.").await {
|
||||
Ok(WorkerResult::Finished) => {
|
||||
println!("✅ Task completed normally");
|
||||
}
|
||||
Ok(WorkerResult::Paused) => {
|
||||
println!("⏸️ Task paused");
|
||||
}
|
||||
Err(e) => {
|
||||
println!("❌ Task error: {}", e);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Task 2: Cancel after 2 seconds
|
||||
tokio::spawn(async move {
|
||||
tokio::time::sleep(Duration::from_secs(2)).await;
|
||||
println!("\n🛑 Cancelling worker...");
|
||||
let _ = cancel_tx.send(()).await;
|
||||
});
|
||||
|
||||
// Wait for task completion
|
||||
task.await?;
|
||||
|
||||
println!("\n✨ Demo complete!");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,17 +1,17 @@
|
||||
//! Worker を用いた対話型 CLI クライアント
|
||||
//! Interactive CLI client using Worker
|
||||
//!
|
||||
//! 複数のLLMプロバイダ(Anthropic, Gemini, OpenAI, Ollama)と対話するCLIアプリケーション。
|
||||
//! ツールの登録と実行、ストリーミングレスポンスの表示をデモする。
|
||||
//! A CLI application for interacting with multiple LLM providers (Anthropic, Gemini, OpenAI, Ollama).
|
||||
//! Demonstrates tool registration and execution, and streaming response display.
|
||||
//!
|
||||
//! ## 使用方法
|
||||
//! ## Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # .envファイルにAPIキーを設定
|
||||
//! # Set API keys in .env file
|
||||
//! echo "ANTHROPIC_API_KEY=your-api-key" > .env
|
||||
//! echo "GEMINI_API_KEY=your-api-key" >> .env
|
||||
//! echo "OPENAI_API_KEY=your-api-key" >> .env
|
||||
//!
|
||||
//! # Anthropic (デフォルト)
|
||||
//! # Anthropic (default)
|
||||
//! cargo run --example worker_cli
|
||||
//!
|
||||
//! # Gemini
|
||||
@@ -20,13 +20,13 @@
|
||||
//! # OpenAI
|
||||
//! cargo run --example worker_cli -- --provider openai --model gpt-4o
|
||||
//!
|
||||
//! # Ollama (ローカル)
|
||||
//! # Ollama (local)
|
||||
//! cargo run --example worker_cli -- --provider ollama --model llama3.2
|
||||
//!
|
||||
//! # オプション指定
|
||||
//! # With options
|
||||
//! cargo run --example worker_cli -- --provider anthropic --model claude-3-haiku-20240307 --system "You are a helpful assistant."
|
||||
//!
|
||||
//! # ヘルプ表示
|
||||
//! # Show help
|
||||
//! cargo run --example worker_cli -- --help
|
||||
//! ```
|
||||
|
||||
@@ -39,9 +39,9 @@ use tracing::info;
|
||||
use tracing_subscriber::EnvFilter;
|
||||
|
||||
use clap::{Parser, ValueEnum};
|
||||
use worker::{
|
||||
use llm_worker::{
|
||||
Worker,
|
||||
hook::{ControlFlow, HookError, ToolResult, WorkerHook},
|
||||
hook::{Hook, HookError, PostToolCall, PostToolCallContext, PostToolCallResult},
|
||||
llm_client::{
|
||||
LlmClient,
|
||||
providers::{
|
||||
@@ -51,17 +51,17 @@ use worker::{
|
||||
},
|
||||
timeline::{Handler, TextBlockEvent, TextBlockKind, ToolUseBlockEvent, ToolUseBlockKind},
|
||||
};
|
||||
use worker_macros::tool_registry;
|
||||
use llm_worker_macros::tool_registry;
|
||||
|
||||
// 必要なマクロ展開用インポート
|
||||
// Required imports for macro expansion
|
||||
use schemars;
|
||||
use serde;
|
||||
|
||||
// =============================================================================
|
||||
// プロバイダ定義
|
||||
// Provider Definition
|
||||
// =============================================================================
|
||||
|
||||
/// 利用可能なLLMプロバイダ
|
||||
/// Available LLM providers
|
||||
#[derive(Debug, Clone, Copy, ValueEnum, Default)]
|
||||
enum Provider {
|
||||
/// Anthropic Claude
|
||||
@@ -71,12 +71,12 @@ enum Provider {
|
||||
Gemini,
|
||||
/// OpenAI GPT
|
||||
Openai,
|
||||
/// Ollama (ローカル)
|
||||
/// Ollama (local)
|
||||
Ollama,
|
||||
}
|
||||
|
||||
impl Provider {
|
||||
/// プロバイダのデフォルトモデル
|
||||
/// Default model for the provider
|
||||
fn default_model(&self) -> &'static str {
|
||||
match self {
|
||||
Provider::Anthropic => "claude-sonnet-4-20250514",
|
||||
@@ -86,7 +86,7 @@ impl Provider {
|
||||
}
|
||||
}
|
||||
|
||||
/// プロバイダの表示名
|
||||
/// Display name for the provider
|
||||
fn display_name(&self) -> &'static str {
|
||||
match self {
|
||||
Provider::Anthropic => "Anthropic Claude",
|
||||
@@ -96,78 +96,78 @@ impl Provider {
|
||||
}
|
||||
}
|
||||
|
||||
/// APIキーの環境変数名
|
||||
/// Environment variable name for API key
|
||||
fn env_var_name(&self) -> Option<&'static str> {
|
||||
match self {
|
||||
Provider::Anthropic => Some("ANTHROPIC_API_KEY"),
|
||||
Provider::Gemini => Some("GEMINI_API_KEY"),
|
||||
Provider::Openai => Some("OPENAI_API_KEY"),
|
||||
Provider::Ollama => None, // Ollamaはローカルなので不要
|
||||
Provider::Ollama => None, // Ollama is local, no key needed
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// CLI引数定義
|
||||
// CLI Argument Definition
|
||||
// =============================================================================
|
||||
|
||||
/// 複数のLLMプロバイダに対応した対話型CLIクライアント
|
||||
/// Interactive CLI client supporting multiple LLM providers
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "worker-cli")]
|
||||
#[command(about = "Interactive CLI client for multiple LLM providers using Worker")]
|
||||
#[command(version)]
|
||||
struct Args {
|
||||
/// 使用するプロバイダ
|
||||
/// Provider to use
|
||||
#[arg(long, value_enum, default_value_t = Provider::Anthropic)]
|
||||
provider: Provider,
|
||||
|
||||
/// 使用するモデル名(未指定時はプロバイダのデフォルト)
|
||||
/// Model name to use (defaults to provider's default if not specified)
|
||||
#[arg(short, long)]
|
||||
model: Option<String>,
|
||||
|
||||
/// システムプロンプト
|
||||
/// System prompt
|
||||
#[arg(short, long)]
|
||||
system: Option<String>,
|
||||
|
||||
/// ツールを無効化
|
||||
/// Disable tools
|
||||
#[arg(long, default_value = "false")]
|
||||
no_tools: bool,
|
||||
|
||||
/// 最初のメッセージ(指定するとそれを送信して終了)
|
||||
/// Initial message (if specified, sends it and exits)
|
||||
#[arg(short = 'p', long)]
|
||||
prompt: Option<String>,
|
||||
|
||||
/// APIキー(環境変数より優先)
|
||||
/// API key (takes precedence over environment variable)
|
||||
#[arg(long)]
|
||||
api_key: Option<String>,
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ツール定義
|
||||
// Tool Definition
|
||||
// =============================================================================
|
||||
|
||||
/// アプリケーションコンテキスト
|
||||
/// Application context
|
||||
#[derive(Clone)]
|
||||
struct AppContext;
|
||||
|
||||
#[tool_registry]
|
||||
impl AppContext {
|
||||
/// 現在の日時を取得する
|
||||
/// Get the current date and time
|
||||
///
|
||||
/// システムの現在の日付と時刻を返します。
|
||||
/// Returns the system's current date and time.
|
||||
#[tool]
|
||||
fn get_current_time(&self) -> String {
|
||||
let now = std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs();
|
||||
// シンプルなUnixタイムスタンプからの変換
|
||||
// Simple conversion from Unix timestamp
|
||||
format!("Current Unix timestamp: {}", now)
|
||||
}
|
||||
|
||||
/// 簡単な計算を行う
|
||||
/// Perform a simple calculation
|
||||
///
|
||||
/// 2つの数値の四則演算を実行します。
|
||||
/// Executes arithmetic operations on two numbers.
|
||||
#[tool]
|
||||
fn calculate(&self, a: f64, b: f64, operation: String) -> Result<String, String> {
|
||||
let result = match operation.as_str() {
|
||||
@@ -187,10 +187,10 @@ impl AppContext {
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ストリーミング表示用ハンドラー
|
||||
// Streaming Display Handlers
|
||||
// =============================================================================
|
||||
|
||||
/// テキストをリアルタイムで出力するハンドラー
|
||||
/// Handler that outputs text in real-time
|
||||
struct StreamingPrinter {
|
||||
is_first_delta: Arc<Mutex<bool>>,
|
||||
}
|
||||
@@ -226,7 +226,7 @@ impl Handler<TextBlockKind> for StreamingPrinter {
|
||||
}
|
||||
}
|
||||
|
||||
/// ツール呼び出しを表示するハンドラー
|
||||
/// Handler that displays tool calls
|
||||
struct ToolCallPrinter {
|
||||
call_names: Arc<Mutex<HashMap<String, String>>>,
|
||||
}
|
||||
@@ -270,7 +270,7 @@ impl Handler<ToolUseBlockKind> for ToolCallPrinter {
|
||||
}
|
||||
}
|
||||
|
||||
/// ツール実行結果を表示するHook
|
||||
/// Hook that displays tool execution results
|
||||
struct ToolResultPrinterHook {
|
||||
call_names: Arc<Mutex<HashMap<String, String>>>,
|
||||
}
|
||||
@@ -282,40 +282,37 @@ impl ToolResultPrinterHook {
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl WorkerHook for ToolResultPrinterHook {
|
||||
async fn after_tool_call(
|
||||
&self,
|
||||
tool_result: &mut ToolResult,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
impl Hook<PostToolCall> for ToolResultPrinterHook {
|
||||
async fn call(&self, ctx: &mut PostToolCallContext) -> Result<PostToolCallResult, HookError> {
|
||||
let name = self
|
||||
.call_names
|
||||
.lock()
|
||||
.unwrap()
|
||||
.remove(&tool_result.tool_use_id)
|
||||
.unwrap_or_else(|| tool_result.tool_use_id.clone());
|
||||
.remove(&ctx.result.tool_use_id)
|
||||
.unwrap_or_else(|| ctx.result.tool_use_id.clone());
|
||||
|
||||
if tool_result.is_error {
|
||||
println!(" Result ({}): ❌ {}", name, tool_result.content);
|
||||
if ctx.result.is_error {
|
||||
println!(" Result ({}): ❌ {}", name, ctx.result.content);
|
||||
} else {
|
||||
println!(" Result ({}): ✅ {}", name, tool_result.content);
|
||||
println!(" Result ({}): ✅ {}", name, ctx.result.content);
|
||||
}
|
||||
|
||||
Ok(ControlFlow::Continue)
|
||||
Ok(PostToolCallResult::Continue)
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// クライアント作成
|
||||
// Client Creation
|
||||
// =============================================================================
|
||||
|
||||
/// プロバイダに応じたAPIキーを取得
|
||||
/// Get API key based on provider
|
||||
fn get_api_key(args: &Args) -> Result<String, String> {
|
||||
// CLI引数のAPIキーが優先
|
||||
// CLI argument API key takes precedence
|
||||
if let Some(ref key) = args.api_key {
|
||||
return Ok(key.clone());
|
||||
}
|
||||
|
||||
// プロバイダに応じた環境変数を確認
|
||||
// Check environment variable based on provider
|
||||
if let Some(env_var) = args.provider.env_var_name() {
|
||||
std::env::var(env_var).map_err(|_| {
|
||||
format!(
|
||||
@@ -324,12 +321,12 @@ fn get_api_key(args: &Args) -> Result<String, String> {
|
||||
)
|
||||
})
|
||||
} else {
|
||||
// Ollamaなどはキー不要
|
||||
// Ollama etc. don't need a key
|
||||
Ok(String::new())
|
||||
}
|
||||
}
|
||||
|
||||
/// プロバイダに応じたクライアントを作成
|
||||
/// Create client based on provider
|
||||
fn create_client(args: &Args) -> Result<Box<dyn LlmClient>, String> {
|
||||
let model = args
|
||||
.model
|
||||
@@ -359,17 +356,17 @@ fn create_client(args: &Args) -> Result<Box<dyn LlmClient>, String> {
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// メイン
|
||||
// Main
|
||||
// =============================================================================
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
// .envファイルを読み込む
|
||||
// Load .env file
|
||||
dotenv::dotenv().ok();
|
||||
|
||||
// ロギング初期化
|
||||
// RUST_LOG=debug cargo run --example worker_cli ... で詳細ログ表示
|
||||
// デフォルトは warn レベル、RUST_LOG 環境変数で上書き可能
|
||||
// Initialize logging
|
||||
// Use RUST_LOG=debug cargo run --example worker_cli ... for detailed logs
|
||||
// Default is warn level, can be overridden with RUST_LOG environment variable
|
||||
let filter = EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("warn"));
|
||||
|
||||
tracing_subscriber::fmt()
|
||||
@@ -377,7 +374,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
.with_target(true)
|
||||
.init();
|
||||
|
||||
// CLI引数をパース
|
||||
// Parse CLI arguments
|
||||
let args = Args::parse();
|
||||
|
||||
info!(
|
||||
@@ -386,10 +383,10 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
"Starting worker CLI"
|
||||
);
|
||||
|
||||
// 対話モードかワンショットモードか
|
||||
// Interactive mode or one-shot mode
|
||||
let is_interactive = args.prompt.is_none();
|
||||
|
||||
// モデル名(表示用)
|
||||
// Model name (for display)
|
||||
let model_name = args
|
||||
.model
|
||||
.clone()
|
||||
@@ -419,7 +416,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
println!("─────────────────────────────────────────────────");
|
||||
}
|
||||
|
||||
// クライアント作成
|
||||
// Create client
|
||||
let client = match create_client(&args) {
|
||||
Ok(c) => c,
|
||||
Err(e) => {
|
||||
@@ -428,32 +425,34 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
}
|
||||
};
|
||||
|
||||
// Worker作成
|
||||
// Create Worker
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
let tool_call_names = Arc::new(Mutex::new(HashMap::new()));
|
||||
|
||||
// システムプロンプトを設定
|
||||
// Set system prompt
|
||||
if let Some(ref system_prompt) = args.system {
|
||||
worker.set_system_prompt(system_prompt);
|
||||
}
|
||||
|
||||
// ツール登録(--no-tools でなければ)
|
||||
// Register tools (unless --no-tools)
|
||||
if !args.no_tools {
|
||||
let app = AppContext;
|
||||
worker.register_tool(app.get_current_time_tool());
|
||||
worker.register_tool(app.calculate_tool());
|
||||
worker
|
||||
.register_tool(app.get_current_time_definition())
|
||||
.unwrap();
|
||||
worker.register_tool(app.calculate_definition()).unwrap();
|
||||
}
|
||||
|
||||
// ストリーミング表示用ハンドラーを登録
|
||||
// Register streaming display handlers
|
||||
worker
|
||||
.timeline_mut()
|
||||
.on_text_block(StreamingPrinter::new())
|
||||
.on_tool_use_block(ToolCallPrinter::new(tool_call_names.clone()));
|
||||
|
||||
worker.add_hook(ToolResultPrinterHook::new(tool_call_names));
|
||||
worker.add_post_tool_call_hook(ToolResultPrinterHook::new(tool_call_names));
|
||||
|
||||
// ワンショットモード
|
||||
// One-shot mode
|
||||
if let Some(prompt) = args.prompt {
|
||||
match worker.run(&prompt).await {
|
||||
Ok(_) => {}
|
||||
@@ -466,7 +465,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// 対話ループ
|
||||
// Interactive loop
|
||||
loop {
|
||||
print!("\n👤 You: ");
|
||||
io::stdout().flush()?;
|
||||
@@ -484,7 +483,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
break;
|
||||
}
|
||||
|
||||
// Workerを実行(Workerが履歴を管理)
|
||||
// Run Worker (Worker manages history)
|
||||
match worker.run(input).await {
|
||||
Ok(_) => {}
|
||||
Err(e) => {
|
||||
@@ -1,6 +1,6 @@
|
||||
//! Worker層の公開イベント型
|
||||
//! Public event types for Worker layer
|
||||
//!
|
||||
//! 外部利用者に公開するためのイベント表現。
|
||||
//! Event representation exposed to external users.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
@@ -8,38 +8,38 @@ use serde::{Deserialize, Serialize};
|
||||
// Core Event Types (from llm_client layer)
|
||||
// =============================================================================
|
||||
|
||||
/// LLMからのストリーミングイベント
|
||||
/// Streaming events from LLM
|
||||
///
|
||||
/// 各LLMプロバイダからのレスポンスは、この`Event`のストリームとして
|
||||
/// 統一的に処理されます。
|
||||
/// Responses from each LLM provider are processed uniformly
|
||||
/// as a stream of `Event`.
|
||||
///
|
||||
/// # イベントの種類
|
||||
/// # Event Types
|
||||
///
|
||||
/// - **メタイベント**: `Ping`, `Usage`, `Status`, `Error`
|
||||
/// - **ブロックイベント**: `BlockStart`, `BlockDelta`, `BlockStop`, `BlockAbort`
|
||||
/// - **Meta events**: `Ping`, `Usage`, `Status`, `Error`
|
||||
/// - **Block events**: `BlockStart`, `BlockDelta`, `BlockStop`, `BlockAbort`
|
||||
///
|
||||
/// # ブロックのライフサイクル
|
||||
/// # Block Lifecycle
|
||||
///
|
||||
/// テキストやツール呼び出しは、`BlockStart` → `BlockDelta`(複数) → `BlockStop`
|
||||
/// の順序でイベントが発生します。
|
||||
/// Text and tool calls have events in the order of
|
||||
/// `BlockStart` → `BlockDelta`(multiple) → `BlockStop`.
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub enum Event {
|
||||
/// ハートビート
|
||||
/// Heartbeat
|
||||
Ping(PingEvent),
|
||||
/// トークン使用量
|
||||
/// Token usage
|
||||
Usage(UsageEvent),
|
||||
/// ストリームのステータス変化
|
||||
/// Stream status change
|
||||
Status(StatusEvent),
|
||||
/// エラー発生
|
||||
/// Error occurred
|
||||
Error(ErrorEvent),
|
||||
|
||||
/// ブロック開始(テキスト、ツール使用等)
|
||||
/// Block start (text, tool use, etc.)
|
||||
BlockStart(BlockStart),
|
||||
/// ブロックの差分データ
|
||||
/// Block delta data
|
||||
BlockDelta(BlockDelta),
|
||||
/// ブロック正常終了
|
||||
/// Block normal end
|
||||
BlockStop(BlockStop),
|
||||
/// ブロック中断
|
||||
/// Block abort
|
||||
BlockAbort(BlockAbort),
|
||||
}
|
||||
|
||||
@@ -47,47 +47,47 @@ pub enum Event {
|
||||
// Meta Events
|
||||
// =============================================================================
|
||||
|
||||
/// Pingイベント(ハートビート)
|
||||
/// Ping event (heartbeat)
|
||||
#[derive(Debug, Clone, PartialEq, Default, Serialize, Deserialize)]
|
||||
pub struct PingEvent {
|
||||
pub timestamp: Option<u64>,
|
||||
}
|
||||
|
||||
/// 使用量イベント
|
||||
/// Usage event
|
||||
#[derive(Debug, Clone, PartialEq, Default, Serialize, Deserialize)]
|
||||
pub struct UsageEvent {
|
||||
/// 入力トークン数
|
||||
/// Input token count
|
||||
pub input_tokens: Option<u64>,
|
||||
/// 出力トークン数
|
||||
/// Output token count
|
||||
pub output_tokens: Option<u64>,
|
||||
/// 合計トークン数
|
||||
/// Total token count
|
||||
pub total_tokens: Option<u64>,
|
||||
/// キャッシュ読み込みトークン数
|
||||
/// Cache read token count
|
||||
pub cache_read_input_tokens: Option<u64>,
|
||||
/// キャッシュ作成トークン数
|
||||
/// Cache creation token count
|
||||
pub cache_creation_input_tokens: Option<u64>,
|
||||
}
|
||||
|
||||
/// ステータスイベント
|
||||
/// Status event
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct StatusEvent {
|
||||
pub status: ResponseStatus,
|
||||
}
|
||||
|
||||
/// レスポンスステータス
|
||||
/// Response status
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub enum ResponseStatus {
|
||||
/// ストリーム開始
|
||||
/// Stream started
|
||||
Started,
|
||||
/// 正常完了
|
||||
/// Completed normally
|
||||
Completed,
|
||||
/// キャンセルされた
|
||||
/// Cancelled
|
||||
Cancelled,
|
||||
/// エラー発生
|
||||
/// Error occurred
|
||||
Failed,
|
||||
}
|
||||
|
||||
/// エラーイベント
|
||||
/// Error event
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ErrorEvent {
|
||||
pub code: Option<String>,
|
||||
@@ -98,27 +98,27 @@ pub struct ErrorEvent {
|
||||
// Block Types
|
||||
// =============================================================================
|
||||
|
||||
/// ブロックの種別
|
||||
/// Block type
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
pub enum BlockType {
|
||||
/// テキスト生成
|
||||
/// Text generation
|
||||
Text,
|
||||
/// 思考 (Claude Extended Thinking等)
|
||||
/// Thinking (Claude Extended Thinking, etc.)
|
||||
Thinking,
|
||||
/// ツール呼び出し
|
||||
/// Tool call
|
||||
ToolUse,
|
||||
/// ツール結果
|
||||
/// Tool result
|
||||
ToolResult,
|
||||
}
|
||||
|
||||
/// ブロック開始イベント
|
||||
/// Block start event
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct BlockStart {
|
||||
/// ブロックのインデックス
|
||||
/// Block index
|
||||
pub index: usize,
|
||||
/// ブロックの種別
|
||||
/// Block type
|
||||
pub block_type: BlockType,
|
||||
/// ブロック固有のメタデータ
|
||||
/// Block-specific metadata
|
||||
pub metadata: BlockMetadata,
|
||||
}
|
||||
|
||||
@@ -128,7 +128,7 @@ impl BlockStart {
|
||||
}
|
||||
}
|
||||
|
||||
/// ブロックのメタデータ
|
||||
/// Block metadata
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub enum BlockMetadata {
|
||||
Text,
|
||||
@@ -137,28 +137,28 @@ pub enum BlockMetadata {
|
||||
ToolResult { tool_use_id: String },
|
||||
}
|
||||
|
||||
/// ブロックデルタイベント
|
||||
/// Block delta event
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct BlockDelta {
|
||||
/// ブロックのインデックス
|
||||
/// Block index
|
||||
pub index: usize,
|
||||
/// デルタの内容
|
||||
/// Delta content
|
||||
pub delta: DeltaContent,
|
||||
}
|
||||
|
||||
/// デルタの内容
|
||||
/// Delta content
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub enum DeltaContent {
|
||||
/// テキストデルタ
|
||||
/// Text delta
|
||||
Text(String),
|
||||
/// 思考デルタ
|
||||
/// Thinking delta
|
||||
Thinking(String),
|
||||
/// ツール引数のJSON部分文字列
|
||||
/// JSON substring of tool arguments
|
||||
InputJson(String),
|
||||
}
|
||||
|
||||
impl DeltaContent {
|
||||
/// デルタのブロック種別を取得
|
||||
/// Get block type of the delta
|
||||
pub fn block_type(&self) -> BlockType {
|
||||
match self {
|
||||
DeltaContent::Text(_) => BlockType::Text,
|
||||
@@ -168,14 +168,14 @@ impl DeltaContent {
|
||||
}
|
||||
}
|
||||
|
||||
/// ブロック停止イベント
|
||||
/// Block stop event
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct BlockStop {
|
||||
/// ブロックのインデックス
|
||||
/// Block index
|
||||
pub index: usize,
|
||||
/// ブロックの種別
|
||||
/// Block type
|
||||
pub block_type: BlockType,
|
||||
/// 停止理由
|
||||
/// Stop reason
|
||||
pub stop_reason: Option<StopReason>,
|
||||
}
|
||||
|
||||
@@ -185,14 +185,14 @@ impl BlockStop {
|
||||
}
|
||||
}
|
||||
|
||||
/// ブロック中断イベント
|
||||
/// Block abort event
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct BlockAbort {
|
||||
/// ブロックのインデックス
|
||||
/// Block index
|
||||
pub index: usize,
|
||||
/// ブロックの種別
|
||||
/// Block type
|
||||
pub block_type: BlockType,
|
||||
/// 中断理由
|
||||
/// Abort reason
|
||||
pub reason: String,
|
||||
}
|
||||
|
||||
@@ -202,16 +202,16 @@ impl BlockAbort {
|
||||
}
|
||||
}
|
||||
|
||||
/// 停止理由
|
||||
/// Stop reason
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub enum StopReason {
|
||||
/// 自然終了
|
||||
/// Natural end
|
||||
EndTurn,
|
||||
/// 最大トークン数到達
|
||||
/// Max tokens reached
|
||||
MaxTokens,
|
||||
/// ストップシーケンス到達
|
||||
/// Stop sequence reached
|
||||
StopSequence,
|
||||
/// ツール使用
|
||||
/// Tool use
|
||||
ToolUse,
|
||||
}
|
||||
|
||||
@@ -220,7 +220,7 @@ pub enum StopReason {
|
||||
// =============================================================================
|
||||
|
||||
impl Event {
|
||||
/// テキストブロック開始イベントを作成
|
||||
/// Create text block start event
|
||||
pub fn text_block_start(index: usize) -> Self {
|
||||
Event::BlockStart(BlockStart {
|
||||
index,
|
||||
@@ -229,7 +229,7 @@ impl Event {
|
||||
})
|
||||
}
|
||||
|
||||
/// テキストデルタイベントを作成
|
||||
/// Create text delta event
|
||||
pub fn text_delta(index: usize, text: impl Into<String>) -> Self {
|
||||
Event::BlockDelta(BlockDelta {
|
||||
index,
|
||||
@@ -237,7 +237,7 @@ impl Event {
|
||||
})
|
||||
}
|
||||
|
||||
/// テキストブロック停止イベントを作成
|
||||
/// Create text block stop event
|
||||
pub fn text_block_stop(index: usize, stop_reason: Option<StopReason>) -> Self {
|
||||
Event::BlockStop(BlockStop {
|
||||
index,
|
||||
@@ -246,7 +246,7 @@ impl Event {
|
||||
})
|
||||
}
|
||||
|
||||
/// ツール使用ブロック開始イベントを作成
|
||||
/// Create tool use block start event
|
||||
pub fn tool_use_start(index: usize, id: impl Into<String>, name: impl Into<String>) -> Self {
|
||||
Event::BlockStart(BlockStart {
|
||||
index,
|
||||
@@ -258,7 +258,7 @@ impl Event {
|
||||
})
|
||||
}
|
||||
|
||||
/// ツール引数デルタイベントを作成
|
||||
/// Create tool input delta event
|
||||
pub fn tool_input_delta(index: usize, json: impl Into<String>) -> Self {
|
||||
Event::BlockDelta(BlockDelta {
|
||||
index,
|
||||
@@ -266,7 +266,7 @@ impl Event {
|
||||
})
|
||||
}
|
||||
|
||||
/// ツール使用ブロック停止イベントを作成
|
||||
/// Create tool use block stop event
|
||||
pub fn tool_use_stop(index: usize) -> Self {
|
||||
Event::BlockStop(BlockStop {
|
||||
index,
|
||||
@@ -275,7 +275,7 @@ impl Event {
|
||||
})
|
||||
}
|
||||
|
||||
/// 使用量イベントを作成
|
||||
/// Create usage event
|
||||
pub fn usage(input_tokens: u64, output_tokens: u64) -> Self {
|
||||
Event::Usage(UsageEvent {
|
||||
input_tokens: Some(input_tokens),
|
||||
@@ -286,7 +286,7 @@ impl Event {
|
||||
})
|
||||
}
|
||||
|
||||
/// Pingイベントを作成
|
||||
/// Create ping event
|
||||
pub fn ping() -> Self {
|
||||
Event::Ping(PingEvent { timestamp: None })
|
||||
}
|
||||
@@ -1,8 +1,8 @@
|
||||
//! Handler/Kind型
|
||||
//! Handler/Kind Types
|
||||
//!
|
||||
//! Timeline層でイベントを処理するためのトレイト。
|
||||
//! カスタムハンドラを実装してTimelineに登録することで、
|
||||
//! ストリームイベントを受信できます。
|
||||
//! Traits for processing events in the Timeline layer.
|
||||
//! By implementing custom handlers and registering them with Timeline,
|
||||
//! you can receive stream events.
|
||||
|
||||
use crate::timeline::event::*;
|
||||
|
||||
@@ -10,13 +10,13 @@ use crate::timeline::event::*;
|
||||
// Kind Trait
|
||||
// =============================================================================
|
||||
|
||||
/// イベント種別を定義するマーカートレイト
|
||||
/// Marker trait defining event types
|
||||
///
|
||||
/// 各Kindは対応するイベント型を指定します。
|
||||
/// HandlerはこのKindに対して実装され、同じKindに対して
|
||||
/// 異なるScope型を持つ複数のHandlerを登録できます。
|
||||
/// Each Kind specifies its corresponding event type.
|
||||
/// Handlers are implemented for this Kind, and multiple Handlers
|
||||
/// with different Scope types can be registered for the same Kind.
|
||||
pub trait Kind {
|
||||
/// このKindに対応するイベント型
|
||||
/// Event type corresponding to this Kind
|
||||
type Event;
|
||||
}
|
||||
|
||||
@@ -24,22 +24,22 @@ pub trait Kind {
|
||||
// Handler Trait
|
||||
// =============================================================================
|
||||
|
||||
/// イベントを処理するハンドラトレイト
|
||||
/// Handler trait for processing events
|
||||
///
|
||||
/// 特定の`Kind`に対するイベント処理を定義します。
|
||||
/// `Scope`はブロックのライフサイクル中に保持される状態です。
|
||||
/// Defines event processing for a specific `Kind`.
|
||||
/// `Scope` is state held during the block's lifecycle.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::timeline::{Handler, TextBlockEvent, TextBlockKind};
|
||||
/// use llm_worker::timeline::{Handler, TextBlockEvent, TextBlockKind};
|
||||
///
|
||||
/// struct TextCollector {
|
||||
/// texts: Vec<String>,
|
||||
/// }
|
||||
///
|
||||
/// impl Handler<TextBlockKind> for TextCollector {
|
||||
/// type Scope = String; // ブロックごとのバッファ
|
||||
/// type Scope = String; // Buffer per block
|
||||
///
|
||||
/// fn on_event(&mut self, buffer: &mut String, event: &TextBlockEvent) {
|
||||
/// match event {
|
||||
@@ -53,13 +53,13 @@ pub trait Kind {
|
||||
/// }
|
||||
/// ```
|
||||
pub trait Handler<K: Kind> {
|
||||
/// Handler固有のスコープ型
|
||||
/// Handler-specific scope type
|
||||
///
|
||||
/// ブロック開始時に`Default::default()`で生成され、
|
||||
/// ブロック終了時に破棄されます。
|
||||
/// Generated with `Default::default()` at block start,
|
||||
/// and destroyed at block end.
|
||||
type Scope: Default;
|
||||
|
||||
/// イベントを処理する
|
||||
/// Process the event
|
||||
fn on_event(&mut self, scope: &mut Self::Scope, event: &K::Event);
|
||||
}
|
||||
|
||||
@@ -67,25 +67,25 @@ pub trait Handler<K: Kind> {
|
||||
// Meta Kind Definitions
|
||||
// =============================================================================
|
||||
|
||||
/// Usage Kind - 使用量イベント用
|
||||
/// Usage Kind - for usage events
|
||||
pub struct UsageKind;
|
||||
impl Kind for UsageKind {
|
||||
type Event = UsageEvent;
|
||||
}
|
||||
|
||||
/// Ping Kind - Pingイベント用
|
||||
/// Ping Kind - for ping events
|
||||
pub struct PingKind;
|
||||
impl Kind for PingKind {
|
||||
type Event = PingEvent;
|
||||
}
|
||||
|
||||
/// Status Kind - ステータスイベント用
|
||||
/// Status Kind - for status events
|
||||
pub struct StatusKind;
|
||||
impl Kind for StatusKind {
|
||||
type Event = StatusEvent;
|
||||
}
|
||||
|
||||
/// Error Kind - エラーイベント用
|
||||
/// Error Kind - for error events
|
||||
pub struct ErrorKind;
|
||||
impl Kind for ErrorKind {
|
||||
type Event = ErrorEvent;
|
||||
@@ -95,13 +95,13 @@ impl Kind for ErrorKind {
|
||||
// Block Kind Definitions
|
||||
// =============================================================================
|
||||
|
||||
/// TextBlock Kind - テキストブロック用
|
||||
/// TextBlock Kind - for text blocks
|
||||
pub struct TextBlockKind;
|
||||
impl Kind for TextBlockKind {
|
||||
type Event = TextBlockEvent;
|
||||
}
|
||||
|
||||
/// テキストブロックのイベント
|
||||
/// Text block events
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum TextBlockEvent {
|
||||
Start(TextBlockStart),
|
||||
@@ -120,13 +120,13 @@ pub struct TextBlockStop {
|
||||
pub stop_reason: Option<StopReason>,
|
||||
}
|
||||
|
||||
/// ThinkingBlock Kind - 思考ブロック用
|
||||
/// ThinkingBlock Kind - for thinking blocks
|
||||
pub struct ThinkingBlockKind;
|
||||
impl Kind for ThinkingBlockKind {
|
||||
type Event = ThinkingBlockEvent;
|
||||
}
|
||||
|
||||
/// 思考ブロックのイベント
|
||||
/// Thinking block events
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum ThinkingBlockEvent {
|
||||
Start(ThinkingBlockStart),
|
||||
@@ -144,17 +144,17 @@ pub struct ThinkingBlockStop {
|
||||
pub index: usize,
|
||||
}
|
||||
|
||||
/// ToolUseBlock Kind - ツール使用ブロック用
|
||||
/// ToolUseBlock Kind - for tool use blocks
|
||||
pub struct ToolUseBlockKind;
|
||||
impl Kind for ToolUseBlockKind {
|
||||
type Event = ToolUseBlockEvent;
|
||||
}
|
||||
|
||||
/// ツール使用ブロックのイベント
|
||||
/// Tool use block events
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum ToolUseBlockEvent {
|
||||
Start(ToolUseBlockStart),
|
||||
/// ツール引数のJSON部分文字列
|
||||
/// JSON substring of tool arguments
|
||||
InputJsonDelta(String),
|
||||
Stop(ToolUseBlockStop),
|
||||
}
|
||||
@@ -0,0 +1,233 @@
|
||||
//! Hook-related type definitions
|
||||
//!
|
||||
//! Types used for turn control and intervention in the Worker layer
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
|
||||
// =============================================================================
|
||||
// Hook Event Kinds
|
||||
// =============================================================================
|
||||
|
||||
pub trait HookEventKind: Send + Sync + 'static {
|
||||
type Input;
|
||||
type Output;
|
||||
}
|
||||
|
||||
pub struct OnPromptSubmit;
|
||||
pub struct PreLlmRequest;
|
||||
pub struct PreToolCall;
|
||||
pub struct PostToolCall;
|
||||
pub struct OnTurnEnd;
|
||||
pub struct OnAbort;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum OnPromptSubmitResult {
|
||||
Continue,
|
||||
Cancel(String),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum PreLlmRequestResult {
|
||||
Continue,
|
||||
Cancel(String),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum PreToolCallResult {
|
||||
Continue,
|
||||
Skip,
|
||||
Abort(String),
|
||||
Pause,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum PostToolCallResult {
|
||||
Continue,
|
||||
Abort(String),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum OnTurnEndResult {
|
||||
Finish,
|
||||
ContinueWithMessages(Vec<crate::Item>),
|
||||
Paused,
|
||||
}
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::tool::{Tool, ToolMeta};
|
||||
|
||||
/// Input context for PreToolCall
|
||||
pub struct ToolCallContext {
|
||||
/// Tool call information (modifiable)
|
||||
pub call: ToolCall,
|
||||
/// Tool meta information (immutable)
|
||||
pub meta: ToolMeta,
|
||||
/// Tool instance (for state access)
|
||||
pub tool: Arc<dyn Tool>,
|
||||
}
|
||||
|
||||
/// Input context for PostToolCall
|
||||
pub struct PostToolCallContext {
|
||||
/// Tool call information
|
||||
pub call: ToolCall,
|
||||
/// Tool execution result (modifiable)
|
||||
pub result: ToolResult,
|
||||
/// Tool meta information (immutable)
|
||||
pub meta: ToolMeta,
|
||||
/// Tool instance (for state access)
|
||||
pub tool: Arc<dyn Tool>,
|
||||
}
|
||||
|
||||
impl HookEventKind for OnPromptSubmit {
|
||||
type Input = crate::Item;
|
||||
type Output = OnPromptSubmitResult;
|
||||
}
|
||||
|
||||
impl HookEventKind for PreLlmRequest {
|
||||
type Input = Vec<crate::Item>;
|
||||
type Output = PreLlmRequestResult;
|
||||
}
|
||||
|
||||
impl HookEventKind for PreToolCall {
|
||||
type Input = ToolCallContext;
|
||||
type Output = PreToolCallResult;
|
||||
}
|
||||
|
||||
impl HookEventKind for PostToolCall {
|
||||
type Input = PostToolCallContext;
|
||||
type Output = PostToolCallResult;
|
||||
}
|
||||
|
||||
impl HookEventKind for OnTurnEnd {
|
||||
type Input = Vec<crate::Item>;
|
||||
type Output = OnTurnEndResult;
|
||||
}
|
||||
|
||||
impl HookEventKind for OnAbort {
|
||||
type Input = String;
|
||||
type Output = ();
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Tool Call / Result Types
|
||||
// =============================================================================
|
||||
|
||||
/// Tool call information
|
||||
///
|
||||
/// Represents a ToolUse block from LLM, modifiable in Hook processing
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolCall {
|
||||
/// Tool call ID (used for linking with response)
|
||||
pub id: String,
|
||||
/// Tool name
|
||||
pub name: String,
|
||||
/// Input arguments (JSON)
|
||||
pub input: Value,
|
||||
}
|
||||
|
||||
/// Tool execution result
|
||||
///
|
||||
/// Represents the result after tool execution, modifiable in Hook processing
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolResult {
|
||||
/// Corresponding tool call ID
|
||||
pub tool_use_id: String,
|
||||
/// Result content
|
||||
pub content: String,
|
||||
/// Whether this is an error
|
||||
#[serde(default)]
|
||||
pub is_error: bool,
|
||||
}
|
||||
|
||||
impl ToolResult {
|
||||
/// Create a success result
|
||||
pub fn success(tool_use_id: impl Into<String>, content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
tool_use_id: tool_use_id.into(),
|
||||
content: content.into(),
|
||||
is_error: false,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create an error result
|
||||
pub fn error(tool_use_id: impl Into<String>, content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
tool_use_id: tool_use_id.into(),
|
||||
content: content.into(),
|
||||
is_error: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Hook Error
|
||||
// =============================================================================
|
||||
|
||||
/// Hook error
|
||||
#[derive(Debug, Error)]
|
||||
pub enum HookError {
|
||||
/// Processing was aborted
|
||||
#[error("Aborted: {0}")]
|
||||
Aborted(String),
|
||||
/// Internal error
|
||||
#[error("Hook error: {0}")]
|
||||
Internal(String),
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Hook Trait
|
||||
// =============================================================================
|
||||
|
||||
/// Trait for handling Hook events
|
||||
///
|
||||
/// Each event type has a different return type, constrained via `HookEventKind`.
|
||||
#[async_trait]
|
||||
pub trait Hook<E: HookEventKind>: Send + Sync {
|
||||
async fn call(&self, input: &mut E::Input) -> Result<E::Output, HookError>;
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Hook Registry
|
||||
// =============================================================================
|
||||
|
||||
/// Registry holding all Hooks
|
||||
///
|
||||
/// Used internally by Worker to manage all Hook types.
|
||||
pub struct HookRegistry {
|
||||
/// on_prompt_submit Hook
|
||||
pub(crate) on_prompt_submit: Vec<Box<dyn Hook<OnPromptSubmit>>>,
|
||||
/// pre_llm_request Hook
|
||||
pub(crate) pre_llm_request: Vec<Box<dyn Hook<PreLlmRequest>>>,
|
||||
/// pre_tool_call Hook
|
||||
pub(crate) pre_tool_call: Vec<Box<dyn Hook<PreToolCall>>>,
|
||||
/// post_tool_call Hook
|
||||
pub(crate) post_tool_call: Vec<Box<dyn Hook<PostToolCall>>>,
|
||||
/// on_turn_end Hook
|
||||
pub(crate) on_turn_end: Vec<Box<dyn Hook<OnTurnEnd>>>,
|
||||
/// on_abort Hook
|
||||
pub(crate) on_abort: Vec<Box<dyn Hook<OnAbort>>>,
|
||||
}
|
||||
|
||||
impl Default for HookRegistry {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl HookRegistry {
|
||||
/// Create an empty HookRegistry
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
on_prompt_submit: Vec::new(),
|
||||
pre_llm_request: Vec::new(),
|
||||
pre_tool_call: Vec::new(),
|
||||
post_tool_call: Vec::new(),
|
||||
on_turn_end: Vec::new(),
|
||||
on_abort: Vec::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
//! llm-worker - LLM Worker Library
|
||||
//!
|
||||
//! Provides components for managing interactions with LLMs.
|
||||
//!
|
||||
//! # Main Components
|
||||
//!
|
||||
//! - [`Worker`] - Central component for managing LLM interactions
|
||||
//! - [`tool::Tool`] - Tools that can be invoked by the LLM
|
||||
//! - [`hook::Hook`] - Hooks for intercepting turn progression
|
||||
//! - [`subscriber::WorkerSubscriber`] - Subscribing to streaming events
|
||||
//!
|
||||
//! # Quick Start
|
||||
//!
|
||||
//! ```ignore
|
||||
//! use llm_worker::{Worker, Item};
|
||||
//!
|
||||
//! // Create a Worker
|
||||
//! let mut worker = Worker::new(client)
|
||||
//! .system_prompt("You are a helpful assistant.");
|
||||
//!
|
||||
//! // Register tools (optional)
|
||||
//! // worker.register_tool(my_tool_definition)?;
|
||||
//!
|
||||
//! // Run the interaction
|
||||
//! let history = worker.run("Hello!").await?;
|
||||
//! ```
|
||||
//!
|
||||
//! # Cache Protection
|
||||
//!
|
||||
//! To maximize KV cache hit rate, transition to the locked state
|
||||
//! with [`Worker::lock()`] before execution.
|
||||
//!
|
||||
//! ```ignore
|
||||
//! let mut locked = worker.lock();
|
||||
//! locked.run("user input").await?;
|
||||
//! ```
|
||||
|
||||
mod handler;
|
||||
mod message;
|
||||
mod worker;
|
||||
|
||||
pub mod event;
|
||||
pub mod hook;
|
||||
pub mod llm_client;
|
||||
pub mod state;
|
||||
pub mod subscriber;
|
||||
pub mod timeline;
|
||||
pub mod tool;
|
||||
pub mod tool_server;
|
||||
|
||||
pub use message::{ContentPart, Item, Message, Role};
|
||||
pub use worker::{ToolRegistryError, Worker, WorkerConfig, WorkerError, WorkerResult};
|
||||
@@ -0,0 +1,86 @@
|
||||
//! LLMクライアント共通trait定義
|
||||
|
||||
use std::pin::Pin;
|
||||
|
||||
use crate::llm_client::{ClientError, Request, RequestConfig, event::Event};
|
||||
use async_trait::async_trait;
|
||||
use futures::Stream;
|
||||
|
||||
/// 設定に関する警告
|
||||
///
|
||||
/// プロバイダがサポートしていない設定を使用した場合に返される。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ConfigWarning {
|
||||
/// 設定オプション名
|
||||
pub option_name: &'static str,
|
||||
/// 警告メッセージ
|
||||
pub message: String,
|
||||
}
|
||||
|
||||
impl ConfigWarning {
|
||||
/// 新しい警告を作成
|
||||
pub fn unsupported(option_name: &'static str, provider_name: &str) -> Self {
|
||||
Self {
|
||||
option_name,
|
||||
message: format!(
|
||||
"'{}' is not supported by {} and will be ignored",
|
||||
option_name, provider_name
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ConfigWarning {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
write!(f, "{}: {}", self.option_name, self.message)
|
||||
}
|
||||
}
|
||||
|
||||
/// LLMクライアントのtrait
|
||||
///
|
||||
/// 各プロバイダはこのtraitを実装し、統一されたインターフェースを提供する。
|
||||
#[async_trait]
|
||||
pub trait LlmClient: Send + Sync {
|
||||
/// ストリーミングリクエストを送信し、Eventストリームを返す
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `request` - リクエスト情報
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(Stream)` - イベントストリーム
|
||||
/// * `Err(ClientError)` - エラー
|
||||
async fn stream(
|
||||
&self,
|
||||
request: Request,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<Event, ClientError>> + Send>>, ClientError>;
|
||||
|
||||
/// 設定をバリデーションし、未サポートの設定があれば警告を返す
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `config` - バリデーション対象の設定
|
||||
///
|
||||
/// # Returns
|
||||
/// サポートされていない設定に対する警告のリスト
|
||||
fn validate_config(&self, config: &RequestConfig) -> Vec<ConfigWarning> {
|
||||
// デフォルト実装: 全ての設定をサポート
|
||||
let _ = config;
|
||||
Vec::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// `Box<dyn LlmClient>` に対する `LlmClient` の実装
|
||||
///
|
||||
/// これにより、動的ディスパッチを使用するクライアントも `Worker` で利用可能になる。
|
||||
#[async_trait]
|
||||
impl LlmClient for Box<dyn LlmClient> {
|
||||
async fn stream(
|
||||
&self,
|
||||
request: Request,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<Event, ClientError>> + Send>>, ClientError> {
|
||||
(**self).stream(request).await
|
||||
}
|
||||
|
||||
fn validate_config(&self, config: &RequestConfig) -> Vec<ConfigWarning> {
|
||||
(**self).validate_config(config)
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
//! LLMクライアント層
|
||||
//!
|
||||
//! 各LLMプロバイダと通信し、統一された[`Event`](crate::llm_client::event::Event)
|
||||
//! 各LLMプロバイダと通信し、統一された[`Event`]
|
||||
//! ストリームを出力します。
|
||||
//!
|
||||
//! # サポートするプロバイダ
|
||||
+13
-1
@@ -5,7 +5,8 @@
|
||||
use std::pin::Pin;
|
||||
|
||||
use crate::llm_client::{
|
||||
ClientError, LlmClient, Request, event::Event, scheme::openai::OpenAIScheme,
|
||||
ClientError, ConfigWarning, LlmClient, Request, RequestConfig, event::Event,
|
||||
scheme::openai::OpenAIScheme,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use eventsource_stream::Eventsource;
|
||||
@@ -197,4 +198,15 @@ impl LlmClient for OpenAIClient {
|
||||
|
||||
Ok(Box::pin(stream))
|
||||
}
|
||||
|
||||
fn validate_config(&self, config: &RequestConfig) -> Vec<ConfigWarning> {
|
||||
let mut warnings = Vec::new();
|
||||
|
||||
// OpenAI does not support top_k
|
||||
if config.top_k.is_some() {
|
||||
warnings.push(ConfigWarning::unsupported("top_k", "OpenAI"));
|
||||
}
|
||||
|
||||
warnings
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,323 @@
|
||||
//! Anthropic Request Builder
|
||||
//!
|
||||
//! Converts Open Responses native Item model to Anthropic Messages API format.
|
||||
|
||||
use serde::Serialize;
|
||||
|
||||
use crate::llm_client::{
|
||||
types::{ContentPart, Item, Role, ToolDefinition},
|
||||
Request,
|
||||
};
|
||||
|
||||
use super::AnthropicScheme;
|
||||
|
||||
/// Anthropic API request body
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct AnthropicRequest {
|
||||
pub model: String,
|
||||
pub max_tokens: u32,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub system: Option<String>,
|
||||
pub messages: Vec<AnthropicMessage>,
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub tools: Vec<AnthropicTool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub top_p: Option<f32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub top_k: Option<u32>,
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub stop_sequences: Vec<String>,
|
||||
pub stream: bool,
|
||||
}
|
||||
|
||||
/// Anthropic message
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct AnthropicMessage {
|
||||
pub role: String,
|
||||
pub content: AnthropicContent,
|
||||
}
|
||||
|
||||
/// Anthropic content
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum AnthropicContent {
|
||||
Text(String),
|
||||
Parts(Vec<AnthropicContentPart>),
|
||||
}
|
||||
|
||||
/// Anthropic content part
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub(crate) enum AnthropicContentPart {
|
||||
#[serde(rename = "text")]
|
||||
Text { text: String },
|
||||
#[serde(rename = "tool_use")]
|
||||
ToolUse {
|
||||
id: String,
|
||||
name: String,
|
||||
input: serde_json::Value,
|
||||
},
|
||||
#[serde(rename = "tool_result")]
|
||||
ToolResult { tool_use_id: String, content: String },
|
||||
}
|
||||
|
||||
/// Anthropic tool definition
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct AnthropicTool {
|
||||
pub name: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
pub input_schema: serde_json::Value,
|
||||
}
|
||||
|
||||
impl AnthropicScheme {
|
||||
/// Build Anthropic request from Request
|
||||
pub(crate) fn build_request(&self, model: &str, request: &Request) -> AnthropicRequest {
|
||||
let messages = self.convert_items_to_messages(&request.items);
|
||||
let tools = request.tools.iter().map(|t| self.convert_tool(t)).collect();
|
||||
|
||||
AnthropicRequest {
|
||||
model: model.to_string(),
|
||||
max_tokens: request.config.max_tokens.unwrap_or(4096),
|
||||
system: request.system_prompt.clone(),
|
||||
messages,
|
||||
tools,
|
||||
temperature: request.config.temperature,
|
||||
top_p: request.config.top_p,
|
||||
top_k: request.config.top_k,
|
||||
stop_sequences: request.config.stop_sequences.clone(),
|
||||
stream: true,
|
||||
}
|
||||
}
|
||||
|
||||
/// Convert Open Responses Items to Anthropic Messages
|
||||
///
|
||||
/// Anthropic uses a message-based model where:
|
||||
/// - User messages have role "user"
|
||||
/// - Assistant messages have role "assistant"
|
||||
/// - Tool calls are content parts within assistant messages
|
||||
/// - Tool results are content parts within user messages
|
||||
fn convert_items_to_messages(&self, items: &[Item]) -> Vec<AnthropicMessage> {
|
||||
let mut messages = Vec::new();
|
||||
let mut pending_assistant_parts: Vec<AnthropicContentPart> = Vec::new();
|
||||
let mut pending_user_parts: Vec<AnthropicContentPart> = Vec::new();
|
||||
|
||||
for item in items {
|
||||
match item {
|
||||
Item::Message { role, content, .. } => {
|
||||
// Flush pending parts before a new message
|
||||
self.flush_pending_parts(
|
||||
&mut messages,
|
||||
&mut pending_assistant_parts,
|
||||
&mut pending_user_parts,
|
||||
);
|
||||
|
||||
let anthropic_role = match role {
|
||||
Role::User => "user",
|
||||
Role::Assistant => "assistant",
|
||||
Role::System => continue, // Skip system role items
|
||||
};
|
||||
|
||||
let parts: Vec<AnthropicContentPart> = content
|
||||
.iter()
|
||||
.map(|p| match p {
|
||||
ContentPart::InputText { text } => {
|
||||
AnthropicContentPart::Text { text: text.clone() }
|
||||
}
|
||||
ContentPart::OutputText { text } => {
|
||||
AnthropicContentPart::Text { text: text.clone() }
|
||||
}
|
||||
ContentPart::Refusal { refusal } => {
|
||||
AnthropicContentPart::Text {
|
||||
text: refusal.clone(),
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
if parts.len() == 1 {
|
||||
if let AnthropicContentPart::Text { text } = &parts[0] {
|
||||
messages.push(AnthropicMessage {
|
||||
role: anthropic_role.to_string(),
|
||||
content: AnthropicContent::Text(text.clone()),
|
||||
});
|
||||
} else {
|
||||
messages.push(AnthropicMessage {
|
||||
role: anthropic_role.to_string(),
|
||||
content: AnthropicContent::Parts(parts),
|
||||
});
|
||||
}
|
||||
} else {
|
||||
messages.push(AnthropicMessage {
|
||||
role: anthropic_role.to_string(),
|
||||
content: AnthropicContent::Parts(parts),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
Item::FunctionCall {
|
||||
call_id,
|
||||
name,
|
||||
arguments,
|
||||
..
|
||||
} => {
|
||||
// Flush pending user parts first
|
||||
if !pending_user_parts.is_empty() {
|
||||
messages.push(AnthropicMessage {
|
||||
role: "user".to_string(),
|
||||
content: AnthropicContent::Parts(std::mem::take(
|
||||
&mut pending_user_parts,
|
||||
)),
|
||||
});
|
||||
}
|
||||
|
||||
// Parse arguments JSON string to Value
|
||||
let input = serde_json::from_str(arguments)
|
||||
.unwrap_or_else(|_| serde_json::Value::Object(serde_json::Map::new()));
|
||||
|
||||
pending_assistant_parts.push(AnthropicContentPart::ToolUse {
|
||||
id: call_id.clone(),
|
||||
name: name.clone(),
|
||||
input,
|
||||
});
|
||||
}
|
||||
|
||||
Item::FunctionCallOutput { call_id, output, .. } => {
|
||||
// Flush pending assistant parts first
|
||||
if !pending_assistant_parts.is_empty() {
|
||||
messages.push(AnthropicMessage {
|
||||
role: "assistant".to_string(),
|
||||
content: AnthropicContent::Parts(std::mem::take(
|
||||
&mut pending_assistant_parts,
|
||||
)),
|
||||
});
|
||||
}
|
||||
|
||||
pending_user_parts.push(AnthropicContentPart::ToolResult {
|
||||
tool_use_id: call_id.clone(),
|
||||
content: output.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
Item::Reasoning { text, .. } => {
|
||||
// Flush pending user parts first
|
||||
if !pending_user_parts.is_empty() {
|
||||
messages.push(AnthropicMessage {
|
||||
role: "user".to_string(),
|
||||
content: AnthropicContent::Parts(std::mem::take(
|
||||
&mut pending_user_parts,
|
||||
)),
|
||||
});
|
||||
}
|
||||
|
||||
// Reasoning is treated as assistant text in Anthropic
|
||||
// (actual thinking blocks are handled differently in streaming)
|
||||
pending_assistant_parts.push(AnthropicContentPart::Text { text: text.clone() });
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Flush remaining pending parts
|
||||
self.flush_pending_parts(
|
||||
&mut messages,
|
||||
&mut pending_assistant_parts,
|
||||
&mut pending_user_parts,
|
||||
);
|
||||
|
||||
messages
|
||||
}
|
||||
|
||||
fn flush_pending_parts(
|
||||
&self,
|
||||
messages: &mut Vec<AnthropicMessage>,
|
||||
pending_assistant_parts: &mut Vec<AnthropicContentPart>,
|
||||
pending_user_parts: &mut Vec<AnthropicContentPart>,
|
||||
) {
|
||||
if !pending_assistant_parts.is_empty() {
|
||||
messages.push(AnthropicMessage {
|
||||
role: "assistant".to_string(),
|
||||
content: AnthropicContent::Parts(std::mem::take(pending_assistant_parts)),
|
||||
});
|
||||
}
|
||||
if !pending_user_parts.is_empty() {
|
||||
messages.push(AnthropicMessage {
|
||||
role: "user".to_string(),
|
||||
content: AnthropicContent::Parts(std::mem::take(pending_user_parts)),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_tool(&self, tool: &ToolDefinition) -> AnthropicTool {
|
||||
AnthropicTool {
|
||||
name: tool.name.clone(),
|
||||
description: tool.description.clone(),
|
||||
input_schema: tool.input_schema.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_build_simple_request() {
|
||||
let scheme = AnthropicScheme::new();
|
||||
let request = Request::new()
|
||||
.system("You are a helpful assistant.")
|
||||
.user("Hello!");
|
||||
|
||||
let anthropic_req = scheme.build_request("claude-sonnet-4-20250514", &request);
|
||||
|
||||
assert_eq!(anthropic_req.model, "claude-sonnet-4-20250514");
|
||||
assert_eq!(
|
||||
anthropic_req.system,
|
||||
Some("You are a helpful assistant.".to_string())
|
||||
);
|
||||
assert_eq!(anthropic_req.messages.len(), 1);
|
||||
assert!(anthropic_req.stream);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_request_with_tool() {
|
||||
let scheme = AnthropicScheme::new();
|
||||
let request = Request::new().user("What's the weather?").tool(
|
||||
ToolDefinition::new("get_weather")
|
||||
.description("Get current weather")
|
||||
.input_schema(serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": { "type": "string" }
|
||||
},
|
||||
"required": ["location"]
|
||||
})),
|
||||
);
|
||||
|
||||
let anthropic_req = scheme.build_request("claude-sonnet-4-20250514", &request);
|
||||
|
||||
assert_eq!(anthropic_req.tools.len(), 1);
|
||||
assert_eq!(anthropic_req.tools[0].name, "get_weather");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_function_call_and_output() {
|
||||
let scheme = AnthropicScheme::new();
|
||||
let request = Request::new()
|
||||
.user("What's the weather?")
|
||||
.item(Item::function_call(
|
||||
"call_123",
|
||||
"get_weather",
|
||||
r#"{"city":"Tokyo"}"#,
|
||||
))
|
||||
.item(Item::function_call_output("call_123", "Sunny, 25°C"));
|
||||
|
||||
let anthropic_req = scheme.build_request("claude-sonnet-4-20250514", &request);
|
||||
|
||||
assert_eq!(anthropic_req.messages.len(), 3);
|
||||
assert_eq!(anthropic_req.messages[0].role, "user");
|
||||
assert_eq!(anthropic_req.messages[1].role, "assistant");
|
||||
assert_eq!(anthropic_req.messages[2].role, "user");
|
||||
}
|
||||
}
|
||||
+175
-87
@@ -1,130 +1,130 @@
|
||||
//! Gemini リクエスト生成
|
||||
//! Gemini Request Builder
|
||||
//!
|
||||
//! Google Gemini APIへのリクエストボディを構築
|
||||
//! Converts Open Responses native Item model to Google Gemini API format.
|
||||
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::llm_client::{
|
||||
types::{Item, Role, ToolDefinition},
|
||||
Request,
|
||||
types::{ContentPart, Message, MessageContent, Role, ToolDefinition},
|
||||
};
|
||||
|
||||
use super::GeminiScheme;
|
||||
|
||||
/// Gemini APIへのリクエストボディ
|
||||
/// Gemini API request body
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub(crate) struct GeminiRequest {
|
||||
/// コンテンツ(会話履歴)
|
||||
/// Contents (conversation history)
|
||||
pub contents: Vec<GeminiContent>,
|
||||
/// システム指示
|
||||
/// System instruction
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub system_instruction: Option<GeminiContent>,
|
||||
/// ツール定義
|
||||
/// Tool definitions
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub tools: Vec<GeminiTool>,
|
||||
/// ツール設定
|
||||
/// Tool config
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_config: Option<GeminiToolConfig>,
|
||||
/// 生成設定
|
||||
/// Generation config
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub generation_config: Option<GeminiGenerationConfig>,
|
||||
}
|
||||
|
||||
/// Gemini コンテンツ
|
||||
/// Gemini content
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct GeminiContent {
|
||||
/// ロール
|
||||
/// Role
|
||||
pub role: String,
|
||||
/// パーツ
|
||||
/// Parts
|
||||
pub parts: Vec<GeminiPart>,
|
||||
}
|
||||
|
||||
/// Gemini パーツ
|
||||
/// Gemini part
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum GeminiPart {
|
||||
/// テキストパーツ
|
||||
/// Text part
|
||||
Text { text: String },
|
||||
/// 関数呼び出しパーツ
|
||||
/// Function call part
|
||||
FunctionCall {
|
||||
#[serde(rename = "functionCall")]
|
||||
function_call: GeminiFunctionCall,
|
||||
},
|
||||
/// 関数レスポンスパーツ
|
||||
/// Function response part
|
||||
FunctionResponse {
|
||||
#[serde(rename = "functionResponse")]
|
||||
function_response: GeminiFunctionResponse,
|
||||
},
|
||||
}
|
||||
|
||||
/// Gemini 関数呼び出し
|
||||
/// Gemini function call
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct GeminiFunctionCall {
|
||||
pub name: String,
|
||||
pub args: Value,
|
||||
}
|
||||
|
||||
/// Gemini 関数レスポンス
|
||||
/// Gemini function response
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct GeminiFunctionResponse {
|
||||
pub name: String,
|
||||
pub response: GeminiFunctionResponseContent,
|
||||
}
|
||||
|
||||
/// Gemini 関数レスポンス内容
|
||||
/// Gemini function response content
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct GeminiFunctionResponseContent {
|
||||
pub name: String,
|
||||
pub content: Value,
|
||||
}
|
||||
|
||||
/// Gemini ツール定義
|
||||
/// Gemini tool definition
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub(crate) struct GeminiTool {
|
||||
/// 関数宣言
|
||||
/// Function declarations
|
||||
pub function_declarations: Vec<GeminiFunctionDeclaration>,
|
||||
}
|
||||
|
||||
/// Gemini 関数宣言
|
||||
/// Gemini function declaration
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct GeminiFunctionDeclaration {
|
||||
/// 関数名
|
||||
/// Function name
|
||||
pub name: String,
|
||||
/// 説明
|
||||
/// Description
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
/// パラメータスキーマ
|
||||
/// Parameter schema
|
||||
pub parameters: Value,
|
||||
}
|
||||
|
||||
/// Gemini ツール設定
|
||||
/// Gemini tool config
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub(crate) struct GeminiToolConfig {
|
||||
/// 関数呼び出し設定
|
||||
/// Function calling config
|
||||
pub function_calling_config: GeminiFunctionCallingConfig,
|
||||
}
|
||||
|
||||
/// Gemini 関数呼び出し設定
|
||||
/// Gemini function calling config
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub(crate) struct GeminiFunctionCallingConfig {
|
||||
/// モード: AUTO, ANY, NONE
|
||||
/// Mode: AUTO, ANY, NONE
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub mode: Option<String>,
|
||||
/// ストリーミング関数呼び出し引数を有効にするか
|
||||
/// Enable streaming function call arguments
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stream_function_call_arguments: Option<bool>,
|
||||
}
|
||||
|
||||
/// Gemini 生成設定
|
||||
/// Gemini generation config
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub(crate) struct GeminiGenerationConfig {
|
||||
/// 最大出力トークン数
|
||||
/// Max output tokens
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_output_tokens: Option<u32>,
|
||||
/// Temperature
|
||||
@@ -133,27 +133,26 @@ pub(crate) struct GeminiGenerationConfig {
|
||||
/// Top P
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub top_p: Option<f32>,
|
||||
/// ストップシーケンス
|
||||
/// Top K
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub top_k: Option<u32>,
|
||||
/// Stop sequences
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub stop_sequences: Vec<String>,
|
||||
}
|
||||
|
||||
impl GeminiScheme {
|
||||
/// RequestからGeminiのリクエストボディを構築
|
||||
/// Build Gemini request from Request
|
||||
pub(crate) fn build_request(&self, request: &Request) -> GeminiRequest {
|
||||
let mut contents = Vec::new();
|
||||
let contents = self.convert_items_to_contents(&request.items);
|
||||
|
||||
for message in &request.messages {
|
||||
contents.push(self.convert_message(message));
|
||||
}
|
||||
|
||||
// システムプロンプト
|
||||
// System prompt
|
||||
let system_instruction = request.system_prompt.as_ref().map(|s| GeminiContent {
|
||||
role: "user".to_string(), // system_instructionではroleは"user"か省略
|
||||
role: "user".to_string(),
|
||||
parts: vec![GeminiPart::Text { text: s.clone() }],
|
||||
});
|
||||
|
||||
// ツール
|
||||
// Tools
|
||||
let tools = if request.tools.is_empty() {
|
||||
vec![]
|
||||
} else {
|
||||
@@ -162,7 +161,7 @@ impl GeminiScheme {
|
||||
}]
|
||||
};
|
||||
|
||||
// ツール設定
|
||||
// Tool config
|
||||
let tool_config = if !request.tools.is_empty() {
|
||||
Some(GeminiToolConfig {
|
||||
function_calling_config: GeminiFunctionCallingConfig {
|
||||
@@ -178,11 +177,12 @@ impl GeminiScheme {
|
||||
None
|
||||
};
|
||||
|
||||
// 生成設定
|
||||
// Generation config
|
||||
let generation_config = Some(GeminiGenerationConfig {
|
||||
max_output_tokens: request.config.max_tokens,
|
||||
temperature: request.config.temperature,
|
||||
top_p: request.config.top_p,
|
||||
top_k: request.config.top_k,
|
||||
stop_sequences: request.config.stop_sequences.clone(),
|
||||
});
|
||||
|
||||
@@ -195,58 +195,126 @@ impl GeminiScheme {
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_message(&self, message: &Message) -> GeminiContent {
|
||||
let role = match message.role {
|
||||
/// Convert Open Responses Items to Gemini Contents
|
||||
///
|
||||
/// Gemini uses:
|
||||
/// - role "user" for user messages and function responses
|
||||
/// - role "model" for assistant messages and function calls
|
||||
fn convert_items_to_contents(&self, items: &[Item]) -> Vec<GeminiContent> {
|
||||
let mut contents = Vec::new();
|
||||
let mut pending_model_parts: Vec<GeminiPart> = Vec::new();
|
||||
let mut pending_user_parts: Vec<GeminiPart> = Vec::new();
|
||||
|
||||
for item in items {
|
||||
match item {
|
||||
Item::Message { role, content, .. } => {
|
||||
// Flush pending parts
|
||||
self.flush_pending_parts(
|
||||
&mut contents,
|
||||
&mut pending_model_parts,
|
||||
&mut pending_user_parts,
|
||||
);
|
||||
|
||||
let gemini_role = match role {
|
||||
Role::User => "user",
|
||||
Role::Assistant => "model",
|
||||
Role::System => continue, // Skip system role items
|
||||
};
|
||||
|
||||
let parts = match &message.content {
|
||||
MessageContent::Text(text) => vec![GeminiPart::Text { text: text.clone() }],
|
||||
MessageContent::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
} => {
|
||||
// Geminiでは関数レスポンスとしてマップ
|
||||
vec![GeminiPart::FunctionResponse {
|
||||
function_response: GeminiFunctionResponse {
|
||||
name: tool_use_id.clone(),
|
||||
response: GeminiFunctionResponseContent {
|
||||
name: tool_use_id.clone(),
|
||||
content: serde_json::Value::String(content.clone()),
|
||||
},
|
||||
},
|
||||
}]
|
||||
}
|
||||
MessageContent::Parts(parts) => parts
|
||||
let parts: Vec<GeminiPart> = content
|
||||
.iter()
|
||||
.map(|p| match p {
|
||||
ContentPart::Text { text } => GeminiPart::Text { text: text.clone() },
|
||||
ContentPart::ToolUse { id: _, name, input } => GeminiPart::FunctionCall {
|
||||
.map(|p| GeminiPart::Text {
|
||||
text: p.as_text().to_string(),
|
||||
})
|
||||
.collect();
|
||||
|
||||
contents.push(GeminiContent {
|
||||
role: gemini_role.to_string(),
|
||||
parts,
|
||||
});
|
||||
}
|
||||
|
||||
Item::FunctionCall {
|
||||
name, arguments, ..
|
||||
} => {
|
||||
// Flush pending user parts first
|
||||
if !pending_user_parts.is_empty() {
|
||||
contents.push(GeminiContent {
|
||||
role: "user".to_string(),
|
||||
parts: std::mem::take(&mut pending_user_parts),
|
||||
});
|
||||
}
|
||||
|
||||
// Parse arguments
|
||||
let args = serde_json::from_str(arguments)
|
||||
.unwrap_or_else(|_| Value::Object(serde_json::Map::new()));
|
||||
|
||||
pending_model_parts.push(GeminiPart::FunctionCall {
|
||||
function_call: GeminiFunctionCall {
|
||||
name: name.clone(),
|
||||
args: input.clone(),
|
||||
args,
|
||||
},
|
||||
},
|
||||
ContentPart::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
} => GeminiPart::FunctionResponse {
|
||||
function_response: GeminiFunctionResponse {
|
||||
name: tool_use_id.clone(),
|
||||
response: GeminiFunctionResponseContent {
|
||||
name: tool_use_id.clone(),
|
||||
content: serde_json::Value::String(content.clone()),
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
.collect(),
|
||||
};
|
||||
});
|
||||
}
|
||||
|
||||
GeminiContent {
|
||||
role: role.to_string(),
|
||||
parts,
|
||||
Item::FunctionCallOutput { call_id, output, .. } => {
|
||||
// Flush pending model parts first
|
||||
if !pending_model_parts.is_empty() {
|
||||
contents.push(GeminiContent {
|
||||
role: "model".to_string(),
|
||||
parts: std::mem::take(&mut pending_model_parts),
|
||||
});
|
||||
}
|
||||
|
||||
pending_user_parts.push(GeminiPart::FunctionResponse {
|
||||
function_response: GeminiFunctionResponse {
|
||||
name: call_id.clone(),
|
||||
response: GeminiFunctionResponseContent {
|
||||
name: call_id.clone(),
|
||||
content: Value::String(output.clone()),
|
||||
},
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
Item::Reasoning { text, .. } => {
|
||||
// Flush pending user parts first
|
||||
if !pending_user_parts.is_empty() {
|
||||
contents.push(GeminiContent {
|
||||
role: "user".to_string(),
|
||||
parts: std::mem::take(&mut pending_user_parts),
|
||||
});
|
||||
}
|
||||
|
||||
// Reasoning is treated as model text in Gemini
|
||||
pending_model_parts.push(GeminiPart::Text { text: text.clone() });
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Flush remaining pending parts
|
||||
self.flush_pending_parts(&mut contents, &mut pending_model_parts, &mut pending_user_parts);
|
||||
|
||||
contents
|
||||
}
|
||||
|
||||
fn flush_pending_parts(
|
||||
&self,
|
||||
contents: &mut Vec<GeminiContent>,
|
||||
pending_model_parts: &mut Vec<GeminiPart>,
|
||||
pending_user_parts: &mut Vec<GeminiPart>,
|
||||
) {
|
||||
if !pending_model_parts.is_empty() {
|
||||
contents.push(GeminiContent {
|
||||
role: "model".to_string(),
|
||||
parts: std::mem::take(pending_model_parts),
|
||||
});
|
||||
}
|
||||
if !pending_user_parts.is_empty() {
|
||||
contents.push(GeminiContent {
|
||||
role: "user".to_string(),
|
||||
parts: std::mem::take(pending_user_parts),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -314,4 +382,24 @@ mod tests {
|
||||
assert_eq!(gemini_req.contents[0].role, "user");
|
||||
assert_eq!(gemini_req.contents[1].role, "model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_function_call_and_output() {
|
||||
let scheme = GeminiScheme::new();
|
||||
let request = Request::new()
|
||||
.user("What's the weather?")
|
||||
.item(Item::function_call(
|
||||
"call_123",
|
||||
"get_weather",
|
||||
r#"{"city":"Tokyo"}"#,
|
||||
))
|
||||
.item(Item::function_call_output("call_123", "Sunny, 25°C"));
|
||||
|
||||
let gemini_req = scheme.build_request(&request);
|
||||
|
||||
assert_eq!(gemini_req.contents.len(), 3);
|
||||
assert_eq!(gemini_req.contents[0].role, "user");
|
||||
assert_eq!(gemini_req.contents[1].role, "model");
|
||||
assert_eq!(gemini_req.contents[2].role, "user");
|
||||
}
|
||||
}
|
||||
+133
-94
@@ -1,21 +1,23 @@
|
||||
//! OpenAI リクエスト生成
|
||||
//! OpenAI Request Builder
|
||||
//!
|
||||
//! Converts Open Responses native Item model to OpenAI Chat Completions API format.
|
||||
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::llm_client::{
|
||||
types::{Item, Role, ToolDefinition},
|
||||
Request,
|
||||
types::{ContentPart, Message, MessageContent, Role, ToolDefinition},
|
||||
};
|
||||
|
||||
use super::OpenAIScheme;
|
||||
|
||||
/// OpenAI APIへのリクエストボディ
|
||||
/// OpenAI API request body
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct OpenAIRequest {
|
||||
pub model: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_completion_tokens: Option<u32>, // max_tokens is deprecated for newer models, generally max_completion_tokens is preferred
|
||||
pub max_completion_tokens: Option<u32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_tokens: Option<u32>, // Legacy field for compatibility (e.g. Ollama)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
@@ -31,7 +33,7 @@ pub(crate) struct OpenAIRequest {
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub tools: Vec<OpenAITool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_choice: Option<String>, // "auto", "none", or specific
|
||||
pub tool_choice: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
@@ -39,20 +41,21 @@ pub(crate) struct StreamOptions {
|
||||
pub include_usage: bool,
|
||||
}
|
||||
|
||||
/// OpenAI メッセージ
|
||||
/// OpenAI message
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct OpenAIMessage {
|
||||
pub role: String,
|
||||
pub content: Option<OpenAIContent>, // Optional for assistant tool calls
|
||||
pub content: Option<OpenAIContent>,
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub tool_calls: Vec<OpenAIToolCall>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_call_id: Option<String>, // For tool_result (role: tool)
|
||||
pub tool_call_id: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub name: Option<String>, // Optional name
|
||||
pub name: Option<String>,
|
||||
}
|
||||
|
||||
/// OpenAI コンテンツ
|
||||
/// OpenAI content
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum OpenAIContent {
|
||||
@@ -60,7 +63,7 @@ pub(crate) enum OpenAIContent {
|
||||
Parts(Vec<OpenAIContentPart>),
|
||||
}
|
||||
|
||||
/// OpenAI コンテンツパーツ
|
||||
/// OpenAI content part
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(tag = "type")]
|
||||
@@ -76,7 +79,7 @@ pub(crate) struct ImageUrl {
|
||||
pub url: String,
|
||||
}
|
||||
|
||||
/// OpenAI ツール定義
|
||||
/// OpenAI tool definition
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct OpenAITool {
|
||||
pub r#type: String,
|
||||
@@ -91,7 +94,7 @@ pub(crate) struct OpenAIToolFunction {
|
||||
pub parameters: Value,
|
||||
}
|
||||
|
||||
/// OpenAI ツール呼び出し(メッセージ内)
|
||||
/// OpenAI tool call in message
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct OpenAIToolCall {
|
||||
pub id: String,
|
||||
@@ -106,10 +109,11 @@ pub(crate) struct OpenAIToolCallFunction {
|
||||
}
|
||||
|
||||
impl OpenAIScheme {
|
||||
/// RequestからOpenAIのリクエストボディを構築
|
||||
/// Build OpenAI request from Request
|
||||
pub(crate) fn build_request(&self, model: &str, request: &Request) -> OpenAIRequest {
|
||||
let mut messages = Vec::new();
|
||||
|
||||
// Add system message if present
|
||||
if let Some(system) = &request.system_prompt {
|
||||
messages.push(OpenAIMessage {
|
||||
role: "system".to_string(),
|
||||
@@ -120,7 +124,8 @@ impl OpenAIScheme {
|
||||
});
|
||||
}
|
||||
|
||||
messages.extend(request.messages.iter().map(|m| self.convert_message(m)));
|
||||
// Convert items to messages
|
||||
messages.extend(self.convert_items_to_messages(&request.items));
|
||||
|
||||
let tools = request.tools.iter().map(|t| self.convert_tool(t)).collect();
|
||||
|
||||
@@ -143,106 +148,122 @@ impl OpenAIScheme {
|
||||
}),
|
||||
messages,
|
||||
tools,
|
||||
tool_choice: None, // Default to auto if tools are present? Or let API decide (which is auto)
|
||||
tool_choice: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_message(&self, message: &Message) -> OpenAIMessage {
|
||||
match &message.content {
|
||||
MessageContent::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
} => OpenAIMessage {
|
||||
role: "tool".to_string(),
|
||||
content: Some(OpenAIContent::Text(content.clone())),
|
||||
tool_calls: vec![],
|
||||
tool_call_id: Some(tool_use_id.clone()),
|
||||
name: None,
|
||||
},
|
||||
MessageContent::Text(text) => {
|
||||
let role = match message.role {
|
||||
/// Convert Open Responses Items to OpenAI Messages
|
||||
///
|
||||
/// OpenAI uses a message-based model where:
|
||||
/// - User messages have role "user"
|
||||
/// - Assistant messages have role "assistant"
|
||||
/// - Tool calls are within assistant messages as tool_calls array
|
||||
/// - Tool results have role "tool" with tool_call_id
|
||||
fn convert_items_to_messages(&self, items: &[Item]) -> Vec<OpenAIMessage> {
|
||||
let mut messages = Vec::new();
|
||||
let mut pending_tool_calls: Vec<OpenAIToolCall> = Vec::new();
|
||||
let mut pending_assistant_text: Option<String> = None;
|
||||
|
||||
for item in items {
|
||||
match item {
|
||||
Item::Message { role, content, .. } => {
|
||||
// Flush pending tool calls
|
||||
self.flush_pending_assistant(
|
||||
&mut messages,
|
||||
&mut pending_tool_calls,
|
||||
&mut pending_assistant_text,
|
||||
);
|
||||
|
||||
let openai_role = match role {
|
||||
Role::User => "user",
|
||||
Role::Assistant => "assistant",
|
||||
Role::System => "system",
|
||||
};
|
||||
OpenAIMessage {
|
||||
role: role.to_string(),
|
||||
content: Some(OpenAIContent::Text(text.clone())),
|
||||
|
||||
let text_content: String = content
|
||||
.iter()
|
||||
.map(|p| p.as_text())
|
||||
.collect::<Vec<_>>()
|
||||
.join("");
|
||||
|
||||
messages.push(OpenAIMessage {
|
||||
role: openai_role.to_string(),
|
||||
content: Some(OpenAIContent::Text(text_content)),
|
||||
tool_calls: vec![],
|
||||
tool_call_id: None,
|
||||
name: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
MessageContent::Parts(parts) => {
|
||||
let role = match message.role {
|
||||
Role::User => "user",
|
||||
Role::Assistant => "assistant",
|
||||
};
|
||||
|
||||
let mut content_parts = Vec::new();
|
||||
let mut tool_calls = Vec::new();
|
||||
let mut is_tool_result = false;
|
||||
let mut tool_result_id = None;
|
||||
let mut tool_result_content = String::new();
|
||||
|
||||
for part in parts {
|
||||
match part {
|
||||
ContentPart::Text { text } => {
|
||||
content_parts.push(OpenAIContentPart::Text { text: text.clone() });
|
||||
}
|
||||
ContentPart::ToolUse { id, name, input } => {
|
||||
tool_calls.push(OpenAIToolCall {
|
||||
id: id.clone(),
|
||||
Item::FunctionCall {
|
||||
call_id,
|
||||
name,
|
||||
arguments,
|
||||
..
|
||||
} => {
|
||||
pending_tool_calls.push(OpenAIToolCall {
|
||||
id: call_id.clone(),
|
||||
r#type: "function".to_string(),
|
||||
function: OpenAIToolCallFunction {
|
||||
name: name.clone(),
|
||||
arguments: input.to_string(),
|
||||
arguments: arguments.clone(),
|
||||
},
|
||||
});
|
||||
}
|
||||
ContentPart::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
} => {
|
||||
// OpenAI doesn't support mixed content with ToolResult in the same message easily if not careful
|
||||
// But strictly speaking, a Message with ToolResult should be its own message with role "tool"
|
||||
is_tool_result = true;
|
||||
tool_result_id = Some(tool_use_id.clone());
|
||||
tool_result_content = content.clone();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if is_tool_result {
|
||||
OpenAIMessage {
|
||||
Item::FunctionCallOutput { call_id, output, .. } => {
|
||||
// Flush pending tool calls before tool result
|
||||
self.flush_pending_assistant(
|
||||
&mut messages,
|
||||
&mut pending_tool_calls,
|
||||
&mut pending_assistant_text,
|
||||
);
|
||||
|
||||
messages.push(OpenAIMessage {
|
||||
role: "tool".to_string(),
|
||||
content: Some(OpenAIContent::Text(tool_result_content)),
|
||||
content: Some(OpenAIContent::Text(output.clone())),
|
||||
tool_calls: vec![],
|
||||
tool_call_id: tool_result_id,
|
||||
tool_call_id: Some(call_id.clone()),
|
||||
name: None,
|
||||
});
|
||||
}
|
||||
} else {
|
||||
let content = if content_parts.is_empty() {
|
||||
None
|
||||
} else if content_parts.len() == 1 {
|
||||
// Simplify single text part to just Text content if preferred, or keep as Parts
|
||||
if let OpenAIContentPart::Text { text } = &content_parts[0] {
|
||||
Some(OpenAIContent::Text(text.clone()))
|
||||
} else {
|
||||
Some(OpenAIContent::Parts(content_parts))
|
||||
}
|
||||
} else {
|
||||
Some(OpenAIContent::Parts(content_parts))
|
||||
};
|
||||
|
||||
OpenAIMessage {
|
||||
role: role.to_string(),
|
||||
content,
|
||||
tool_calls,
|
||||
Item::Reasoning { text, .. } => {
|
||||
// Reasoning is treated as assistant text in OpenAI
|
||||
// (OpenAI doesn't have native reasoning support like Claude)
|
||||
if let Some(ref mut existing) = pending_assistant_text {
|
||||
existing.push_str(text);
|
||||
} else {
|
||||
pending_assistant_text = Some(text.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Flush remaining pending items
|
||||
self.flush_pending_assistant(
|
||||
&mut messages,
|
||||
&mut pending_tool_calls,
|
||||
&mut pending_assistant_text,
|
||||
);
|
||||
|
||||
messages
|
||||
}
|
||||
|
||||
fn flush_pending_assistant(
|
||||
&self,
|
||||
messages: &mut Vec<OpenAIMessage>,
|
||||
pending_tool_calls: &mut Vec<OpenAIToolCall>,
|
||||
pending_assistant_text: &mut Option<String>,
|
||||
) {
|
||||
if !pending_tool_calls.is_empty() || pending_assistant_text.is_some() {
|
||||
messages.push(OpenAIMessage {
|
||||
role: "assistant".to_string(),
|
||||
content: pending_assistant_text.take().map(OpenAIContent::Text),
|
||||
tool_calls: std::mem::take(pending_tool_calls),
|
||||
tool_call_id: None,
|
||||
name: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -274,7 +295,6 @@ mod tests {
|
||||
assert_eq!(body.messages[0].role, "system");
|
||||
assert_eq!(body.messages[1].role, "user");
|
||||
|
||||
// Check system content
|
||||
if let Some(OpenAIContent::Text(text)) = &body.messages[0].content {
|
||||
assert_eq!(text, "System prompt");
|
||||
} else {
|
||||
@@ -301,20 +321,39 @@ mod tests {
|
||||
|
||||
let body = scheme.build_request("llama3", &request);
|
||||
|
||||
// max_tokens should be set, max_completion_tokens should be None
|
||||
assert_eq!(body.max_tokens, Some(100));
|
||||
assert!(body.max_completion_tokens.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_request_modern_max_tokens() {
|
||||
let scheme = OpenAIScheme::new(); // Default matches modern (legacy=false)
|
||||
let scheme = OpenAIScheme::new();
|
||||
let request = Request::new().user("Hello").max_tokens(100);
|
||||
|
||||
let body = scheme.build_request("gpt-4o", &request);
|
||||
|
||||
// max_completion_tokens should be set, max_tokens should be None
|
||||
assert_eq!(body.max_completion_tokens, Some(100));
|
||||
assert!(body.max_tokens.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_function_call_and_output() {
|
||||
let scheme = OpenAIScheme::new();
|
||||
let request = Request::new()
|
||||
.user("Check weather")
|
||||
.item(Item::function_call(
|
||||
"call_123",
|
||||
"get_weather",
|
||||
r#"{"city":"Tokyo"}"#,
|
||||
))
|
||||
.item(Item::function_call_output("call_123", "Sunny, 25°C"));
|
||||
|
||||
let body = scheme.build_request("gpt-4o", &request);
|
||||
|
||||
assert_eq!(body.messages.len(), 3);
|
||||
assert_eq!(body.messages[0].role, "user");
|
||||
assert_eq!(body.messages[1].role, "assistant");
|
||||
assert_eq!(body.messages[1].tool_calls.len(), 1);
|
||||
assert_eq!(body.messages[2].role, "tool");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,494 @@
|
||||
//! Open Responses Event Parser
|
||||
//!
|
||||
//! Parses SSE events from the Open Responses API into internal Event types.
|
||||
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::llm_client::{
|
||||
event::{
|
||||
BlockMetadata, BlockStart, BlockStop, DeltaContent, ErrorEvent, Event, ResponseStatus,
|
||||
StatusEvent, StopReason, UsageEvent,
|
||||
},
|
||||
ClientError,
|
||||
};
|
||||
|
||||
// =============================================================================
|
||||
// Open Responses SSE Event Types
|
||||
// =============================================================================
|
||||
|
||||
/// Response created event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ResponseCreatedEvent {
|
||||
pub response: ResponseObject,
|
||||
}
|
||||
|
||||
/// Response object
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ResponseObject {
|
||||
pub id: String,
|
||||
pub status: String,
|
||||
#[serde(default)]
|
||||
pub output: Vec<OutputItem>,
|
||||
pub usage: Option<UsageObject>,
|
||||
}
|
||||
|
||||
/// Output item in response
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum OutputItem {
|
||||
Message {
|
||||
id: String,
|
||||
role: String,
|
||||
#[serde(default)]
|
||||
content: Vec<ContentPartObject>,
|
||||
},
|
||||
FunctionCall {
|
||||
id: String,
|
||||
call_id: String,
|
||||
name: String,
|
||||
arguments: String,
|
||||
},
|
||||
Reasoning {
|
||||
id: String,
|
||||
#[serde(default)]
|
||||
text: String,
|
||||
},
|
||||
}
|
||||
|
||||
/// Content part object
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum ContentPartObject {
|
||||
OutputText { text: String },
|
||||
InputText { text: String },
|
||||
Refusal { refusal: String },
|
||||
}
|
||||
|
||||
/// Usage object
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct UsageObject {
|
||||
pub input_tokens: Option<u64>,
|
||||
pub output_tokens: Option<u64>,
|
||||
pub total_tokens: Option<u64>,
|
||||
}
|
||||
|
||||
/// Output item added event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OutputItemAddedEvent {
|
||||
pub output_index: usize,
|
||||
pub item: OutputItem,
|
||||
}
|
||||
|
||||
/// Text delta event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct TextDeltaEvent {
|
||||
pub output_index: usize,
|
||||
pub content_index: usize,
|
||||
pub delta: String,
|
||||
}
|
||||
|
||||
/// Text done event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct TextDoneEvent {
|
||||
pub output_index: usize,
|
||||
pub content_index: usize,
|
||||
pub text: String,
|
||||
}
|
||||
|
||||
/// Function call arguments delta event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct FunctionCallArgumentsDeltaEvent {
|
||||
pub output_index: usize,
|
||||
pub call_id: String,
|
||||
pub delta: String,
|
||||
}
|
||||
|
||||
/// Function call arguments done event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct FunctionCallArgumentsDoneEvent {
|
||||
pub output_index: usize,
|
||||
pub call_id: String,
|
||||
pub arguments: String,
|
||||
}
|
||||
|
||||
/// Reasoning delta event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ReasoningDeltaEvent {
|
||||
pub output_index: usize,
|
||||
pub delta: String,
|
||||
}
|
||||
|
||||
/// Reasoning done event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ReasoningDoneEvent {
|
||||
pub output_index: usize,
|
||||
pub text: String,
|
||||
}
|
||||
|
||||
/// Content part done event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ContentPartDoneEvent {
|
||||
pub output_index: usize,
|
||||
pub content_index: usize,
|
||||
pub part: ContentPartObject,
|
||||
}
|
||||
|
||||
/// Output item done event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OutputItemDoneEvent {
|
||||
pub output_index: usize,
|
||||
pub item: OutputItem,
|
||||
}
|
||||
|
||||
/// Response done event
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ResponseDoneEvent {
|
||||
pub response: ResponseObject,
|
||||
}
|
||||
|
||||
/// Error event from API
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ApiErrorEvent {
|
||||
pub error: ApiError,
|
||||
}
|
||||
|
||||
/// API error details
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct ApiError {
|
||||
pub code: Option<String>,
|
||||
pub message: String,
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Event Parsing
|
||||
// =============================================================================
|
||||
|
||||
/// Parse SSE event into internal Event(s)
|
||||
///
|
||||
/// Returns `Ok(None)` for events that should be ignored (e.g., heartbeats)
|
||||
/// Returns `Ok(Some(vec))` for events that produce one or more internal Events
|
||||
pub fn parse_event(event_type: &str, data: &str) -> Result<Option<Vec<Event>>, ClientError> {
|
||||
// Skip empty data
|
||||
if data.is_empty() || data == "[DONE]" {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let events = match event_type {
|
||||
// Response lifecycle
|
||||
"response.created" => {
|
||||
let _event: ResponseCreatedEvent = parse_json(data)?;
|
||||
Some(vec![Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Started,
|
||||
})])
|
||||
}
|
||||
|
||||
"response.in_progress" => {
|
||||
// Just a status update, no action needed
|
||||
None
|
||||
}
|
||||
|
||||
"response.completed" | "response.done" => {
|
||||
let event: ResponseDoneEvent = parse_json(data)?;
|
||||
let mut events = Vec::new();
|
||||
|
||||
// Emit usage if present
|
||||
if let Some(usage) = event.response.usage {
|
||||
events.push(Event::Usage(UsageEvent {
|
||||
input_tokens: usage.input_tokens,
|
||||
output_tokens: usage.output_tokens,
|
||||
total_tokens: usage.total_tokens,
|
||||
cache_read_input_tokens: None,
|
||||
cache_creation_input_tokens: None,
|
||||
}));
|
||||
}
|
||||
|
||||
events.push(Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}));
|
||||
Some(events)
|
||||
}
|
||||
|
||||
"response.failed" => {
|
||||
// Try to parse error
|
||||
if let Ok(error_event) = parse_json::<ApiErrorEvent>(data) {
|
||||
Some(vec![
|
||||
Event::Error(ErrorEvent {
|
||||
code: error_event.error.code,
|
||||
message: error_event.error.message,
|
||||
}),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Failed,
|
||||
}),
|
||||
])
|
||||
} else {
|
||||
Some(vec![Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Failed,
|
||||
})])
|
||||
}
|
||||
}
|
||||
|
||||
// Output item events
|
||||
"response.output_item.added" => {
|
||||
let event: OutputItemAddedEvent = parse_json(data)?;
|
||||
Some(vec![convert_item_added(&event)])
|
||||
}
|
||||
|
||||
"response.output_item.done" => {
|
||||
let event: OutputItemDoneEvent = parse_json(data)?;
|
||||
Some(vec![convert_item_done(&event)])
|
||||
}
|
||||
|
||||
// Text content events
|
||||
"response.output_text.delta" => {
|
||||
let event: TextDeltaEvent = parse_json(data)?;
|
||||
Some(vec![Event::text_delta(event.output_index, &event.delta)])
|
||||
}
|
||||
|
||||
"response.output_text.done" => {
|
||||
// Text done - we'll handle stop in output_item.done
|
||||
let _event: TextDoneEvent = parse_json(data)?;
|
||||
None
|
||||
}
|
||||
|
||||
// Content part events
|
||||
"response.content_part.added" => {
|
||||
// Content part added - we handle this via output_item.added
|
||||
None
|
||||
}
|
||||
|
||||
"response.content_part.done" => {
|
||||
// Content part done - we handle stop in output_item.done
|
||||
None
|
||||
}
|
||||
|
||||
// Function call events
|
||||
"response.function_call_arguments.delta" => {
|
||||
let event: FunctionCallArgumentsDeltaEvent = parse_json(data)?;
|
||||
Some(vec![Event::BlockDelta(crate::llm_client::event::BlockDelta {
|
||||
index: event.output_index,
|
||||
delta: DeltaContent::InputJson(event.delta),
|
||||
})])
|
||||
}
|
||||
|
||||
"response.function_call_arguments.done" => {
|
||||
// Arguments done - we handle stop in output_item.done
|
||||
let _event: FunctionCallArgumentsDoneEvent = parse_json(data)?;
|
||||
None
|
||||
}
|
||||
|
||||
// Reasoning events
|
||||
"response.reasoning.delta" | "response.reasoning_summary_text.delta" => {
|
||||
let event: ReasoningDeltaEvent = parse_json(data)?;
|
||||
Some(vec![Event::BlockDelta(crate::llm_client::event::BlockDelta {
|
||||
index: event.output_index,
|
||||
delta: DeltaContent::Thinking(event.delta),
|
||||
})])
|
||||
}
|
||||
|
||||
"response.reasoning.done" | "response.reasoning_summary_text.done" => {
|
||||
// Reasoning done - we handle stop in output_item.done
|
||||
let _event: ReasoningDoneEvent = parse_json(data)?;
|
||||
None
|
||||
}
|
||||
|
||||
// Error event
|
||||
"error" => {
|
||||
let event: ApiErrorEvent = parse_json(data)?;
|
||||
Some(vec![Event::Error(ErrorEvent {
|
||||
code: event.error.code,
|
||||
message: event.error.message,
|
||||
})])
|
||||
}
|
||||
|
||||
// Unknown event type - ignore
|
||||
_ => {
|
||||
tracing::debug!(event_type = event_type, "Unknown Open Responses event type");
|
||||
None
|
||||
}
|
||||
};
|
||||
|
||||
Ok(events)
|
||||
}
|
||||
|
||||
fn parse_json<T: serde::de::DeserializeOwned>(data: &str) -> Result<T, ClientError> {
|
||||
serde_json::from_str(data).map_err(|e| ClientError::Parse(e.to_string()))
|
||||
}
|
||||
|
||||
fn convert_item_added(event: &OutputItemAddedEvent) -> Event {
|
||||
match &event.item {
|
||||
OutputItem::Message { id, role: _, content: _ } => Event::BlockStart(BlockStart {
|
||||
index: event.output_index,
|
||||
block_type: crate::llm_client::event::BlockType::Text,
|
||||
metadata: BlockMetadata::Text,
|
||||
}),
|
||||
|
||||
OutputItem::FunctionCall {
|
||||
id,
|
||||
call_id,
|
||||
name,
|
||||
arguments: _,
|
||||
} => Event::BlockStart(BlockStart {
|
||||
index: event.output_index,
|
||||
block_type: crate::llm_client::event::BlockType::ToolUse,
|
||||
metadata: BlockMetadata::ToolUse {
|
||||
id: call_id.clone(),
|
||||
name: name.clone(),
|
||||
},
|
||||
}),
|
||||
|
||||
OutputItem::Reasoning { id, text: _ } => Event::BlockStart(BlockStart {
|
||||
index: event.output_index,
|
||||
block_type: crate::llm_client::event::BlockType::Thinking,
|
||||
metadata: BlockMetadata::Thinking,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_item_done(event: &OutputItemDoneEvent) -> Event {
|
||||
let stop_reason = match &event.item {
|
||||
OutputItem::FunctionCall { .. } => Some(StopReason::ToolUse),
|
||||
_ => Some(StopReason::EndTurn),
|
||||
};
|
||||
|
||||
Event::BlockStop(BlockStop {
|
||||
index: event.output_index,
|
||||
stop_reason,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_parse_response_created() {
|
||||
let data = r#"{"response":{"id":"resp_123","status":"in_progress","output":[]}}"#;
|
||||
let events = parse_event("response.created", data).unwrap().unwrap();
|
||||
assert_eq!(events.len(), 1);
|
||||
assert!(matches!(
|
||||
events[0],
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Started
|
||||
})
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_text_delta() {
|
||||
let data = r#"{"output_index":0,"content_index":0,"delta":"Hello"}"#;
|
||||
let events = parse_event("response.output_text.delta", data)
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(events.len(), 1);
|
||||
if let Event::BlockDelta(delta) = &events[0] {
|
||||
assert_eq!(delta.index, 0);
|
||||
assert!(matches!(&delta.delta, DeltaContent::Text(t) if t == "Hello"));
|
||||
} else {
|
||||
panic!("Expected BlockDelta");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_output_item_added_message() {
|
||||
let data = r#"{"output_index":0,"item":{"type":"message","id":"msg_123","role":"assistant","content":[]}}"#;
|
||||
let events = parse_event("response.output_item.added", data)
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(events.len(), 1);
|
||||
if let Event::BlockStart(start) = &events[0] {
|
||||
assert_eq!(start.index, 0);
|
||||
assert!(matches!(
|
||||
start.block_type,
|
||||
crate::llm_client::event::BlockType::Text
|
||||
));
|
||||
} else {
|
||||
panic!("Expected BlockStart");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_output_item_added_function_call() {
|
||||
let data = r#"{"output_index":1,"item":{"type":"function_call","id":"fc_123","call_id":"call_456","name":"get_weather","arguments":""}}"#;
|
||||
let events = parse_event("response.output_item.added", data)
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(events.len(), 1);
|
||||
if let Event::BlockStart(start) = &events[0] {
|
||||
assert_eq!(start.index, 1);
|
||||
assert!(matches!(
|
||||
start.block_type,
|
||||
crate::llm_client::event::BlockType::ToolUse
|
||||
));
|
||||
if let BlockMetadata::ToolUse { id, name } = &start.metadata {
|
||||
assert_eq!(id, "call_456");
|
||||
assert_eq!(name, "get_weather");
|
||||
} else {
|
||||
panic!("Expected ToolUse metadata");
|
||||
}
|
||||
} else {
|
||||
panic!("Expected BlockStart");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_function_call_arguments_delta() {
|
||||
let data = r#"{"output_index":1,"call_id":"call_456","delta":"{\"city\":"}"#;
|
||||
let events = parse_event("response.function_call_arguments.delta", data)
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(events.len(), 1);
|
||||
if let Event::BlockDelta(delta) = &events[0] {
|
||||
assert_eq!(delta.index, 1);
|
||||
assert!(matches!(
|
||||
&delta.delta,
|
||||
DeltaContent::InputJson(s) if s == "{\"city\":"
|
||||
));
|
||||
} else {
|
||||
panic!("Expected BlockDelta");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_response_completed() {
|
||||
let data = r#"{"response":{"id":"resp_123","status":"completed","output":[],"usage":{"input_tokens":10,"output_tokens":20,"total_tokens":30}}}"#;
|
||||
let events = parse_event("response.completed", data).unwrap().unwrap();
|
||||
assert_eq!(events.len(), 2);
|
||||
|
||||
// First event should be usage
|
||||
if let Event::Usage(usage) = &events[0] {
|
||||
assert_eq!(usage.input_tokens, Some(10));
|
||||
assert_eq!(usage.output_tokens, Some(20));
|
||||
assert_eq!(usage.total_tokens, Some(30));
|
||||
} else {
|
||||
panic!("Expected Usage event");
|
||||
}
|
||||
|
||||
// Second event should be status
|
||||
assert!(matches!(
|
||||
events[1],
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed
|
||||
})
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_error() {
|
||||
let data = r#"{"error":{"code":"rate_limit","message":"Too many requests"}}"#;
|
||||
let events = parse_event("error", data).unwrap().unwrap();
|
||||
assert_eq!(events.len(), 1);
|
||||
if let Event::Error(err) = &events[0] {
|
||||
assert_eq!(err.code, Some("rate_limit".to_string()));
|
||||
assert_eq!(err.message, "Too many requests");
|
||||
} else {
|
||||
panic!("Expected Error event");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_unknown_event() {
|
||||
let data = r#"{}"#;
|
||||
let events = parse_event("some.unknown.event", data).unwrap();
|
||||
assert!(events.is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
//! Open Responses Scheme
|
||||
//!
|
||||
//! Handles request/response conversion for the Open Responses API.
|
||||
//! Since our internal types are already Open Responses native, this scheme
|
||||
//! primarily passes through data with minimal transformation.
|
||||
|
||||
mod events;
|
||||
mod request;
|
||||
|
||||
use crate::llm_client::{ClientError, Request};
|
||||
|
||||
pub use events::*;
|
||||
pub use request::*;
|
||||
|
||||
/// Open Responses Scheme
|
||||
///
|
||||
/// Handles conversion between internal types and the Open Responses wire format.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct OpenResponsesScheme {
|
||||
/// Optional model override
|
||||
pub model: Option<String>,
|
||||
}
|
||||
|
||||
impl OpenResponsesScheme {
|
||||
/// Create a new OpenResponsesScheme
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
/// Set the model
|
||||
pub fn with_model(mut self, model: impl Into<String>) -> Self {
|
||||
self.model = Some(model.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// Build Open Responses request from internal Request
|
||||
pub fn build_request(&self, model: &str, request: &Request) -> OpenResponsesRequest {
|
||||
build_request(model, request)
|
||||
}
|
||||
|
||||
/// Parse SSE event data into internal Event(s)
|
||||
pub fn parse_event(
|
||||
&self,
|
||||
event_type: &str,
|
||||
data: &str,
|
||||
) -> Result<Option<Vec<crate::llm_client::Event>>, ClientError> {
|
||||
parse_event(event_type, data)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,285 @@
|
||||
//! Open Responses Request Builder
|
||||
//!
|
||||
//! Converts internal Request/Item types to Open Responses API format.
|
||||
//! Since our internal types are already Open Responses native, this is
|
||||
//! mostly a direct serialization with some field renaming.
|
||||
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::llm_client::{types::Item, Request, ToolDefinition};
|
||||
|
||||
/// Open Responses API request body
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct OpenResponsesRequest {
|
||||
/// Model identifier
|
||||
pub model: String,
|
||||
|
||||
/// Input items (conversation history)
|
||||
pub input: Vec<OpenResponsesItem>,
|
||||
|
||||
/// System instructions
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub instructions: Option<String>,
|
||||
|
||||
/// Tool definitions
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub tools: Vec<OpenResponsesTool>,
|
||||
|
||||
/// Enable streaming
|
||||
pub stream: bool,
|
||||
|
||||
/// Maximum output tokens
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_output_tokens: Option<u32>,
|
||||
|
||||
/// Temperature
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f32>,
|
||||
|
||||
/// Top P (nucleus sampling)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub top_p: Option<f32>,
|
||||
}
|
||||
|
||||
/// Open Responses input item
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum OpenResponsesItem {
|
||||
/// Message item
|
||||
Message {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
id: Option<String>,
|
||||
role: String,
|
||||
content: Vec<OpenResponsesContentPart>,
|
||||
},
|
||||
|
||||
/// Function call item
|
||||
FunctionCall {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
id: Option<String>,
|
||||
call_id: String,
|
||||
name: String,
|
||||
arguments: String,
|
||||
},
|
||||
|
||||
/// Function call output item
|
||||
FunctionCallOutput {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
id: Option<String>,
|
||||
call_id: String,
|
||||
output: String,
|
||||
},
|
||||
|
||||
/// Reasoning item
|
||||
Reasoning {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
id: Option<String>,
|
||||
text: String,
|
||||
},
|
||||
}
|
||||
|
||||
/// Open Responses content part
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum OpenResponsesContentPart {
|
||||
/// Input text (for user messages)
|
||||
InputText { text: String },
|
||||
|
||||
/// Output text (for assistant messages)
|
||||
OutputText { text: String },
|
||||
|
||||
/// Refusal
|
||||
Refusal { refusal: String },
|
||||
}
|
||||
|
||||
/// Open Responses tool definition
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct OpenResponsesTool {
|
||||
/// Tool type (always "function")
|
||||
pub r#type: String,
|
||||
|
||||
/// Function definition
|
||||
pub name: String,
|
||||
|
||||
/// Description
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
|
||||
/// Parameters schema
|
||||
pub parameters: Value,
|
||||
}
|
||||
|
||||
/// Build Open Responses request from internal Request
|
||||
pub fn build_request(model: &str, request: &Request) -> OpenResponsesRequest {
|
||||
let input = request.items.iter().map(convert_item).collect();
|
||||
let tools = request.tools.iter().map(convert_tool).collect();
|
||||
|
||||
OpenResponsesRequest {
|
||||
model: model.to_string(),
|
||||
input,
|
||||
instructions: request.system_prompt.clone(),
|
||||
tools,
|
||||
stream: true,
|
||||
max_output_tokens: request.config.max_tokens,
|
||||
temperature: request.config.temperature,
|
||||
top_p: request.config.top_p,
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_item(item: &Item) -> OpenResponsesItem {
|
||||
match item {
|
||||
Item::Message {
|
||||
id,
|
||||
role,
|
||||
content,
|
||||
status: _,
|
||||
} => {
|
||||
let role_str = match role {
|
||||
crate::llm_client::types::Role::User => "user",
|
||||
crate::llm_client::types::Role::Assistant => "assistant",
|
||||
crate::llm_client::types::Role::System => "system",
|
||||
};
|
||||
|
||||
let parts = content
|
||||
.iter()
|
||||
.map(|p| match p {
|
||||
crate::llm_client::types::ContentPart::InputText { text } => {
|
||||
OpenResponsesContentPart::InputText { text: text.clone() }
|
||||
}
|
||||
crate::llm_client::types::ContentPart::OutputText { text } => {
|
||||
OpenResponsesContentPart::OutputText { text: text.clone() }
|
||||
}
|
||||
crate::llm_client::types::ContentPart::Refusal { refusal } => {
|
||||
OpenResponsesContentPart::Refusal {
|
||||
refusal: refusal.clone(),
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
OpenResponsesItem::Message {
|
||||
id: id.clone(),
|
||||
role: role_str.to_string(),
|
||||
content: parts,
|
||||
}
|
||||
}
|
||||
|
||||
Item::FunctionCall {
|
||||
id,
|
||||
call_id,
|
||||
name,
|
||||
arguments,
|
||||
status: _,
|
||||
} => OpenResponsesItem::FunctionCall {
|
||||
id: id.clone(),
|
||||
call_id: call_id.clone(),
|
||||
name: name.clone(),
|
||||
arguments: arguments.clone(),
|
||||
},
|
||||
|
||||
Item::FunctionCallOutput {
|
||||
id,
|
||||
call_id,
|
||||
output,
|
||||
} => OpenResponsesItem::FunctionCallOutput {
|
||||
id: id.clone(),
|
||||
call_id: call_id.clone(),
|
||||
output: output.clone(),
|
||||
},
|
||||
|
||||
Item::Reasoning {
|
||||
id,
|
||||
text,
|
||||
status: _,
|
||||
} => OpenResponsesItem::Reasoning {
|
||||
id: id.clone(),
|
||||
text: text.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_tool(tool: &ToolDefinition) -> OpenResponsesTool {
|
||||
OpenResponsesTool {
|
||||
r#type: "function".to_string(),
|
||||
name: tool.name.clone(),
|
||||
description: tool.description.clone(),
|
||||
parameters: tool.input_schema.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::llm_client::types::Item;
|
||||
|
||||
#[test]
|
||||
fn test_build_simple_request() {
|
||||
let request = Request::new()
|
||||
.system("You are a helpful assistant.")
|
||||
.user("Hello!");
|
||||
|
||||
let or_req = build_request("gpt-4o", &request);
|
||||
|
||||
assert_eq!(or_req.model, "gpt-4o");
|
||||
assert_eq!(
|
||||
or_req.instructions,
|
||||
Some("You are a helpful assistant.".to_string())
|
||||
);
|
||||
assert_eq!(or_req.input.len(), 1);
|
||||
assert!(or_req.stream);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_request_with_tool() {
|
||||
let request = Request::new().user("What's the weather?").tool(
|
||||
ToolDefinition::new("get_weather")
|
||||
.description("Get current weather")
|
||||
.input_schema(serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": { "type": "string" }
|
||||
},
|
||||
"required": ["location"]
|
||||
})),
|
||||
);
|
||||
|
||||
let or_req = build_request("gpt-4o", &request);
|
||||
|
||||
assert_eq!(or_req.tools.len(), 1);
|
||||
assert_eq!(or_req.tools[0].name, "get_weather");
|
||||
assert_eq!(or_req.tools[0].r#type, "function");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_function_call_and_output() {
|
||||
let request = Request::new()
|
||||
.user("What's the weather?")
|
||||
.item(Item::function_call(
|
||||
"call_123",
|
||||
"get_weather",
|
||||
r#"{"city":"Tokyo"}"#,
|
||||
))
|
||||
.item(Item::function_call_output("call_123", "Sunny, 25°C"));
|
||||
|
||||
let or_req = build_request("gpt-4o", &request);
|
||||
|
||||
assert_eq!(or_req.input.len(), 3);
|
||||
|
||||
// Check function call
|
||||
if let OpenResponsesItem::FunctionCall { call_id, name, .. } = &or_req.input[1] {
|
||||
assert_eq!(call_id, "call_123");
|
||||
assert_eq!(name, "get_weather");
|
||||
} else {
|
||||
panic!("Expected FunctionCall");
|
||||
}
|
||||
|
||||
// Check function call output
|
||||
if let OpenResponsesItem::FunctionCallOutput { call_id, output, .. } = &or_req.input[2] {
|
||||
assert_eq!(call_id, "call_123");
|
||||
assert_eq!(output, "Sunny, 25°C");
|
||||
} else {
|
||||
panic!("Expected FunctionCallOutput");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,567 @@
|
||||
//! LLM Client Common Types - Open Responses Native
|
||||
//!
|
||||
//! This module defines types that are natively aligned with the Open Responses specification.
|
||||
//! The core abstraction is `Item` which represents different types of conversation elements:
|
||||
//! - Message items (user/assistant messages with content parts)
|
||||
//! - FunctionCall items (tool invocations)
|
||||
//! - FunctionCallOutput items (tool results)
|
||||
//! - Reasoning items (extended thinking)
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
// ============================================================================
|
||||
// Item - The core unit of conversation
|
||||
// ============================================================================
|
||||
|
||||
/// Item ID type for tracking items in a conversation
|
||||
pub type ItemId = String;
|
||||
|
||||
/// Call ID type for linking function calls to their outputs
|
||||
pub type CallId = String;
|
||||
|
||||
/// Conversation item - the primary unit in Open Responses
|
||||
///
|
||||
/// Items represent discrete elements in a conversation. Unlike traditional
|
||||
/// message-based APIs, Open Responses treats tool calls and reasoning as
|
||||
/// first-class items rather than parts of messages.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use llm_worker::Item;
|
||||
///
|
||||
/// // User message
|
||||
/// let user_item = Item::user_message("Hello!");
|
||||
///
|
||||
/// // Assistant message
|
||||
/// let assistant_item = Item::assistant_message("Hi there!");
|
||||
///
|
||||
/// // Function call
|
||||
/// let call = Item::function_call("call_123", "get_weather", json!({"city": "Tokyo"}));
|
||||
///
|
||||
/// // Function call output
|
||||
/// let result = Item::function_call_output("call_123", "Sunny, 25°C");
|
||||
/// ```
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum Item {
|
||||
/// User or assistant message with content parts
|
||||
Message {
|
||||
/// Optional item ID
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
id: Option<ItemId>,
|
||||
/// Message role
|
||||
role: Role,
|
||||
/// Content parts
|
||||
content: Vec<ContentPart>,
|
||||
/// Item status
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
status: Option<ItemStatus>,
|
||||
},
|
||||
|
||||
/// Function (tool) call from the assistant
|
||||
FunctionCall {
|
||||
/// Optional item ID
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
id: Option<ItemId>,
|
||||
/// Call ID for linking to output
|
||||
call_id: CallId,
|
||||
/// Function name
|
||||
name: String,
|
||||
/// Function arguments as JSON string
|
||||
arguments: String,
|
||||
/// Item status
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
status: Option<ItemStatus>,
|
||||
},
|
||||
|
||||
/// Function (tool) call output/result
|
||||
FunctionCallOutput {
|
||||
/// Optional item ID
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
id: Option<ItemId>,
|
||||
/// Call ID linking to the function call
|
||||
call_id: CallId,
|
||||
/// Output content
|
||||
output: String,
|
||||
},
|
||||
|
||||
/// Reasoning/thinking item
|
||||
Reasoning {
|
||||
/// Optional item ID
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
id: Option<ItemId>,
|
||||
/// Reasoning text
|
||||
text: String,
|
||||
/// Item status
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
status: Option<ItemStatus>,
|
||||
},
|
||||
}
|
||||
|
||||
impl Item {
|
||||
// ========================================================================
|
||||
// Message constructors
|
||||
// ========================================================================
|
||||
|
||||
/// Create a user message item with text content
|
||||
pub fn user_message(text: impl Into<String>) -> Self {
|
||||
Self::Message {
|
||||
id: None,
|
||||
role: Role::User,
|
||||
content: vec![ContentPart::InputText {
|
||||
text: text.into(),
|
||||
}],
|
||||
status: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a user message item with multiple content parts
|
||||
pub fn user_message_parts(parts: Vec<ContentPart>) -> Self {
|
||||
Self::Message {
|
||||
id: None,
|
||||
role: Role::User,
|
||||
content: parts,
|
||||
status: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create an assistant message item with text content
|
||||
pub fn assistant_message(text: impl Into<String>) -> Self {
|
||||
Self::Message {
|
||||
id: None,
|
||||
role: Role::Assistant,
|
||||
content: vec![ContentPart::OutputText {
|
||||
text: text.into(),
|
||||
}],
|
||||
status: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create an assistant message item with multiple content parts
|
||||
pub fn assistant_message_parts(parts: Vec<ContentPart>) -> Self {
|
||||
Self::Message {
|
||||
id: None,
|
||||
role: Role::Assistant,
|
||||
content: parts,
|
||||
status: None,
|
||||
}
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// Function call constructors
|
||||
// ========================================================================
|
||||
|
||||
/// Create a function call item
|
||||
pub fn function_call(
|
||||
call_id: impl Into<String>,
|
||||
name: impl Into<String>,
|
||||
arguments: impl Into<String>,
|
||||
) -> Self {
|
||||
Self::FunctionCall {
|
||||
id: None,
|
||||
call_id: call_id.into(),
|
||||
name: name.into(),
|
||||
arguments: arguments.into(),
|
||||
status: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a function call item from a JSON value
|
||||
pub fn function_call_json(
|
||||
call_id: impl Into<String>,
|
||||
name: impl Into<String>,
|
||||
arguments: serde_json::Value,
|
||||
) -> Self {
|
||||
Self::function_call(call_id, name, arguments.to_string())
|
||||
}
|
||||
|
||||
/// Create a function call output item
|
||||
pub fn function_call_output(call_id: impl Into<String>, output: impl Into<String>) -> Self {
|
||||
Self::FunctionCallOutput {
|
||||
id: None,
|
||||
call_id: call_id.into(),
|
||||
output: output.into(),
|
||||
}
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// Reasoning constructors
|
||||
// ========================================================================
|
||||
|
||||
/// Create a reasoning item
|
||||
pub fn reasoning(text: impl Into<String>) -> Self {
|
||||
Self::Reasoning {
|
||||
id: None,
|
||||
text: text.into(),
|
||||
status: None,
|
||||
}
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// Builder methods
|
||||
// ========================================================================
|
||||
|
||||
/// Set the item ID
|
||||
pub fn with_id(mut self, id: impl Into<String>) -> Self {
|
||||
match &mut self {
|
||||
Self::Message { id: item_id, .. } => *item_id = Some(id.into()),
|
||||
Self::FunctionCall { id: item_id, .. } => *item_id = Some(id.into()),
|
||||
Self::FunctionCallOutput { id: item_id, .. } => *item_id = Some(id.into()),
|
||||
Self::Reasoning { id: item_id, .. } => *item_id = Some(id.into()),
|
||||
}
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the item status
|
||||
pub fn with_status(mut self, new_status: ItemStatus) -> Self {
|
||||
match &mut self {
|
||||
Self::Message { status, .. } => *status = Some(new_status),
|
||||
Self::FunctionCall { status, .. } => *status = Some(new_status),
|
||||
Self::FunctionCallOutput { .. } => {} // Output items don't have status
|
||||
Self::Reasoning { status, .. } => *status = Some(new_status),
|
||||
}
|
||||
self
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// Accessors
|
||||
// ========================================================================
|
||||
|
||||
/// Get the item ID if set
|
||||
pub fn id(&self) -> Option<&str> {
|
||||
match self {
|
||||
Self::Message { id, .. } => id.as_deref(),
|
||||
Self::FunctionCall { id, .. } => id.as_deref(),
|
||||
Self::FunctionCallOutput { id, .. } => id.as_deref(),
|
||||
Self::Reasoning { id, .. } => id.as_deref(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the item type as a string
|
||||
pub fn item_type(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Message { .. } => "message",
|
||||
Self::FunctionCall { .. } => "function_call",
|
||||
Self::FunctionCallOutput { .. } => "function_call_output",
|
||||
Self::Reasoning { .. } => "reasoning",
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if this is a user message
|
||||
pub fn is_user_message(&self) -> bool {
|
||||
matches!(self, Self::Message { role: Role::User, .. })
|
||||
}
|
||||
|
||||
/// Check if this is an assistant message
|
||||
pub fn is_assistant_message(&self) -> bool {
|
||||
matches!(self, Self::Message { role: Role::Assistant, .. })
|
||||
}
|
||||
|
||||
/// Check if this is a function call
|
||||
pub fn is_function_call(&self) -> bool {
|
||||
matches!(self, Self::FunctionCall { .. })
|
||||
}
|
||||
|
||||
/// Check if this is a function call output
|
||||
pub fn is_function_call_output(&self) -> bool {
|
||||
matches!(self, Self::FunctionCallOutput { .. })
|
||||
}
|
||||
|
||||
/// Check if this is a reasoning item
|
||||
pub fn is_reasoning(&self) -> bool {
|
||||
matches!(self, Self::Reasoning { .. })
|
||||
}
|
||||
|
||||
/// Get text content if this is a simple text message
|
||||
pub fn as_text(&self) -> Option<&str> {
|
||||
match self {
|
||||
Self::Message { content, .. } if content.len() == 1 => match &content[0] {
|
||||
ContentPart::InputText { text } => Some(text),
|
||||
ContentPart::OutputText { text } => Some(text),
|
||||
_ => None,
|
||||
},
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Content Parts - Components within message items
|
||||
// ============================================================================
|
||||
|
||||
/// Content part within a message item
|
||||
///
|
||||
/// Open Responses distinguishes between input and output content types.
|
||||
/// Input types are used in user messages, output types in assistant messages.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum ContentPart {
|
||||
/// Input text (for user messages)
|
||||
InputText {
|
||||
/// The text content
|
||||
text: String,
|
||||
},
|
||||
|
||||
/// Output text (for assistant messages)
|
||||
OutputText {
|
||||
/// The text content
|
||||
text: String,
|
||||
},
|
||||
|
||||
/// Refusal content (for assistant messages)
|
||||
Refusal {
|
||||
/// The refusal message
|
||||
refusal: String,
|
||||
},
|
||||
// Future: InputAudio, OutputAudio, etc.
|
||||
}
|
||||
|
||||
impl ContentPart {
|
||||
/// Create an input text part
|
||||
pub fn input_text(text: impl Into<String>) -> Self {
|
||||
Self::InputText { text: text.into() }
|
||||
}
|
||||
|
||||
/// Create an output text part
|
||||
pub fn output_text(text: impl Into<String>) -> Self {
|
||||
Self::OutputText { text: text.into() }
|
||||
}
|
||||
|
||||
/// Create a refusal part
|
||||
pub fn refusal(refusal: impl Into<String>) -> Self {
|
||||
Self::Refusal {
|
||||
refusal: refusal.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the text content regardless of type
|
||||
pub fn as_text(&self) -> &str {
|
||||
match self {
|
||||
Self::InputText { text } => text,
|
||||
Self::OutputText { text } => text,
|
||||
Self::Refusal { refusal } => refusal,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Role and Status
|
||||
// ============================================================================
|
||||
|
||||
/// Message role
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum Role {
|
||||
/// User
|
||||
User,
|
||||
/// Assistant
|
||||
Assistant,
|
||||
/// System (for system prompts, not typically used in items)
|
||||
System,
|
||||
}
|
||||
|
||||
/// Item status
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum ItemStatus {
|
||||
/// Item is being generated
|
||||
InProgress,
|
||||
/// Item completed successfully
|
||||
Completed,
|
||||
/// Item was truncated (e.g., max tokens)
|
||||
Incomplete,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Request Types
|
||||
// ============================================================================
|
||||
|
||||
/// LLM Request
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct Request {
|
||||
/// System prompt (instructions)
|
||||
pub system_prompt: Option<String>,
|
||||
/// Input items (conversation history)
|
||||
pub items: Vec<Item>,
|
||||
/// Tool definitions
|
||||
pub tools: Vec<ToolDefinition>,
|
||||
/// Request configuration
|
||||
pub config: RequestConfig,
|
||||
}
|
||||
|
||||
impl Request {
|
||||
/// Create a new empty request
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
/// Set the system prompt
|
||||
pub fn system(mut self, prompt: impl Into<String>) -> Self {
|
||||
self.system_prompt = Some(prompt.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// Add a user message
|
||||
pub fn user(mut self, content: impl Into<String>) -> Self {
|
||||
self.items.push(Item::user_message(content));
|
||||
self
|
||||
}
|
||||
|
||||
/// Add an assistant message
|
||||
pub fn assistant(mut self, content: impl Into<String>) -> Self {
|
||||
self.items.push(Item::assistant_message(content));
|
||||
self
|
||||
}
|
||||
|
||||
/// Add an item
|
||||
pub fn item(mut self, item: Item) -> Self {
|
||||
self.items.push(item);
|
||||
self
|
||||
}
|
||||
|
||||
/// Add multiple items
|
||||
pub fn items(mut self, items: impl IntoIterator<Item = Item>) -> Self {
|
||||
self.items.extend(items);
|
||||
self
|
||||
}
|
||||
|
||||
/// Add a tool definition
|
||||
pub fn tool(mut self, tool: ToolDefinition) -> Self {
|
||||
self.tools.push(tool);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the request config
|
||||
pub fn config(mut self, config: RequestConfig) -> Self {
|
||||
self.config = config;
|
||||
self
|
||||
}
|
||||
|
||||
/// Set max tokens
|
||||
pub fn max_tokens(mut self, max_tokens: u32) -> Self {
|
||||
self.config.max_tokens = Some(max_tokens);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set temperature
|
||||
pub fn temperature(mut self, temperature: f32) -> Self {
|
||||
self.config.temperature = Some(temperature);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set top_p
|
||||
pub fn top_p(mut self, top_p: f32) -> Self {
|
||||
self.config.top_p = Some(top_p);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set top_k
|
||||
pub fn top_k(mut self, top_k: u32) -> Self {
|
||||
self.config.top_k = Some(top_k);
|
||||
self
|
||||
}
|
||||
|
||||
/// Add a stop sequence
|
||||
pub fn stop_sequence(mut self, sequence: impl Into<String>) -> Self {
|
||||
self.config.stop_sequences.push(sequence.into());
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Tool Definition
|
||||
// ============================================================================
|
||||
|
||||
/// Tool (function) definition
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolDefinition {
|
||||
/// Tool name
|
||||
pub name: String,
|
||||
/// Tool description
|
||||
pub description: Option<String>,
|
||||
/// Input schema (JSON Schema)
|
||||
pub input_schema: serde_json::Value,
|
||||
}
|
||||
|
||||
impl ToolDefinition {
|
||||
/// Create a new tool definition
|
||||
pub fn new(name: impl Into<String>) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
description: None,
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the description
|
||||
pub fn description(mut self, desc: impl Into<String>) -> Self {
|
||||
self.description = Some(desc.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the input schema
|
||||
pub fn input_schema(mut self, schema: serde_json::Value) -> Self {
|
||||
self.input_schema = schema;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Request Config
|
||||
// ============================================================================
|
||||
|
||||
/// Request configuration
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct RequestConfig {
|
||||
/// Maximum tokens to generate
|
||||
pub max_tokens: Option<u32>,
|
||||
/// Temperature (randomness)
|
||||
pub temperature: Option<f32>,
|
||||
/// Top P (nucleus sampling)
|
||||
pub top_p: Option<f32>,
|
||||
/// Top K
|
||||
pub top_k: Option<u32>,
|
||||
/// Stop sequences
|
||||
pub stop_sequences: Vec<String>,
|
||||
}
|
||||
|
||||
impl RequestConfig {
|
||||
/// Create a new default config
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
/// Set max tokens
|
||||
pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
|
||||
self.max_tokens = Some(max_tokens);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set temperature
|
||||
pub fn with_temperature(mut self, temperature: f32) -> Self {
|
||||
self.temperature = Some(temperature);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set top_p
|
||||
pub fn with_top_p(mut self, top_p: f32) -> Self {
|
||||
self.top_p = Some(top_p);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set top_k
|
||||
pub fn with_top_k(mut self, top_k: u32) -> Self {
|
||||
self.top_k = Some(top_k);
|
||||
self
|
||||
}
|
||||
|
||||
/// Add a stop sequence
|
||||
pub fn with_stop_sequence(mut self, sequence: impl Into<String>) -> Self {
|
||||
self.stop_sequences.push(sequence.into());
|
||||
self
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
//! Message and Item Types
|
||||
//!
|
||||
//! This module provides the core types for representing conversation items
|
||||
//! in the Open Responses format.
|
||||
//!
|
||||
//! The primary type is [`Item`], which represents different kinds of conversation
|
||||
//! elements: messages, function calls, function call outputs, and reasoning.
|
||||
|
||||
// Re-export all types from llm_client::types
|
||||
pub use crate::llm_client::types::{ContentPart, Item, Role};
|
||||
|
||||
/// Convenience alias for backward compatibility
|
||||
///
|
||||
/// In the Open Responses model, messages are just one type of Item.
|
||||
/// This alias allows code that expects a "Message" type to continue working.
|
||||
pub type Message = Item;
|
||||
@@ -0,0 +1,60 @@
|
||||
//! Worker State
|
||||
//!
|
||||
//! State marker types for cache protection using the Type-state pattern.
|
||||
//! Worker has state transitions from `Mutable` → `CacheLocked`.
|
||||
|
||||
/// Marker trait representing Worker state
|
||||
///
|
||||
/// This trait is sealed and cannot be implemented externally.
|
||||
pub trait WorkerState: private::Sealed + Send + Sync + 'static {}
|
||||
|
||||
mod private {
|
||||
pub trait Sealed {}
|
||||
}
|
||||
|
||||
/// Mutable state (editable)
|
||||
///
|
||||
/// In this state, the following operations are available:
|
||||
/// - Setting/changing system prompt
|
||||
/// - Editing message history (add, delete, clear)
|
||||
/// - Registering tools and hooks
|
||||
///
|
||||
/// Can transition to [`CacheLocked`] state via `Worker::lock()`.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use llm_worker::Worker;
|
||||
///
|
||||
/// let mut worker = Worker::new(client)
|
||||
/// .system_prompt("You are helpful.");
|
||||
///
|
||||
/// // History can be edited
|
||||
/// worker.push_message(Message::user("Hello"));
|
||||
/// worker.clear_history();
|
||||
///
|
||||
/// // Lock to protected state
|
||||
/// let locked = worker.lock();
|
||||
/// ```
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
pub struct Mutable;
|
||||
|
||||
impl private::Sealed for Mutable {}
|
||||
impl WorkerState for Mutable {}
|
||||
|
||||
/// Cache locked state (cache protected)
|
||||
///
|
||||
/// In this state, the following restrictions apply:
|
||||
/// - System prompt cannot be changed
|
||||
/// - Existing message history cannot be modified (only appending to the end)
|
||||
///
|
||||
/// To ensure LLM API KV cache hits,
|
||||
/// using this state during execution is recommended.
|
||||
///
|
||||
/// Can return to [`Mutable`] state via `Worker::unlock()`,
|
||||
/// but note that cache protection will be released.
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
pub struct CacheLocked;
|
||||
|
||||
impl private::Sealed for CacheLocked {}
|
||||
impl WorkerState for CacheLocked {}
|
||||
@@ -1,7 +1,7 @@
|
||||
//! イベント購読
|
||||
//! Event Subscription
|
||||
//!
|
||||
//! LLMからのストリーミングイベントをリアルタイムで受信するためのトレイト。
|
||||
//! UIへのストリーム表示やプログレス表示に使用します。
|
||||
//! Trait for receiving streaming events from LLM in real-time.
|
||||
//! Used for stream display to UI and progress display.
|
||||
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
@@ -18,23 +18,23 @@ use crate::{
|
||||
// WorkerSubscriber Trait
|
||||
// =============================================================================
|
||||
|
||||
/// LLMからのストリーミングイベントを購読するトレイト
|
||||
/// Trait for subscribing to streaming events from LLM
|
||||
///
|
||||
/// Workerに登録すると、テキスト生成やツール呼び出しのイベントを
|
||||
/// リアルタイムで受信できます。UIへのストリーム表示に最適です。
|
||||
/// When registered with Worker, you can receive events from text generation
|
||||
/// and tool calls in real-time. Ideal for stream display to UI.
|
||||
///
|
||||
/// # 受信できるイベント
|
||||
/// # Available Events
|
||||
///
|
||||
/// - **ブロックイベント**: テキスト、ツール使用(スコープ付き)
|
||||
/// - **メタイベント**: 使用量、ステータス、エラー
|
||||
/// - **完了イベント**: テキスト完了、ツール呼び出し完了
|
||||
/// - **ターン制御**: ターン開始、ターン終了
|
||||
/// - **Block events**: Text, tool use (with scope)
|
||||
/// - **Meta events**: Usage, status, error
|
||||
/// - **Completion events**: Text complete, tool call complete
|
||||
/// - **Turn control**: Turn start, turn end
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::subscriber::WorkerSubscriber;
|
||||
/// use worker::timeline::TextBlockEvent;
|
||||
/// use llm_worker::subscriber::WorkerSubscriber;
|
||||
/// use llm_worker::timeline::TextBlockEvent;
|
||||
///
|
||||
/// struct StreamPrinter;
|
||||
///
|
||||
@@ -44,7 +44,7 @@ use crate::{
|
||||
///
|
||||
/// fn on_text_block(&mut self, _: &mut (), event: &TextBlockEvent) {
|
||||
/// if let TextBlockEvent::Delta(text) = event {
|
||||
/// print!("{}", text); // リアルタイム出力
|
||||
/// print!("{}", text); // Real-time output
|
||||
/// }
|
||||
/// }
|
||||
///
|
||||
@@ -53,37 +53,37 @@ use crate::{
|
||||
/// }
|
||||
/// }
|
||||
///
|
||||
/// // Workerに登録
|
||||
/// // Register with Worker
|
||||
/// worker.subscribe(StreamPrinter);
|
||||
/// ```
|
||||
pub trait WorkerSubscriber: Send {
|
||||
// =========================================================================
|
||||
// スコープ型(ブロックイベント用)
|
||||
// Scope Types (for block events)
|
||||
// =========================================================================
|
||||
|
||||
/// テキストブロック処理用のスコープ型
|
||||
/// Scope type for text block processing
|
||||
///
|
||||
/// ブロック開始時にDefault::default()で生成され、
|
||||
/// ブロック終了時に破棄される。
|
||||
type TextBlockScope: Default + Send;
|
||||
/// Generated with Default::default() at block start,
|
||||
/// destroyed at block end.
|
||||
type TextBlockScope: Default + Send + Sync;
|
||||
|
||||
/// ツール使用ブロック処理用のスコープ型
|
||||
type ToolUseBlockScope: Default + Send;
|
||||
/// Scope type for tool use block processing
|
||||
type ToolUseBlockScope: Default + Send + Sync;
|
||||
|
||||
// =========================================================================
|
||||
// ブロックイベント(スコープ管理あり)
|
||||
// Block Events (with scope management)
|
||||
// =========================================================================
|
||||
|
||||
/// テキストブロックイベント
|
||||
/// Text block event
|
||||
///
|
||||
/// Start/Delta/Stopのライフサイクルを持つ。
|
||||
/// scopeはブロック開始時に生成され、終了時に破棄される。
|
||||
/// Has Start/Delta/Stop lifecycle.
|
||||
/// Scope is generated at block start and destroyed at end.
|
||||
#[allow(unused_variables)]
|
||||
fn on_text_block(&mut self, scope: &mut Self::TextBlockScope, event: &TextBlockEvent) {}
|
||||
|
||||
/// ツール使用ブロックイベント
|
||||
/// Tool use block event
|
||||
///
|
||||
/// Start/InputJsonDelta/Stopのライフサイクルを持つ。
|
||||
/// Has Start/InputJsonDelta/Stop lifecycle.
|
||||
#[allow(unused_variables)]
|
||||
fn on_tool_use_block(
|
||||
&mut self,
|
||||
@@ -93,62 +93,62 @@ pub trait WorkerSubscriber: Send {
|
||||
}
|
||||
|
||||
// =========================================================================
|
||||
// 単発イベント(スコープ不要)
|
||||
// Single Events (no scope needed)
|
||||
// =========================================================================
|
||||
|
||||
/// 使用量イベント
|
||||
/// Usage event
|
||||
#[allow(unused_variables)]
|
||||
fn on_usage(&mut self, event: &UsageEvent) {}
|
||||
|
||||
/// ステータスイベント
|
||||
/// Status event
|
||||
#[allow(unused_variables)]
|
||||
fn on_status(&mut self, event: &StatusEvent) {}
|
||||
|
||||
/// エラーイベント
|
||||
/// Error event
|
||||
#[allow(unused_variables)]
|
||||
fn on_error(&mut self, event: &ErrorEvent) {}
|
||||
|
||||
// =========================================================================
|
||||
// 累積イベント(Worker層で追加)
|
||||
// Accumulated Events (added in Worker layer)
|
||||
// =========================================================================
|
||||
|
||||
/// テキスト完了イベント
|
||||
/// Text complete event
|
||||
///
|
||||
/// テキストブロックが完了した時点で、累積されたテキスト全体が渡される。
|
||||
/// ブロック処理後の最終結果を受け取るのに便利。
|
||||
/// When a text block completes, the entire accumulated text is passed.
|
||||
/// Convenient for receiving the final result after block processing.
|
||||
#[allow(unused_variables)]
|
||||
fn on_text_complete(&mut self, text: &str) {}
|
||||
|
||||
/// ツール呼び出し完了イベント
|
||||
/// Tool call complete event
|
||||
///
|
||||
/// ツール使用ブロックが完了した時点で、完全なToolCallが渡される。
|
||||
/// When a tool use block completes, the complete ToolCall is passed.
|
||||
#[allow(unused_variables)]
|
||||
fn on_tool_call_complete(&mut self, call: &ToolCall) {}
|
||||
|
||||
// =========================================================================
|
||||
// ターン制御
|
||||
// Turn Control
|
||||
// =========================================================================
|
||||
|
||||
/// ターン開始時
|
||||
/// On turn start
|
||||
///
|
||||
/// `turn`は0から始まるターン番号。
|
||||
/// `turn` is a 0-based turn number.
|
||||
#[allow(unused_variables)]
|
||||
fn on_turn_start(&mut self, turn: usize) {}
|
||||
|
||||
/// ターン終了時
|
||||
/// On turn end
|
||||
#[allow(unused_variables)]
|
||||
fn on_turn_end(&mut self, turn: usize) {}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// SubscriberAdapter - WorkerSubscriberをTimelineハンドラにブリッジ
|
||||
// SubscriberAdapter - Bridge WorkerSubscriber to Timeline handlers
|
||||
// =============================================================================
|
||||
|
||||
// =============================================================================
|
||||
// TextBlock Handler Adapter
|
||||
// =============================================================================
|
||||
|
||||
/// TextBlockKind用のSubscriberアダプター
|
||||
/// Subscriber adapter for TextBlockKind
|
||||
pub(crate) struct TextBlockSubscriberAdapter<S: WorkerSubscriber> {
|
||||
subscriber: Arc<Mutex<S>>,
|
||||
}
|
||||
@@ -167,10 +167,10 @@ impl<S: WorkerSubscriber> Clone for TextBlockSubscriberAdapter<S> {
|
||||
}
|
||||
}
|
||||
|
||||
/// TextBlockのスコープをラップ
|
||||
/// Wrapper for TextBlock scope
|
||||
pub struct TextBlockScopeWrapper<S: WorkerSubscriber> {
|
||||
inner: S::TextBlockScope,
|
||||
buffer: String, // on_text_complete用のバッファ
|
||||
buffer: String, // Buffer for on_text_complete
|
||||
}
|
||||
|
||||
impl<S: WorkerSubscriber> Default for TextBlockScopeWrapper<S> {
|
||||
@@ -186,16 +186,16 @@ impl<S: WorkerSubscriber + 'static> Handler<TextBlockKind> for TextBlockSubscrib
|
||||
type Scope = TextBlockScopeWrapper<S>;
|
||||
|
||||
fn on_event(&mut self, scope: &mut Self::Scope, event: &TextBlockEvent) {
|
||||
// Deltaの場合はバッファに蓄積
|
||||
// Accumulate deltas into buffer
|
||||
if let TextBlockEvent::Delta(text) = event {
|
||||
scope.buffer.push_str(text);
|
||||
}
|
||||
|
||||
// SubscriberのTextBlockイベントハンドラを呼び出し
|
||||
// Call Subscriber's TextBlock event handler
|
||||
if let Ok(mut subscriber) = self.subscriber.lock() {
|
||||
subscriber.on_text_block(&mut scope.inner, event);
|
||||
|
||||
// Stopの場合はon_text_completeも呼び出し
|
||||
// Also call on_text_complete on Stop
|
||||
if matches!(event, TextBlockEvent::Stop(_)) {
|
||||
subscriber.on_text_complete(&scope.buffer);
|
||||
}
|
||||
@@ -207,7 +207,7 @@ impl<S: WorkerSubscriber + 'static> Handler<TextBlockKind> for TextBlockSubscrib
|
||||
// ToolUseBlock Handler Adapter
|
||||
// =============================================================================
|
||||
|
||||
/// ToolUseBlockKind用のSubscriberアダプター
|
||||
/// Subscriber adapter for ToolUseBlockKind
|
||||
pub(crate) struct ToolUseBlockSubscriberAdapter<S: WorkerSubscriber> {
|
||||
subscriber: Arc<Mutex<S>>,
|
||||
}
|
||||
@@ -226,12 +226,12 @@ impl<S: WorkerSubscriber> Clone for ToolUseBlockSubscriberAdapter<S> {
|
||||
}
|
||||
}
|
||||
|
||||
/// ToolUseBlockのスコープをラップ
|
||||
/// Wrapper for ToolUseBlock scope
|
||||
pub struct ToolUseBlockScopeWrapper<S: WorkerSubscriber> {
|
||||
inner: S::ToolUseBlockScope,
|
||||
id: String,
|
||||
name: String,
|
||||
input_json: String, // JSON蓄積用
|
||||
input_json: String, // JSON accumulation
|
||||
}
|
||||
|
||||
impl<S: WorkerSubscriber> Default for ToolUseBlockScopeWrapper<S> {
|
||||
@@ -249,22 +249,22 @@ impl<S: WorkerSubscriber + 'static> Handler<ToolUseBlockKind> for ToolUseBlockSu
|
||||
type Scope = ToolUseBlockScopeWrapper<S>;
|
||||
|
||||
fn on_event(&mut self, scope: &mut Self::Scope, event: &ToolUseBlockEvent) {
|
||||
// Start時にメタデータを保存
|
||||
// Save metadata on Start
|
||||
if let ToolUseBlockEvent::Start(start) = event {
|
||||
scope.id = start.id.clone();
|
||||
scope.name = start.name.clone();
|
||||
}
|
||||
|
||||
// InputJsonDeltaの場合はバッファに蓄積
|
||||
// Accumulate InputJsonDelta into buffer
|
||||
if let ToolUseBlockEvent::InputJsonDelta(json) = event {
|
||||
scope.input_json.push_str(json);
|
||||
}
|
||||
|
||||
// SubscriberのToolUseBlockイベントハンドラを呼び出し
|
||||
// Call Subscriber's ToolUseBlock event handler
|
||||
if let Ok(mut subscriber) = self.subscriber.lock() {
|
||||
subscriber.on_tool_use_block(&mut scope.inner, event);
|
||||
|
||||
// Stopの場合はon_tool_call_completeも呼び出し
|
||||
// Also call on_tool_call_complete on Stop
|
||||
if matches!(event, ToolUseBlockEvent::Stop(_)) {
|
||||
let input: serde_json::Value =
|
||||
serde_json::from_str(&scope.input_json).unwrap_or_default();
|
||||
@@ -283,7 +283,7 @@ impl<S: WorkerSubscriber + 'static> Handler<ToolUseBlockKind> for ToolUseBlockSu
|
||||
// Meta Event Handler Adapters
|
||||
// =============================================================================
|
||||
|
||||
/// UsageKind用のSubscriberアダプター
|
||||
/// Subscriber adapter for UsageKind
|
||||
pub(crate) struct UsageSubscriberAdapter<S: WorkerSubscriber> {
|
||||
subscriber: Arc<Mutex<S>>,
|
||||
}
|
||||
@@ -312,7 +312,7 @@ impl<S: WorkerSubscriber + 'static> Handler<UsageKind> for UsageSubscriberAdapte
|
||||
}
|
||||
}
|
||||
|
||||
/// StatusKind用のSubscriberアダプター
|
||||
/// Subscriber adapter for StatusKind
|
||||
pub(crate) struct StatusSubscriberAdapter<S: WorkerSubscriber> {
|
||||
subscriber: Arc<Mutex<S>>,
|
||||
}
|
||||
@@ -341,7 +341,7 @@ impl<S: WorkerSubscriber + 'static> Handler<StatusKind> for StatusSubscriberAdap
|
||||
}
|
||||
}
|
||||
|
||||
/// ErrorKind用のSubscriberアダプター
|
||||
/// Subscriber adapter for ErrorKind
|
||||
pub(crate) struct ErrorSubscriberAdapter<S: WorkerSubscriber> {
|
||||
subscriber: Arc<Mutex<S>>,
|
||||
}
|
||||
@@ -17,7 +17,7 @@ use crate::handler::*;
|
||||
/// 各Handlerは独自のScope型を持つため、Timelineで保持するには型消去が必要です。
|
||||
/// 通常は直接使用せず、`Timeline::on_text_block()`などのメソッド経由で
|
||||
/// 自動的にラップされます。
|
||||
pub trait ErasedHandler<K: Kind>: Send {
|
||||
pub trait ErasedHandler<K: Kind>: Send + Sync {
|
||||
/// イベントをディスパッチ
|
||||
fn dispatch(&mut self, event: &K::Event);
|
||||
/// スコープを開始(Block開始時)
|
||||
@@ -54,9 +54,9 @@ where
|
||||
|
||||
impl<H, K> ErasedHandler<K> for HandlerWrapper<H, K>
|
||||
where
|
||||
H: Handler<K> + Send,
|
||||
H: Handler<K> + Send + Sync,
|
||||
K: Kind,
|
||||
H::Scope: Send,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
fn dispatch(&mut self, event: &K::Event) {
|
||||
if let Some(scope) = &mut self.scope {
|
||||
@@ -78,7 +78,7 @@ where
|
||||
// =============================================================================
|
||||
|
||||
/// ブロックハンドラーの型消去trait
|
||||
trait ErasedBlockHandler: Send {
|
||||
trait ErasedBlockHandler: Send + Sync {
|
||||
fn dispatch_start(&mut self, start: &BlockStart);
|
||||
fn dispatch_delta(&mut self, delta: &BlockDelta);
|
||||
fn dispatch_stop(&mut self, stop: &BlockStop);
|
||||
@@ -112,8 +112,8 @@ where
|
||||
|
||||
impl<H> ErasedBlockHandler for TextBlockHandlerWrapper<H>
|
||||
where
|
||||
H: Handler<TextBlockKind> + Send,
|
||||
H::Scope: Send,
|
||||
H: Handler<TextBlockKind> + Send + Sync,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
fn dispatch_start(&mut self, start: &BlockStart) {
|
||||
if let Some(scope) = &mut self.scope {
|
||||
@@ -185,8 +185,8 @@ where
|
||||
|
||||
impl<H> ErasedBlockHandler for ThinkingBlockHandlerWrapper<H>
|
||||
where
|
||||
H: Handler<ThinkingBlockKind> + Send,
|
||||
H::Scope: Send,
|
||||
H: Handler<ThinkingBlockKind> + Send + Sync,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
fn dispatch_start(&mut self, start: &BlockStart) {
|
||||
if let Some(scope) = &mut self.scope {
|
||||
@@ -255,8 +255,8 @@ where
|
||||
|
||||
impl<H> ErasedBlockHandler for ToolUseBlockHandlerWrapper<H>
|
||||
where
|
||||
H: Handler<ToolUseBlockKind> + Send,
|
||||
H::Scope: Send,
|
||||
H: Handler<ToolUseBlockKind> + Send + Sync,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
fn dispatch_start(&mut self, start: &BlockStart) {
|
||||
if let Some(scope) = &mut self.scope {
|
||||
@@ -328,7 +328,7 @@ where
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::{Timeline, Handler, TextBlockKind, TextBlockEvent};
|
||||
/// use llm_worker::{Timeline, Handler, TextBlockKind, TextBlockEvent};
|
||||
///
|
||||
/// struct MyHandler;
|
||||
/// impl Handler<TextBlockKind> for MyHandler {
|
||||
@@ -391,8 +391,8 @@ impl Timeline {
|
||||
/// UsageKind用のHandlerを登録
|
||||
pub fn on_usage<H>(&mut self, handler: H) -> &mut Self
|
||||
where
|
||||
H: Handler<UsageKind> + Send + 'static,
|
||||
H::Scope: Send,
|
||||
H: Handler<UsageKind> + Send + Sync + 'static,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
// Meta系はデフォルトでスコープを開始しておく
|
||||
let mut wrapper = HandlerWrapper::new(handler);
|
||||
@@ -404,8 +404,8 @@ impl Timeline {
|
||||
/// PingKind用のHandlerを登録
|
||||
pub fn on_ping<H>(&mut self, handler: H) -> &mut Self
|
||||
where
|
||||
H: Handler<PingKind> + Send + 'static,
|
||||
H::Scope: Send,
|
||||
H: Handler<PingKind> + Send + Sync + 'static,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
let mut wrapper = HandlerWrapper::new(handler);
|
||||
wrapper.start_scope();
|
||||
@@ -416,8 +416,8 @@ impl Timeline {
|
||||
/// StatusKind用のHandlerを登録
|
||||
pub fn on_status<H>(&mut self, handler: H) -> &mut Self
|
||||
where
|
||||
H: Handler<StatusKind> + Send + 'static,
|
||||
H::Scope: Send,
|
||||
H: Handler<StatusKind> + Send + Sync + 'static,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
let mut wrapper = HandlerWrapper::new(handler);
|
||||
wrapper.start_scope();
|
||||
@@ -428,8 +428,8 @@ impl Timeline {
|
||||
/// ErrorKind用のHandlerを登録
|
||||
pub fn on_error<H>(&mut self, handler: H) -> &mut Self
|
||||
where
|
||||
H: Handler<ErrorKind> + Send + 'static,
|
||||
H::Scope: Send,
|
||||
H: Handler<ErrorKind> + Send + Sync + 'static,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
let mut wrapper = HandlerWrapper::new(handler);
|
||||
wrapper.start_scope();
|
||||
@@ -440,8 +440,8 @@ impl Timeline {
|
||||
/// TextBlockKind用のHandlerを登録
|
||||
pub fn on_text_block<H>(&mut self, handler: H) -> &mut Self
|
||||
where
|
||||
H: Handler<TextBlockKind> + Send + 'static,
|
||||
H::Scope: Send,
|
||||
H: Handler<TextBlockKind> + Send + Sync + 'static,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
self.text_block_handlers
|
||||
.push(Box::new(TextBlockHandlerWrapper::new(handler)));
|
||||
@@ -451,8 +451,8 @@ impl Timeline {
|
||||
/// ThinkingBlockKind用のHandlerを登録
|
||||
pub fn on_thinking_block<H>(&mut self, handler: H) -> &mut Self
|
||||
where
|
||||
H: Handler<ThinkingBlockKind> + Send + 'static,
|
||||
H::Scope: Send,
|
||||
H: Handler<ThinkingBlockKind> + Send + Sync + 'static,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
self.thinking_block_handlers
|
||||
.push(Box::new(ThinkingBlockHandlerWrapper::new(handler)));
|
||||
@@ -462,8 +462,8 @@ impl Timeline {
|
||||
/// ToolUseBlockKind用のHandlerを登録
|
||||
pub fn on_tool_use_block<H>(&mut self, handler: H) -> &mut Self
|
||||
where
|
||||
H: Handler<ToolUseBlockKind> + Send + 'static,
|
||||
H::Scope: Send,
|
||||
H: Handler<ToolUseBlockKind> + Send + Sync + 'static,
|
||||
H::Scope: Send + Sync,
|
||||
{
|
||||
self.tool_use_block_handlers
|
||||
.push(Box::new(ToolUseBlockHandlerWrapper::new(handler)));
|
||||
@@ -578,6 +578,21 @@ impl Timeline {
|
||||
pub fn current_block(&self) -> Option<BlockType> {
|
||||
self.current_block
|
||||
}
|
||||
|
||||
/// 現在アクティブなブロックを中断する
|
||||
///
|
||||
/// キャンセルやエラー時に呼び出し、進行中のブロックに対して
|
||||
/// BlockAbortイベントを発火してスコープをクリーンアップする。
|
||||
pub fn abort_current_block(&mut self) {
|
||||
if let Some(block_type) = self.current_block {
|
||||
let abort = crate::timeline::event::BlockAbort {
|
||||
index: 0, // インデックスは不明なので0
|
||||
block_type,
|
||||
reason: "Cancelled".to_string(),
|
||||
};
|
||||
self.handle_block_abort(&abort);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -0,0 +1,154 @@
|
||||
//! Tool Definition
|
||||
//!
|
||||
//! Traits for defining tools callable by LLM.
|
||||
//! Usually auto-implemented using the `#[tool]` macro.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
|
||||
/// Error during tool execution
|
||||
#[derive(Debug, Error)]
|
||||
pub enum ToolError {
|
||||
/// Invalid argument
|
||||
#[error("Invalid argument: {0}")]
|
||||
InvalidArgument(String),
|
||||
/// Execution failed
|
||||
#[error("Execution failed: {0}")]
|
||||
ExecutionFailed(String),
|
||||
/// Internal error
|
||||
#[error("Internal error: {0}")]
|
||||
Internal(String),
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ToolMeta - Immutable Meta Information
|
||||
// =============================================================================
|
||||
|
||||
/// Tool meta information (fixed at registration, immutable)
|
||||
///
|
||||
/// Generated from `ToolDefinition` factory and does not change after registration with Worker.
|
||||
/// Used for sending tool definitions to LLM.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct ToolMeta {
|
||||
/// Tool name (used by LLM for identification)
|
||||
pub name: String,
|
||||
/// Tool description (included in prompt to LLM)
|
||||
pub description: String,
|
||||
/// JSON Schema for arguments
|
||||
pub input_schema: Value,
|
||||
}
|
||||
|
||||
impl ToolMeta {
|
||||
/// Create a new ToolMeta
|
||||
pub fn new(name: impl Into<String>) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
description: String::new(),
|
||||
input_schema: Value::Object(Default::default()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the description
|
||||
pub fn description(mut self, desc: impl Into<String>) -> Self {
|
||||
self.description = desc.into();
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the argument schema
|
||||
pub fn input_schema(mut self, schema: Value) -> Self {
|
||||
self.input_schema = schema;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ToolDefinition - Factory Type
|
||||
// =============================================================================
|
||||
|
||||
/// Tool definition factory
|
||||
///
|
||||
/// When called, returns `(ToolMeta, Arc<dyn Tool>)`.
|
||||
/// Called once during Worker registration, and the meta information and instance
|
||||
/// are cached at session scope.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// let def: ToolDefinition = Arc::new(|| {
|
||||
/// (
|
||||
/// ToolMeta::new("my_tool")
|
||||
/// .description("My tool description")
|
||||
/// .input_schema(json!({"type": "object"})),
|
||||
/// Arc::new(MyToolImpl { state: 0 }) as Arc<dyn Tool>,
|
||||
/// )
|
||||
/// });
|
||||
/// worker.register_tool(def)?;
|
||||
/// ```
|
||||
pub type ToolDefinition = Arc<dyn Fn() -> (ToolMeta, Arc<dyn Tool>) + Send + Sync>;
|
||||
|
||||
// =============================================================================
|
||||
// Tool trait
|
||||
// =============================================================================
|
||||
|
||||
/// Trait for defining tools callable by LLM
|
||||
///
|
||||
/// Tools are used by LLM to access external resources
|
||||
/// or execute computations.
|
||||
/// Can maintain state during the session.
|
||||
///
|
||||
/// # How to Implement
|
||||
///
|
||||
/// Usually auto-implemented using the `#[tool_registry]` macro:
|
||||
///
|
||||
/// ```ignore
|
||||
/// #[tool_registry]
|
||||
/// impl MyApp {
|
||||
/// #[tool]
|
||||
/// async fn search(&self, query: String) -> String {
|
||||
/// format!("Results for: {}", query)
|
||||
/// }
|
||||
/// }
|
||||
///
|
||||
/// // Register
|
||||
/// worker.register_tool(app.search_definition())?;
|
||||
/// ```
|
||||
///
|
||||
/// # Manual Implementation
|
||||
///
|
||||
/// ```ignore
|
||||
/// use llm_worker::tool::{Tool, ToolError, ToolMeta, ToolDefinition};
|
||||
/// use std::sync::Arc;
|
||||
///
|
||||
/// struct MyTool { counter: std::sync::atomic::AtomicUsize }
|
||||
///
|
||||
/// #[async_trait::async_trait]
|
||||
/// impl Tool for MyTool {
|
||||
/// async fn execute(&self, input: &str) -> Result<String, ToolError> {
|
||||
/// self.counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
|
||||
/// Ok("result".to_string())
|
||||
/// }
|
||||
/// }
|
||||
///
|
||||
/// let def: ToolDefinition = Arc::new(|| {
|
||||
/// (
|
||||
/// ToolMeta::new("my_tool")
|
||||
/// .description("My custom tool")
|
||||
/// .input_schema(serde_json::json!({"type": "object"})),
|
||||
/// Arc::new(MyTool { counter: Default::default() }) as Arc<dyn Tool>,
|
||||
/// )
|
||||
/// });
|
||||
/// ```
|
||||
#[async_trait]
|
||||
pub trait Tool: Send + Sync {
|
||||
/// Execute the tool
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `input_json` - JSON-formatted arguments generated by LLM
|
||||
///
|
||||
/// # Returns
|
||||
/// Result string from execution. This content is returned to LLM.
|
||||
async fn execute(&self, input_json: &str) -> Result<String, ToolError>;
|
||||
}
|
||||
@@ -0,0 +1,182 @@
|
||||
use std::collections::HashMap;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::llm_client::ToolDefinition as LlmToolDefinition;
|
||||
use crate::tool::{Tool, ToolDefinition as WorkerToolDefinition, ToolMeta};
|
||||
|
||||
type ToolMap = HashMap<String, (ToolMeta, Arc<dyn Tool>)>;
|
||||
|
||||
/// Errors produced by ToolServer operations.
|
||||
#[derive(Debug, Error, PartialEq, Eq)]
|
||||
pub enum ToolServerError {
|
||||
/// A tool with the same name already exists.
|
||||
#[error("Tool with name '{0}' already registered")]
|
||||
DuplicateName(String),
|
||||
/// Requested tool was not found.
|
||||
#[error("Tool '{0}' not found")]
|
||||
ToolNotFound(String),
|
||||
/// Tool execution failed.
|
||||
#[error("Tool execution failed: {0}")]
|
||||
ToolExecution(String),
|
||||
}
|
||||
|
||||
/// In-memory tool server.
|
||||
#[derive(Clone, Default)]
|
||||
pub struct ToolServer {
|
||||
tools: Arc<Mutex<ToolMap>>,
|
||||
}
|
||||
|
||||
impl ToolServer {
|
||||
/// Create a new empty tool server.
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
/// Create a handle for shared access.
|
||||
pub fn handle(&self) -> ToolServerHandle {
|
||||
ToolServerHandle {
|
||||
tools: Arc::clone(&self.tools),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Shareable handle to a tool server.
|
||||
#[derive(Clone, Default)]
|
||||
pub struct ToolServerHandle {
|
||||
tools: Arc<Mutex<ToolMap>>,
|
||||
}
|
||||
|
||||
impl ToolServerHandle {
|
||||
/// Register one tool.
|
||||
pub(crate) fn register_tool(
|
||||
&self,
|
||||
factory: WorkerToolDefinition,
|
||||
) -> Result<(), ToolServerError> {
|
||||
let (meta, instance) = factory();
|
||||
let mut guard = self.tools.lock().unwrap_or_else(|e| e.into_inner());
|
||||
if guard.contains_key(&meta.name) {
|
||||
return Err(ToolServerError::DuplicateName(meta.name));
|
||||
}
|
||||
guard.insert(meta.name.clone(), (meta, instance));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Register many tools.
|
||||
pub(crate) fn register_tools(
|
||||
&self,
|
||||
factories: impl IntoIterator<Item = WorkerToolDefinition>,
|
||||
) -> Result<(), ToolServerError> {
|
||||
for factory in factories {
|
||||
self.register_tool(factory)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Get a tool by name for hook contexts.
|
||||
pub fn get_tool(&self, name: &str) -> Option<(ToolMeta, Arc<dyn Tool>)> {
|
||||
let guard = self.tools.lock().unwrap_or_else(|e| e.into_inner());
|
||||
guard.get(name).map(|(meta, tool)| (meta.clone(), Arc::clone(tool)))
|
||||
}
|
||||
|
||||
/// Execute a tool by name.
|
||||
pub async fn call_tool(&self, name: &str, input_json: &str) -> Result<String, ToolServerError> {
|
||||
let tool = {
|
||||
let guard = self.tools.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let (_, tool) = guard
|
||||
.get(name)
|
||||
.ok_or_else(|| ToolServerError::ToolNotFound(name.to_string()))?;
|
||||
Arc::clone(tool)
|
||||
};
|
||||
tool.execute(input_json)
|
||||
.await
|
||||
.map_err(|e| ToolServerError::ToolExecution(e.to_string()))
|
||||
}
|
||||
|
||||
/// Build deterministic tool definitions sorted by tool name.
|
||||
pub fn tool_definitions_sorted(&self) -> Vec<LlmToolDefinition> {
|
||||
let guard = self.tools.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let mut defs: Vec<_> = guard
|
||||
.values()
|
||||
.map(|(meta, _)| {
|
||||
LlmToolDefinition::new(&meta.name)
|
||||
.description(&meta.description)
|
||||
.input_schema(meta.input_schema.clone())
|
||||
})
|
||||
.collect();
|
||||
defs.sort_by(|a, b| a.name.cmp(&b.name));
|
||||
defs
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
use crate::tool::{Tool, ToolDefinition, ToolError, ToolMeta};
|
||||
|
||||
struct EchoTool;
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for EchoTool {
|
||||
async fn execute(&self, input_json: &str) -> Result<String, ToolError> {
|
||||
Ok(input_json.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
fn def(name: &'static str) -> ToolDefinition {
|
||||
Arc::new(move || {
|
||||
(
|
||||
ToolMeta::new(name)
|
||||
.description(format!("desc-{name}"))
|
||||
.input_schema(json!({"type":"object"})),
|
||||
Arc::new(EchoTool) as Arc<dyn Tool>,
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn register_duplicate_name_fails() {
|
||||
let handle = ToolServer::new().handle();
|
||||
handle.register_tool(def("alpha")).expect("first register");
|
||||
let err = handle
|
||||
.register_tool(def("alpha"))
|
||||
.expect_err("duplicate should fail");
|
||||
assert_eq!(err, ToolServerError::DuplicateName("alpha".to_string()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn call_tool_success_and_not_found() {
|
||||
let handle = ToolServer::new().handle();
|
||||
handle.register_tool(def("echo")).expect("register");
|
||||
|
||||
let out = handle.call_tool("echo", r#"{"x":1}"#).await.expect("call");
|
||||
assert_eq!(out, r#"{"x":1}"#);
|
||||
|
||||
let err = handle
|
||||
.call_tool("missing", "{}")
|
||||
.await
|
||||
.expect_err("missing tool");
|
||||
assert_eq!(err, ToolServerError::ToolNotFound("missing".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_definitions_are_sorted() {
|
||||
let handle = ToolServer::new().handle();
|
||||
handle.register_tool(def("zeta")).expect("register zeta");
|
||||
handle.register_tool(def("alpha")).expect("register alpha");
|
||||
handle.register_tool(def("beta")).expect("register beta");
|
||||
|
||||
let names: Vec<_> = handle
|
||||
.tool_definitions_sorted()
|
||||
.into_iter()
|
||||
.map(|d| d.name)
|
||||
.collect();
|
||||
assert_eq!(names, vec!["alpha", "beta", "zeta"]);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,4 @@
|
||||
//! Anthropic フィクスチャベースの統合テスト
|
||||
//! Anthropic fixture-based integration tests
|
||||
|
||||
mod common;
|
||||
|
||||
@@ -8,9 +8,9 @@ use std::sync::{Arc, Mutex};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use futures::Stream;
|
||||
use worker::llm_client::event::{BlockType, DeltaContent, Event};
|
||||
use worker::llm_client::{ClientError, LlmClient, Request};
|
||||
use worker::timeline::{Handler, TextBlockEvent, TextBlockKind, Timeline};
|
||||
use llm_worker::llm_client::event::{BlockType, DeltaContent, Event};
|
||||
use llm_worker::llm_client::{ClientError, LlmClient, Request};
|
||||
use llm_worker::timeline::{Handler, TextBlockEvent, TextBlockKind, Timeline};
|
||||
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
@@ -267,7 +267,7 @@ pub fn assert_timeline_integration(subdir: &str) {
|
||||
});
|
||||
|
||||
for event in &events {
|
||||
let timeline_event: worker::timeline::event::Event = event.clone().into();
|
||||
let timeline_event: llm_worker::timeline::event::Event = event.clone().into();
|
||||
timeline.dispatch(&timeline_event);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
#[test]
|
||||
fn compile_fail_state_constraints() {
|
||||
let t = trybuild::TestCases::new();
|
||||
t.compile_fail("tests/ui/cache_locked_register_tool.rs");
|
||||
t.compile_fail("tests/ui/tool_server_handle_register_tool.rs");
|
||||
}
|
||||
Vendored
Vendored
Vendored
Vendored
Vendored
Vendored
Vendored
Vendored
Vendored
@@ -1,4 +1,4 @@
|
||||
//! Gemini フィクスチャベースの統合テスト
|
||||
//! Gemini fixture-based integration tests
|
||||
|
||||
mod common;
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
//! Ollama フィクスチャベースの統合テスト
|
||||
//! Ollama fixture-based integration tests
|
||||
|
||||
mod common;
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
//! OpenAI フィクスチャベースの統合テスト
|
||||
//! OpenAI fixture-based integration tests
|
||||
|
||||
mod common;
|
||||
|
||||
+71
-72
@@ -1,16 +1,19 @@
|
||||
//! 並列ツール実行のテスト
|
||||
//! Parallel tool execution tests
|
||||
//!
|
||||
//! Workerが複数のツールを並列に実行することを確認する。
|
||||
//! Verify that Worker executes multiple tools in parallel.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use worker::Worker;
|
||||
use worker::hook::{ControlFlow, HookError, ToolCall, ToolResult, WorkerHook};
|
||||
use worker::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use worker::tool::{Tool, ToolError};
|
||||
use llm_worker::Worker;
|
||||
use llm_worker::hook::{
|
||||
Hook, HookError, PostToolCall, PostToolCallContext, PostToolCallResult, PreToolCall,
|
||||
PreToolCallResult, ToolCallContext,
|
||||
};
|
||||
use llm_worker::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use llm_worker::tool::{Tool, ToolDefinition, ToolError, ToolMeta};
|
||||
|
||||
mod common;
|
||||
use common::MockLlmClient;
|
||||
@@ -19,7 +22,7 @@ use common::MockLlmClient;
|
||||
// Parallel Execution Test Tools
|
||||
// =============================================================================
|
||||
|
||||
/// 一定時間待機してから応答するツール
|
||||
/// Tool that waits for a specified time before responding
|
||||
#[derive(Clone)]
|
||||
struct SlowTool {
|
||||
name: String,
|
||||
@@ -39,25 +42,24 @@ impl SlowTool {
|
||||
fn call_count(&self) -> usize {
|
||||
self.call_count.load(Ordering::SeqCst)
|
||||
}
|
||||
|
||||
/// Create ToolDefinition
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
let tool = self.clone();
|
||||
Arc::new(move || {
|
||||
let meta = ToolMeta::new(&tool.name)
|
||||
.description("A tool that waits before responding")
|
||||
.input_schema(serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}));
|
||||
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for SlowTool {
|
||||
fn name(&self) -> &str {
|
||||
&self.name
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
"A tool that waits before responding"
|
||||
}
|
||||
|
||||
fn input_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
})
|
||||
}
|
||||
|
||||
async fn execute(&self, _input_json: &str) -> Result<String, ToolError> {
|
||||
self.call_count.fetch_add(1, Ordering::SeqCst);
|
||||
tokio::time::sleep(Duration::from_millis(self.delay_ms)).await;
|
||||
@@ -69,13 +71,13 @@ impl Tool for SlowTool {
|
||||
// Tests
|
||||
// =============================================================================
|
||||
|
||||
/// 複数のツールが並列に実行されることを確認
|
||||
/// Verify that multiple tools are executed in parallel
|
||||
///
|
||||
/// 各ツールが100msかかる場合、逐次実行なら300ms以上かかるが、
|
||||
/// 並列実行なら100ms程度で完了するはず。
|
||||
/// If each tool takes 100ms, sequential execution would take 300ms+,
|
||||
/// but parallel execution should complete in about 100ms.
|
||||
#[tokio::test]
|
||||
async fn test_parallel_tool_execution() {
|
||||
// 3つのツール呼び出しを含むイベントシーケンス
|
||||
// Event sequence containing 3 tool calls
|
||||
let events = vec![
|
||||
Event::tool_use_start(0, "call_1", "slow_tool_1"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
@@ -94,7 +96,7 @@ async fn test_parallel_tool_execution() {
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// 各ツールは100ms待機
|
||||
// Each tool waits 100ms
|
||||
let tool1 = SlowTool::new("slow_tool_1", 100);
|
||||
let tool2 = SlowTool::new("slow_tool_2", 100);
|
||||
let tool3 = SlowTool::new("slow_tool_3", 100);
|
||||
@@ -103,21 +105,21 @@ async fn test_parallel_tool_execution() {
|
||||
let tool2_clone = tool2.clone();
|
||||
let tool3_clone = tool3.clone();
|
||||
|
||||
worker.register_tool(tool1);
|
||||
worker.register_tool(tool2);
|
||||
worker.register_tool(tool3);
|
||||
worker.register_tool(tool1.definition()).unwrap();
|
||||
worker.register_tool(tool2.definition()).unwrap();
|
||||
worker.register_tool(tool3.definition()).unwrap();
|
||||
|
||||
let start = Instant::now();
|
||||
let _result = worker.run("Run all tools").await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
// 全ツールが呼び出されたことを確認
|
||||
// Verify all tools were called
|
||||
assert_eq!(tool1_clone.call_count(), 1, "Tool 1 should be called once");
|
||||
assert_eq!(tool2_clone.call_count(), 1, "Tool 2 should be called once");
|
||||
assert_eq!(tool3_clone.call_count(), 1, "Tool 3 should be called once");
|
||||
|
||||
// 並列実行なら200ms以下で完了するはず(逐次なら300ms以上)
|
||||
// マージン込みで250msをしきい値とする
|
||||
// Parallel execution should complete in under 200ms (sequential would be 300ms+)
|
||||
// Using 250ms as threshold with margin
|
||||
assert!(
|
||||
elapsed < Duration::from_millis(250),
|
||||
"Parallel execution should complete in ~100ms, but took {:?}",
|
||||
@@ -127,7 +129,7 @@ async fn test_parallel_tool_execution() {
|
||||
println!("Parallel execution completed in {:?}", elapsed);
|
||||
}
|
||||
|
||||
/// Hook: before_tool_call でスキップされたツールは実行されないことを確認
|
||||
/// Hook: pre_tool_call - verify that skipped tools are not executed
|
||||
#[tokio::test]
|
||||
async fn test_before_tool_call_skip() {
|
||||
let events = vec![
|
||||
@@ -151,31 +153,28 @@ async fn test_before_tool_call_skip() {
|
||||
let allowed_clone = allowed_tool.clone();
|
||||
let blocked_clone = blocked_tool.clone();
|
||||
|
||||
worker.register_tool(allowed_tool);
|
||||
worker.register_tool(blocked_tool);
|
||||
worker.register_tool(allowed_tool.definition()).unwrap();
|
||||
worker.register_tool(blocked_tool.definition()).unwrap();
|
||||
|
||||
// "blocked_tool" をスキップするHook
|
||||
// Hook to skip "blocked_tool"
|
||||
struct BlockingHook;
|
||||
|
||||
#[async_trait]
|
||||
impl WorkerHook for BlockingHook {
|
||||
async fn before_tool_call(
|
||||
&self,
|
||||
tool_call: &mut ToolCall,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
if tool_call.name == "blocked_tool" {
|
||||
Ok(ControlFlow::Skip)
|
||||
impl Hook<PreToolCall> for BlockingHook {
|
||||
async fn call(&self, ctx: &mut ToolCallContext) -> Result<PreToolCallResult, HookError> {
|
||||
if ctx.call.name == "blocked_tool" {
|
||||
Ok(PreToolCallResult::Skip)
|
||||
} else {
|
||||
Ok(ControlFlow::Continue)
|
||||
Ok(PreToolCallResult::Continue)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
worker.add_hook(BlockingHook);
|
||||
worker.add_pre_tool_call_hook(BlockingHook);
|
||||
|
||||
let _result = worker.run("Test hook").await;
|
||||
|
||||
// allowed_tool は呼び出されるが、blocked_tool は呼び出されない
|
||||
// allowed_tool is called, but blocked_tool is not
|
||||
assert_eq!(
|
||||
allowed_clone.call_count(),
|
||||
1,
|
||||
@@ -188,12 +187,12 @@ async fn test_before_tool_call_skip() {
|
||||
);
|
||||
}
|
||||
|
||||
/// Hook: after_tool_call で結果が改変されることを確認
|
||||
/// Hook: post_tool_call - verify that results can be modified
|
||||
#[tokio::test]
|
||||
async fn test_after_tool_call_modification() {
|
||||
// 複数リクエストに対応するレスポンスを準備
|
||||
async fn test_post_tool_call_modification() {
|
||||
// Prepare responses for multiple requests
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
// 1回目のリクエスト: ツール呼び出し
|
||||
// First request: tool call
|
||||
vec![
|
||||
Event::tool_use_start(0, "call_1", "test_tool"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
@@ -202,7 +201,7 @@ async fn test_after_tool_call_modification() {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
// 2回目のリクエスト: ツール結果を受けてテキストレスポンス
|
||||
// Second request: text response after receiving tool result
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Done!"),
|
||||
@@ -220,41 +219,41 @@ async fn test_after_tool_call_modification() {
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for SimpleTool {
|
||||
fn name(&self) -> &str {
|
||||
"test_tool"
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"Test"
|
||||
}
|
||||
fn input_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({})
|
||||
}
|
||||
async fn execute(&self, _: &str) -> Result<String, ToolError> {
|
||||
Ok("Original Result".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
worker.register_tool(SimpleTool);
|
||||
fn simple_tool_definition() -> ToolDefinition {
|
||||
Arc::new(|| {
|
||||
let meta = ToolMeta::new("test_tool")
|
||||
.description("Test")
|
||||
.input_schema(serde_json::json!({}));
|
||||
(meta, Arc::new(SimpleTool) as Arc<dyn Tool>)
|
||||
})
|
||||
}
|
||||
|
||||
// 結果を改変するHook
|
||||
worker.register_tool(simple_tool_definition()).unwrap();
|
||||
|
||||
// Hook to modify results
|
||||
struct ModifyingHook {
|
||||
modified_content: Arc<std::sync::Mutex<Option<String>>>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl WorkerHook for ModifyingHook {
|
||||
async fn after_tool_call(
|
||||
impl Hook<PostToolCall> for ModifyingHook {
|
||||
async fn call(
|
||||
&self,
|
||||
tool_result: &mut ToolResult,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
tool_result.content = format!("[Modified] {}", tool_result.content);
|
||||
*self.modified_content.lock().unwrap() = Some(tool_result.content.clone());
|
||||
Ok(ControlFlow::Continue)
|
||||
ctx: &mut PostToolCallContext,
|
||||
) -> Result<PostToolCallResult, HookError> {
|
||||
ctx.result.content = format!("[Modified] {}", ctx.result.content);
|
||||
*self.modified_content.lock().unwrap() = Some(ctx.result.content.clone());
|
||||
Ok(PostToolCallResult::Continue)
|
||||
}
|
||||
}
|
||||
|
||||
let modified_content = Arc::new(std::sync::Mutex::new(None));
|
||||
worker.add_hook(ModifyingHook {
|
||||
worker.add_post_tool_call_hook(ModifyingHook {
|
||||
modified_content: modified_content.clone(),
|
||||
});
|
||||
|
||||
@@ -262,7 +261,7 @@ async fn test_after_tool_call_modification() {
|
||||
|
||||
assert!(result.is_ok(), "Worker should complete: {:?}", result);
|
||||
|
||||
// Hookが呼ばれて内容が改変されたことを確認
|
||||
// Verify hook was called and content was modified
|
||||
let content = modified_content.lock().unwrap().clone();
|
||||
assert!(content.is_some(), "Hook should have been called");
|
||||
assert!(
|
||||
@@ -1,26 +1,26 @@
|
||||
//! WorkerSubscriberのテスト
|
||||
//! WorkerSubscriber tests
|
||||
//!
|
||||
//! WorkerSubscriberを使ってイベントを購読するテスト
|
||||
//! Tests for subscribing to events using WorkerSubscriber
|
||||
|
||||
mod common;
|
||||
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use common::MockLlmClient;
|
||||
use worker::Worker;
|
||||
use worker::hook::ToolCall;
|
||||
use worker::llm_client::event::{Event, ResponseStatus, StatusEvent as ClientStatusEvent};
|
||||
use worker::subscriber::WorkerSubscriber;
|
||||
use worker::timeline::event::{ErrorEvent, StatusEvent, UsageEvent};
|
||||
use worker::timeline::{TextBlockEvent, ToolUseBlockEvent};
|
||||
use llm_worker::Worker;
|
||||
use llm_worker::hook::ToolCall;
|
||||
use llm_worker::llm_client::event::{Event, ResponseStatus, StatusEvent as ClientStatusEvent};
|
||||
use llm_worker::subscriber::WorkerSubscriber;
|
||||
use llm_worker::timeline::event::{ErrorEvent, StatusEvent, UsageEvent};
|
||||
use llm_worker::timeline::{TextBlockEvent, ToolUseBlockEvent};
|
||||
|
||||
// =============================================================================
|
||||
// Test Subscriber
|
||||
// =============================================================================
|
||||
|
||||
/// テスト用のシンプルなSubscriber実装
|
||||
/// Simple Subscriber implementation for testing
|
||||
struct TestSubscriber {
|
||||
// 記録用のバッファ
|
||||
// Recording buffers
|
||||
text_deltas: Arc<Mutex<Vec<String>>>,
|
||||
text_completes: Arc<Mutex<Vec<String>>>,
|
||||
tool_call_completes: Arc<Mutex<Vec<ToolCall>>>,
|
||||
@@ -60,7 +60,7 @@ impl WorkerSubscriber for TestSubscriber {
|
||||
}
|
||||
|
||||
fn on_tool_use_block(&mut self, _scope: &mut (), _event: &ToolUseBlockEvent) {
|
||||
// 必要に応じて処理
|
||||
// Process as needed
|
||||
}
|
||||
|
||||
fn on_tool_call_complete(&mut self, call: &ToolCall) {
|
||||
@@ -76,7 +76,7 @@ impl WorkerSubscriber for TestSubscriber {
|
||||
}
|
||||
|
||||
fn on_error(&mut self, _event: &ErrorEvent) {
|
||||
// 必要に応じて処理
|
||||
// Process as needed
|
||||
}
|
||||
|
||||
fn on_turn_start(&mut self, turn: usize) {
|
||||
@@ -92,10 +92,10 @@ impl WorkerSubscriber for TestSubscriber {
|
||||
// Tests
|
||||
// =============================================================================
|
||||
|
||||
/// WorkerSubscriberがテキストブロックイベントを正しく受け取ることを確認
|
||||
/// Verify that WorkerSubscriber correctly receives text block events
|
||||
#[tokio::test]
|
||||
async fn test_subscriber_text_block_events() {
|
||||
// テキストレスポンスを含むイベントシーケンス
|
||||
// Event sequence containing text response
|
||||
let events = vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Hello, "),
|
||||
@@ -109,33 +109,33 @@ async fn test_subscriber_text_block_events() {
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// Subscriberを登録
|
||||
// Register Subscriber
|
||||
let subscriber = TestSubscriber::new();
|
||||
let text_deltas = subscriber.text_deltas.clone();
|
||||
let text_completes = subscriber.text_completes.clone();
|
||||
worker.subscribe(subscriber);
|
||||
|
||||
// 実行
|
||||
// Execute
|
||||
let result = worker.run("Greet me").await;
|
||||
|
||||
assert!(result.is_ok(), "Worker should complete: {:?}", result);
|
||||
|
||||
// デルタが収集されていることを確認
|
||||
// Verify deltas were collected
|
||||
let deltas = text_deltas.lock().unwrap();
|
||||
assert_eq!(deltas.len(), 2);
|
||||
assert_eq!(deltas[0], "Hello, ");
|
||||
assert_eq!(deltas[1], "World!");
|
||||
|
||||
// 完了テキストが収集されていることを確認
|
||||
// Verify complete text was collected
|
||||
let completes = text_completes.lock().unwrap();
|
||||
assert_eq!(completes.len(), 1);
|
||||
assert_eq!(completes[0], "Hello, World!");
|
||||
}
|
||||
|
||||
/// WorkerSubscriberがツール呼び出し完了イベントを正しく受け取ることを確認
|
||||
/// Verify that WorkerSubscriber correctly receives tool call complete events
|
||||
#[tokio::test]
|
||||
async fn test_subscriber_tool_call_complete() {
|
||||
// ツール呼び出しを含むイベントシーケンス
|
||||
// Event sequence containing tool call
|
||||
let events = vec![
|
||||
Event::tool_use_start(0, "call_123", "get_weather"),
|
||||
Event::tool_input_delta(0, r#"{"city":"#),
|
||||
@@ -149,15 +149,15 @@ async fn test_subscriber_tool_call_complete() {
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// Subscriberを登録
|
||||
// Register Subscriber
|
||||
let subscriber = TestSubscriber::new();
|
||||
let tool_call_completes = subscriber.tool_call_completes.clone();
|
||||
worker.subscribe(subscriber);
|
||||
|
||||
// 実行
|
||||
// Execute
|
||||
let _ = worker.run("Weather please").await;
|
||||
|
||||
// ツール呼び出し完了が収集されていることを確認
|
||||
// Verify tool call complete was collected
|
||||
let completes = tool_call_completes.lock().unwrap();
|
||||
assert_eq!(completes.len(), 1);
|
||||
assert_eq!(completes[0].name, "get_weather");
|
||||
@@ -165,7 +165,7 @@ async fn test_subscriber_tool_call_complete() {
|
||||
assert_eq!(completes[0].input["city"], "Tokyo");
|
||||
}
|
||||
|
||||
/// WorkerSubscriberがターンイベントを正しく受け取ることを確認
|
||||
/// Verify that WorkerSubscriber correctly receives turn events
|
||||
#[tokio::test]
|
||||
async fn test_subscriber_turn_events() {
|
||||
let events = vec![
|
||||
@@ -180,29 +180,29 @@ async fn test_subscriber_turn_events() {
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// Subscriberを登録
|
||||
// Register Subscriber
|
||||
let subscriber = TestSubscriber::new();
|
||||
let turn_starts = subscriber.turn_starts.clone();
|
||||
let turn_ends = subscriber.turn_ends.clone();
|
||||
worker.subscribe(subscriber);
|
||||
|
||||
// 実行
|
||||
// Execute
|
||||
let result = worker.run("Do something").await;
|
||||
|
||||
assert!(result.is_ok());
|
||||
|
||||
// ターンイベントが収集されていることを確認
|
||||
// Verify turn events were collected
|
||||
let starts = turn_starts.lock().unwrap();
|
||||
let ends = turn_ends.lock().unwrap();
|
||||
|
||||
assert_eq!(starts.len(), 1);
|
||||
assert_eq!(starts[0], 0); // 最初のターン
|
||||
assert_eq!(starts[0], 0); // First turn
|
||||
|
||||
assert_eq!(ends.len(), 1);
|
||||
assert_eq!(ends[0], 0);
|
||||
}
|
||||
|
||||
/// WorkerSubscriberがUsageイベントを正しく受け取ることを確認
|
||||
/// Verify that WorkerSubscriber correctly receives Usage events
|
||||
#[tokio::test]
|
||||
async fn test_subscriber_usage_events() {
|
||||
let events = vec![
|
||||
@@ -218,15 +218,15 @@ async fn test_subscriber_usage_events() {
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// Subscriberを登録
|
||||
// Register Subscriber
|
||||
let subscriber = TestSubscriber::new();
|
||||
let usage_events = subscriber.usage_events.clone();
|
||||
worker.subscribe(subscriber);
|
||||
|
||||
// 実行
|
||||
// Execute
|
||||
let _ = worker.run("Hello").await;
|
||||
|
||||
// Usageイベントが収集されていることを確認
|
||||
// Verify Usage events were collected
|
||||
let usages = usage_events.lock().unwrap();
|
||||
assert_eq!(usages.len(), 1);
|
||||
assert_eq!(usages[0].input_tokens, Some(100));
|
||||
@@ -1,22 +1,21 @@
|
||||
//! ツールマクロのテスト
|
||||
//! Tool macro tests
|
||||
//!
|
||||
//! `#[tool_registry]` と `#[tool]` マクロの動作を確認する。
|
||||
//! Verify the behavior of `#[tool_registry]` and `#[tool]` macros.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
// マクロ展開に必要なインポート
|
||||
// Imports needed for macro expansion
|
||||
use schemars;
|
||||
use serde;
|
||||
|
||||
use worker::tool::Tool;
|
||||
use worker_macros::tool_registry;
|
||||
use llm_worker_macros::tool_registry;
|
||||
|
||||
// =============================================================================
|
||||
// Test: Basic Tool Generation
|
||||
// =============================================================================
|
||||
|
||||
/// シンプルなコンテキスト構造体
|
||||
/// Simple context struct
|
||||
#[derive(Clone)]
|
||||
struct SimpleContext {
|
||||
prefix: String,
|
||||
@@ -24,21 +23,21 @@ struct SimpleContext {
|
||||
|
||||
#[tool_registry]
|
||||
impl SimpleContext {
|
||||
/// メッセージに挨拶を追加する
|
||||
/// Add greeting to message
|
||||
///
|
||||
/// 指定されたメッセージにプレフィックスを付けて返します。
|
||||
/// Returns the message with a prefix added.
|
||||
#[tool]
|
||||
async fn greet(&self, message: String) -> String {
|
||||
format!("{}: {}", self.prefix, message)
|
||||
}
|
||||
|
||||
/// 二つの数を足す
|
||||
/// Add two numbers
|
||||
#[tool]
|
||||
async fn add(&self, a: i32, b: i32) -> i32 {
|
||||
a + b
|
||||
}
|
||||
|
||||
/// 引数なしのツール
|
||||
/// Tool with no arguments
|
||||
#[tool]
|
||||
async fn get_prefix(&self) -> String {
|
||||
self.prefix.clone()
|
||||
@@ -51,30 +50,31 @@ async fn test_basic_tool_generation() {
|
||||
prefix: "Hello".to_string(),
|
||||
};
|
||||
|
||||
// ファクトリメソッドでツールを取得
|
||||
let greet_tool = ctx.greet_tool();
|
||||
// Get ToolDefinition from factory method
|
||||
let greet_definition = ctx.greet_definition();
|
||||
|
||||
// 名前の確認
|
||||
assert_eq!(greet_tool.name(), "greet");
|
||||
// Call factory to get Meta and Tool
|
||||
let (meta, tool) = greet_definition();
|
||||
|
||||
// 説明の確認(docコメントから取得)
|
||||
let desc = greet_tool.description();
|
||||
// Verify meta information
|
||||
assert_eq!(meta.name, "greet");
|
||||
assert!(
|
||||
desc.contains("メッセージに挨拶を追加する"),
|
||||
meta.description.contains("Add greeting to message"),
|
||||
"Description should contain doc comment: {}",
|
||||
desc
|
||||
meta.description
|
||||
);
|
||||
|
||||
// スキーマの確認
|
||||
let schema = greet_tool.input_schema();
|
||||
println!("Schema: {}", serde_json::to_string_pretty(&schema).unwrap());
|
||||
assert!(
|
||||
schema.get("properties").is_some(),
|
||||
meta.input_schema.get("properties").is_some(),
|
||||
"Schema should have properties"
|
||||
);
|
||||
|
||||
// 実行テスト
|
||||
let result = greet_tool.execute(r#"{"message": "World"}"#).await;
|
||||
println!(
|
||||
"Schema: {}",
|
||||
serde_json::to_string_pretty(&meta.input_schema).unwrap()
|
||||
);
|
||||
|
||||
// Execution test
|
||||
let result = tool.execute(r#"{"message": "World"}"#).await;
|
||||
assert!(result.is_ok(), "Should execute successfully");
|
||||
let output = result.unwrap();
|
||||
assert!(output.contains("Hello"), "Output should contain prefix");
|
||||
@@ -87,11 +87,11 @@ async fn test_multiple_arguments() {
|
||||
prefix: "".to_string(),
|
||||
};
|
||||
|
||||
let add_tool = ctx.add_tool();
|
||||
let (meta, tool) = ctx.add_definition()();
|
||||
|
||||
assert_eq!(add_tool.name(), "add");
|
||||
assert_eq!(meta.name, "add");
|
||||
|
||||
let result = add_tool.execute(r#"{"a": 10, "b": 20}"#).await;
|
||||
let result = tool.execute(r#"{"a": 10, "b": 20}"#).await;
|
||||
assert!(result.is_ok());
|
||||
let output = result.unwrap();
|
||||
assert!(output.contains("30"), "Should contain sum: {}", output);
|
||||
@@ -103,12 +103,12 @@ async fn test_no_arguments() {
|
||||
prefix: "TestPrefix".to_string(),
|
||||
};
|
||||
|
||||
let get_prefix_tool = ctx.get_prefix_tool();
|
||||
let (meta, tool) = ctx.get_prefix_definition()();
|
||||
|
||||
assert_eq!(get_prefix_tool.name(), "get_prefix");
|
||||
assert_eq!(meta.name, "get_prefix");
|
||||
|
||||
// 空のJSONオブジェクトで呼び出し
|
||||
let result = get_prefix_tool.execute(r#"{}"#).await;
|
||||
// Call with empty JSON object
|
||||
let result = tool.execute(r#"{}"#).await;
|
||||
assert!(result.is_ok());
|
||||
let output = result.unwrap();
|
||||
assert!(
|
||||
@@ -124,10 +124,10 @@ async fn test_invalid_arguments() {
|
||||
prefix: "".to_string(),
|
||||
};
|
||||
|
||||
let greet_tool = ctx.greet_tool();
|
||||
let (_, tool) = ctx.greet_definition()();
|
||||
|
||||
// 不正なJSON
|
||||
let result = greet_tool.execute(r#"{"wrong_field": "value"}"#).await;
|
||||
// Invalid JSON
|
||||
let result = tool.execute(r#"{"wrong_field": "value"}"#).await;
|
||||
assert!(result.is_err(), "Should fail with invalid arguments");
|
||||
}
|
||||
|
||||
@@ -149,7 +149,7 @@ impl std::fmt::Display for MyError {
|
||||
|
||||
#[tool_registry]
|
||||
impl FallibleContext {
|
||||
/// 与えられた値を検証する
|
||||
/// Validate the given value
|
||||
#[tool]
|
||||
async fn validate(&self, value: i32) -> Result<String, MyError> {
|
||||
if value > 0 {
|
||||
@@ -163,9 +163,9 @@ impl FallibleContext {
|
||||
#[tokio::test]
|
||||
async fn test_result_return_type_success() {
|
||||
let ctx = FallibleContext;
|
||||
let validate_tool = ctx.validate_tool();
|
||||
let (_, tool) = ctx.validate_definition()();
|
||||
|
||||
let result = validate_tool.execute(r#"{"value": 42}"#).await;
|
||||
let result = tool.execute(r#"{"value": 42}"#).await;
|
||||
assert!(result.is_ok(), "Should succeed for positive value");
|
||||
let output = result.unwrap();
|
||||
assert!(output.contains("Valid"), "Should contain Valid: {}", output);
|
||||
@@ -174,9 +174,9 @@ async fn test_result_return_type_success() {
|
||||
#[tokio::test]
|
||||
async fn test_result_return_type_error() {
|
||||
let ctx = FallibleContext;
|
||||
let validate_tool = ctx.validate_tool();
|
||||
let (_, tool) = ctx.validate_definition()();
|
||||
|
||||
let result = validate_tool.execute(r#"{"value": -1}"#).await;
|
||||
let result = tool.execute(r#"{"value": -1}"#).await;
|
||||
assert!(result.is_err(), "Should fail for negative value");
|
||||
|
||||
let err = result.unwrap_err();
|
||||
@@ -198,7 +198,7 @@ struct SyncContext {
|
||||
|
||||
#[tool_registry]
|
||||
impl SyncContext {
|
||||
/// カウンターをインクリメントして返す (非async)
|
||||
/// Increment counter and return (non-async)
|
||||
#[tool]
|
||||
fn increment(&self) -> usize {
|
||||
self.counter.fetch_add(1, Ordering::SeqCst) + 1
|
||||
@@ -211,17 +211,36 @@ async fn test_sync_method() {
|
||||
counter: Arc::new(AtomicUsize::new(0)),
|
||||
};
|
||||
|
||||
let increment_tool = ctx.increment_tool();
|
||||
let (_, tool) = ctx.increment_definition()();
|
||||
|
||||
// 3回実行
|
||||
let result1 = increment_tool.execute(r#"{}"#).await;
|
||||
let result2 = increment_tool.execute(r#"{}"#).await;
|
||||
let result3 = increment_tool.execute(r#"{}"#).await;
|
||||
// Execute 3 times
|
||||
let result1 = tool.execute(r#"{}"#).await;
|
||||
let result2 = tool.execute(r#"{}"#).await;
|
||||
let result3 = tool.execute(r#"{}"#).await;
|
||||
|
||||
assert!(result1.is_ok());
|
||||
assert!(result2.is_ok());
|
||||
assert!(result3.is_ok());
|
||||
|
||||
// カウンターは3になっているはず
|
||||
// Counter should be 3
|
||||
assert_eq!(ctx.counter.load(Ordering::SeqCst), 3);
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Test: ToolMeta Immutability
|
||||
// =============================================================================
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tool_meta_immutability() {
|
||||
let ctx = SimpleContext {
|
||||
prefix: "Test".to_string(),
|
||||
};
|
||||
|
||||
// Verify same meta info is returned on multiple calls
|
||||
let (meta1, _) = ctx.greet_definition()();
|
||||
let (meta2, _) = ctx.greet_definition()();
|
||||
|
||||
assert_eq!(meta1.name, meta2.name);
|
||||
assert_eq!(meta1.description, meta2.description);
|
||||
assert_eq!(meta1.input_schema, meta2.input_schema);
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
use llm_worker::Worker;
|
||||
use llm_worker::llm_client::providers::ollama::OllamaClient;
|
||||
use std::sync::Arc;
|
||||
|
||||
fn main() {
|
||||
let client = OllamaClient::new("dummy-model");
|
||||
let worker = Worker::new(client);
|
||||
let mut locked = worker.lock();
|
||||
let def: llm_worker::tool::ToolDefinition = Arc::new(|| panic!("unused"));
|
||||
let _ = locked.register_tool(def);
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
error[E0599]: no method named `register_tool` found for struct `Worker<OllamaClient, CacheLocked>` in the current scope
|
||||
--> tests/ui/cache_locked_register_tool.rs:10:20
|
||||
|
|
||||
10 | let _ = locked.register_tool(def);
|
||||
| ^^^^^^^^^^^^^ method not found in `Worker<OllamaClient, CacheLocked>`
|
||||
|
|
||||
= note: the method was found for
|
||||
- `Worker<C>`
|
||||
@@ -0,0 +1,11 @@
|
||||
use llm_worker::Worker;
|
||||
use llm_worker::llm_client::providers::ollama::OllamaClient;
|
||||
use std::sync::Arc;
|
||||
|
||||
fn main() {
|
||||
let client = OllamaClient::new("dummy-model");
|
||||
let worker = Worker::new(client);
|
||||
let handle = worker.tool_server_handle();
|
||||
let def: llm_worker::tool::ToolDefinition = Arc::new(|| panic!("unused"));
|
||||
let _ = handle.register_tool(def);
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
error[E0624]: method `register_tool` is private
|
||||
--> tests/ui/tool_server_handle_register_tool.rs:10:20
|
||||
|
|
||||
10 | let _ = handle.register_tool(def);
|
||||
| ^^^^^^^^^^^^^ private method
|
||||
|
|
||||
::: src/tool_server.rs
|
||||
|
|
||||
| / pub(crate) fn register_tool(
|
||||
| | &self,
|
||||
| | factory: WorkerToolDefinition,
|
||||
| | ) -> Result<(), ToolServerError> {
|
||||
| |____________________________________- private method defined here
|
||||
@@ -0,0 +1,39 @@
|
||||
use llm_worker::llm_client::providers::openai::OpenAIClient;
|
||||
use llm_worker::{Worker, WorkerError};
|
||||
|
||||
#[test]
|
||||
fn test_openai_top_k_warning() {
|
||||
// Create client with dummy key (validate_config doesn't make network calls, so safe)
|
||||
let client = OpenAIClient::new("dummy-key", "gpt-4o");
|
||||
|
||||
// Create Worker with top_k set (OpenAI doesn't support top_k)
|
||||
let worker = Worker::new(client).top_k(50);
|
||||
|
||||
// Run validate()
|
||||
let result = worker.validate();
|
||||
|
||||
// Verify error is returned and ConfigWarnings is included
|
||||
match result {
|
||||
Err(WorkerError::ConfigWarnings(warnings)) => {
|
||||
assert_eq!(warnings.len(), 1);
|
||||
assert_eq!(warnings[0].option_name, "top_k");
|
||||
println!("Got expected warning: {}", warnings[0]);
|
||||
}
|
||||
Ok(_) => panic!("Should have returned validation error"),
|
||||
Err(e) => panic!("Unexpected error type: {:?}", e),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_openai_valid_config() {
|
||||
let client = OpenAIClient::new("dummy-key", "gpt-4o");
|
||||
|
||||
// Valid configuration (temperature only)
|
||||
let worker = Worker::new(client).temperature(0.7);
|
||||
|
||||
// Run validate()
|
||||
let result = worker.validate();
|
||||
|
||||
// Verify success
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
@@ -1,7 +1,7 @@
|
||||
//! Workerフィクスチャベースの統合テスト
|
||||
//! Worker fixture-based integration tests
|
||||
//!
|
||||
//! 記録されたAPIレスポンスを使ってWorkerの動作をテストする。
|
||||
//! APIキー不要でローカルで実行可能。
|
||||
//! Tests Worker behavior using recorded API responses.
|
||||
//! Can run locally without API keys.
|
||||
|
||||
mod common;
|
||||
|
||||
@@ -11,15 +11,15 @@ use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use common::MockLlmClient;
|
||||
use worker::Worker;
|
||||
use worker::tool::{Tool, ToolError};
|
||||
use llm_worker::Worker;
|
||||
use llm_worker::tool::{Tool, ToolDefinition, ToolError, ToolMeta};
|
||||
|
||||
/// フィクスチャディレクトリのパス
|
||||
/// Fixture directory path
|
||||
fn fixtures_dir() -> std::path::PathBuf {
|
||||
Path::new(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/anthropic")
|
||||
}
|
||||
|
||||
/// シンプルなテスト用ツール
|
||||
/// Simple test tool
|
||||
#[derive(Clone)]
|
||||
struct MockWeatherTool {
|
||||
call_count: Arc<AtomicUsize>,
|
||||
@@ -35,20 +35,13 @@ impl MockWeatherTool {
|
||||
fn get_call_count(&self) -> usize {
|
||||
self.call_count.load(Ordering::SeqCst)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for MockWeatherTool {
|
||||
fn name(&self) -> &str {
|
||||
"get_weather"
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
"Get the current weather for a city"
|
||||
}
|
||||
|
||||
fn input_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
let tool = self.clone();
|
||||
Arc::new(move || {
|
||||
let meta = ToolMeta::new("get_weather")
|
||||
.description("Get the current weather for a city")
|
||||
.input_schema(serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"city": {
|
||||
@@ -57,19 +50,24 @@ impl Tool for MockWeatherTool {
|
||||
}
|
||||
},
|
||||
"required": ["city"]
|
||||
}));
|
||||
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for MockWeatherTool {
|
||||
async fn execute(&self, input_json: &str) -> Result<String, ToolError> {
|
||||
self.call_count.fetch_add(1, Ordering::SeqCst);
|
||||
|
||||
// 入力をパース
|
||||
// Parse input
|
||||
let input: serde_json::Value = serde_json::from_str(input_json)
|
||||
.map_err(|e| ToolError::InvalidArgument(e.to_string()))?;
|
||||
|
||||
let city = input["city"].as_str().unwrap_or("Unknown");
|
||||
|
||||
// モックのレスポンスを返す
|
||||
// Return mock response
|
||||
Ok(format!("Weather in {}: Sunny, 22°C", city))
|
||||
}
|
||||
}
|
||||
@@ -78,12 +76,12 @@ impl Tool for MockWeatherTool {
|
||||
// Basic Fixture Tests
|
||||
// =============================================================================
|
||||
|
||||
/// MockLlmClientがJSONLフィクスチャファイルから正しくイベントをロードできることを確認
|
||||
/// Verify that MockLlmClient can correctly load events from JSONL fixture files
|
||||
///
|
||||
/// 既存のanthropic_*.jsonlファイルを使用し、イベントがパース・ロードされることを検証する。
|
||||
/// Uses existing anthropic_*.jsonl files to verify events are parsed and loaded.
|
||||
#[test]
|
||||
fn test_mock_client_from_fixture() {
|
||||
// 既存のフィクスチャをロード
|
||||
// Load existing fixture
|
||||
let fixture_path = fixtures_dir().join("anthropic_1767624445.jsonl");
|
||||
if !fixture_path.exists() {
|
||||
println!("Fixture not found, skipping test");
|
||||
@@ -95,14 +93,14 @@ fn test_mock_client_from_fixture() {
|
||||
println!("Loaded {} events from fixture", client.event_count());
|
||||
}
|
||||
|
||||
/// MockLlmClientが直接指定されたイベントリストで正しく動作することを確認
|
||||
/// Verify that MockLlmClient works correctly with directly specified event lists
|
||||
///
|
||||
/// fixtureファイルを使わず、プログラムでイベントを構築してクライアントを作成する。
|
||||
/// Creates a client with programmatically constructed events instead of using fixture files.
|
||||
#[test]
|
||||
fn test_mock_client_from_events() {
|
||||
use worker::llm_client::event::Event;
|
||||
use llm_worker::llm_client::event::Event;
|
||||
|
||||
// 直接イベントを指定
|
||||
// Specify events directly
|
||||
let events = vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Hello!"),
|
||||
@@ -117,10 +115,10 @@ fn test_mock_client_from_events() {
|
||||
// Worker Tests with Fixtures
|
||||
// =============================================================================
|
||||
|
||||
/// Workerがシンプルなテキストレスポンスを正しく処理できることを確認
|
||||
/// Verify that Worker can correctly process simple text responses
|
||||
///
|
||||
/// simple_text.jsonlフィクスチャを使用し、ツール呼び出しなしのシナリオをテストする。
|
||||
/// フィクスチャがない場合はスキップされる。
|
||||
/// Uses simple_text.jsonl fixture to test scenarios without tool calls.
|
||||
/// Skipped if fixture is not present.
|
||||
#[tokio::test]
|
||||
async fn test_worker_simple_text_response() {
|
||||
let fixture_path = fixtures_dir().join("simple_text.jsonl");
|
||||
@@ -133,16 +131,16 @@ async fn test_worker_simple_text_response() {
|
||||
let client = MockLlmClient::from_fixture(&fixture_path).unwrap();
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// シンプルなメッセージを送信
|
||||
// Send a simple message
|
||||
let result = worker.run("Hello").await;
|
||||
|
||||
assert!(result.is_ok(), "Worker should complete successfully");
|
||||
}
|
||||
|
||||
/// Workerがツール呼び出しを含むレスポンスを正しく処理できることを確認
|
||||
/// Verify that Worker can correctly process responses containing tool calls
|
||||
///
|
||||
/// tool_call.jsonlフィクスチャを使用し、MockWeatherToolが呼び出されることをテストする。
|
||||
/// max_turns=1に設定し、ツール実行後のループを防止。
|
||||
/// Uses tool_call.jsonl fixture to test that MockWeatherTool is called.
|
||||
/// Sets max_turns=1 to prevent loop after tool execution.
|
||||
#[tokio::test]
|
||||
async fn test_worker_tool_call() {
|
||||
let fixture_path = fixtures_dir().join("tool_call.jsonl");
|
||||
@@ -155,32 +153,32 @@ async fn test_worker_tool_call() {
|
||||
let client = MockLlmClient::from_fixture(&fixture_path).unwrap();
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// ツールを登録
|
||||
// Register tool
|
||||
let weather_tool = MockWeatherTool::new();
|
||||
let tool_for_check = weather_tool.clone();
|
||||
worker.register_tool(weather_tool);
|
||||
worker.register_tool(weather_tool.definition()).unwrap();
|
||||
|
||||
// メッセージを送信
|
||||
// Send message
|
||||
let _result = worker.run("What's the weather in Tokyo?").await;
|
||||
|
||||
// ツールが呼び出されたことを確認
|
||||
// Note: max_turns=1なのでツール結果後のリクエストは送信されない
|
||||
// Verify tool was called
|
||||
// Note: max_turns=1 so no request is sent after tool result
|
||||
let call_count = tool_for_check.get_call_count();
|
||||
println!("Tool was called {} times", call_count);
|
||||
|
||||
// フィクスチャにToolUseが含まれていればツールが呼び出されるはず
|
||||
// ただしmax_turns=1なので1回で終了
|
||||
// Tool should be called if fixture contains ToolUse
|
||||
// But ends after 1 turn due to max_turns=1
|
||||
}
|
||||
|
||||
/// fixtureファイルなしでWorkerが動作することを確認
|
||||
/// Verify that Worker works without fixture files
|
||||
///
|
||||
/// プログラムでイベントシーケンスを構築し、MockLlmClientに渡してテストする。
|
||||
/// テストの独立性を高め、外部ファイルへの依存を排除したい場合に有用。
|
||||
/// Constructs event sequence programmatically and passes to MockLlmClient.
|
||||
/// Useful when test independence is needed and external file dependency should be eliminated.
|
||||
#[tokio::test]
|
||||
async fn test_worker_with_programmatic_events() {
|
||||
use worker::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use llm_worker::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
|
||||
// プログラムでイベントシーケンスを構築
|
||||
// Construct event sequence programmatically
|
||||
let events = vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Hello, "),
|
||||
@@ -199,16 +197,16 @@ async fn test_worker_with_programmatic_events() {
|
||||
assert!(result.is_ok(), "Worker should complete successfully");
|
||||
}
|
||||
|
||||
/// ToolCallCollectorがToolUseブロックイベントから正しくToolCallを収集することを確認
|
||||
/// Verify that ToolCallCollector correctly collects ToolCall from ToolUse block events
|
||||
///
|
||||
/// Timelineにイベントをディスパッチし、ToolCallCollectorが
|
||||
/// id, name, input(JSON)を正しく抽出できることを検証する。
|
||||
/// Dispatches events to Timeline and verifies ToolCallCollector
|
||||
/// correctly extracts id, name, and input (JSON).
|
||||
#[tokio::test]
|
||||
async fn test_tool_call_collector_integration() {
|
||||
use worker::llm_client::event::Event;
|
||||
use worker::timeline::{Timeline, ToolCallCollector};
|
||||
use llm_worker::llm_client::event::Event;
|
||||
use llm_worker::timeline::{Timeline, ToolCallCollector};
|
||||
|
||||
// ToolUseブロックを含むイベントシーケンス
|
||||
// Event sequence containing ToolUse block
|
||||
let events = vec![
|
||||
Event::tool_use_start(0, "call_123", "get_weather"),
|
||||
Event::tool_input_delta(0, r#"{"city":"#),
|
||||
@@ -220,13 +218,13 @@ async fn test_tool_call_collector_integration() {
|
||||
let mut timeline = Timeline::new();
|
||||
timeline.on_tool_use_block(collector.clone());
|
||||
|
||||
// イベントをディスパッチ
|
||||
// Dispatch events
|
||||
for event in &events {
|
||||
let timeline_event: worker::timeline::event::Event = event.clone().into();
|
||||
let timeline_event: llm_worker::timeline::event::Event = event.clone().into();
|
||||
timeline.dispatch(&timeline_event);
|
||||
}
|
||||
|
||||
// 収集されたToolCallを確認
|
||||
// Verify collected ToolCall
|
||||
let calls = collector.take_collected();
|
||||
assert_eq!(calls.len(), 1, "Should collect one tool call");
|
||||
assert_eq!(calls[0].name, "get_weather");
|
||||
@@ -1,20 +1,25 @@
|
||||
//! Worker状態管理のテスト
|
||||
//! Worker state management tests
|
||||
//!
|
||||
//! Type-stateパターン(Mutable/Locked)による状態遷移と
|
||||
//! ターン間の状態保持をテストする。
|
||||
//! Tests for state transitions using the Type-state pattern (Mutable/CacheLocked)
|
||||
//! and state preservation between turns.
|
||||
|
||||
mod common;
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use common::MockLlmClient;
|
||||
use worker::Worker;
|
||||
use worker::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use worker::{Message, MessageContent};
|
||||
use llm_worker::Worker;
|
||||
use llm_worker::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use llm_worker::tool::{Tool, ToolDefinition, ToolError, ToolMeta};
|
||||
use llm_worker::Item;
|
||||
|
||||
// =============================================================================
|
||||
// Mutable状態のテスト
|
||||
// Mutable State Tests
|
||||
// =============================================================================
|
||||
|
||||
/// Mutable状態でシステムプロンプトを設定できることを確認
|
||||
/// Verify that system prompt can be set in Mutable state
|
||||
#[test]
|
||||
fn test_mutable_set_system_prompt() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
@@ -29,114 +34,164 @@ fn test_mutable_set_system_prompt() {
|
||||
);
|
||||
}
|
||||
|
||||
/// Mutable状態で履歴を自由に編集できることを確認
|
||||
/// Verify that history can be freely edited in Mutable state
|
||||
#[test]
|
||||
fn test_mutable_history_manipulation() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// 初期状態は空
|
||||
// Initial state is empty
|
||||
assert!(worker.history().is_empty());
|
||||
|
||||
// 履歴を追加
|
||||
worker.push_message(Message::user("Hello"));
|
||||
worker.push_message(Message::assistant("Hi there!"));
|
||||
// Add to history
|
||||
worker.push_item(Item::user_message("Hello"));
|
||||
worker.push_item(Item::assistant_message("Hi there!"));
|
||||
assert_eq!(worker.history().len(), 2);
|
||||
|
||||
// 履歴への可変アクセス
|
||||
worker.history_mut().push(Message::user("How are you?"));
|
||||
// Mutable access to history
|
||||
worker.history_mut().push(Item::user_message("How are you?"));
|
||||
assert_eq!(worker.history().len(), 3);
|
||||
|
||||
// 履歴をクリア
|
||||
// Clear history
|
||||
worker.clear_history();
|
||||
assert!(worker.history().is_empty());
|
||||
|
||||
// 履歴を設定
|
||||
let messages = vec![Message::user("Test"), Message::assistant("Response")];
|
||||
worker.set_history(messages);
|
||||
// Set history
|
||||
let items = vec![Item::user_message("Test"), Item::assistant_message("Response")];
|
||||
worker.set_history(items);
|
||||
assert_eq!(worker.history().len(), 2);
|
||||
}
|
||||
|
||||
/// ビルダーパターンでWorkerを構築できることを確認
|
||||
/// Verify that Worker can be constructed using builder pattern
|
||||
#[test]
|
||||
fn test_mutable_builder_pattern() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let worker = Worker::new(client)
|
||||
.system_prompt("System prompt")
|
||||
.with_message(Message::user("Hello"))
|
||||
.with_message(Message::assistant("Hi!"))
|
||||
.with_messages(vec![
|
||||
Message::user("How are you?"),
|
||||
Message::assistant("I'm fine!"),
|
||||
.with_item(Item::user_message("Hello"))
|
||||
.with_item(Item::assistant_message("Hi!"))
|
||||
.with_items(vec![
|
||||
Item::user_message("How are you?"),
|
||||
Item::assistant_message("I'm fine!"),
|
||||
]);
|
||||
|
||||
assert_eq!(worker.get_system_prompt(), Some("System prompt"));
|
||||
assert_eq!(worker.history().len(), 4);
|
||||
}
|
||||
|
||||
/// extend_historyで複数メッセージを追加できることを確認
|
||||
/// Verify that multiple items can be added with extend_history
|
||||
#[test]
|
||||
fn test_mutable_extend_history() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
worker.push_message(Message::user("First"));
|
||||
worker.push_item(Item::user_message("First"));
|
||||
|
||||
worker.extend_history(vec![
|
||||
Message::assistant("Response 1"),
|
||||
Message::user("Second"),
|
||||
Message::assistant("Response 2"),
|
||||
Item::assistant_message("Response 1"),
|
||||
Item::user_message("Second"),
|
||||
Item::assistant_message("Response 2"),
|
||||
]);
|
||||
|
||||
assert_eq!(worker.history().len(), 4);
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct CountingTool {
|
||||
name: String,
|
||||
calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
impl CountingTool {
|
||||
fn new(name: impl Into<String>) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
calls: Arc::new(AtomicUsize::new(0)),
|
||||
}
|
||||
}
|
||||
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
let tool = self.clone();
|
||||
Arc::new(move || {
|
||||
(
|
||||
ToolMeta::new(&tool.name)
|
||||
.description("Counting tool")
|
||||
.input_schema(serde_json::json!({"type":"object","properties":{}})),
|
||||
Arc::new(tool.clone()) as Arc<dyn Tool>,
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn call_count(&self) -> usize {
|
||||
self.calls.load(Ordering::SeqCst)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for CountingTool {
|
||||
async fn execute(&self, _input_json: &str) -> Result<String, ToolError> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(format!("{}-ok", self.name))
|
||||
}
|
||||
}
|
||||
|
||||
/// Verify that tools can be registered in Mutable state.
|
||||
#[test]
|
||||
fn test_mutable_can_register_tool() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
let tool = CountingTool::new("count_tool");
|
||||
|
||||
let result = worker.register_tool(tool.definition());
|
||||
assert!(result.is_ok(), "Mutable should allow tool registration");
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// 状態遷移テスト
|
||||
// State Transition Tests
|
||||
// =============================================================================
|
||||
|
||||
/// lock()でMutable -> Locked状態に遷移することを確認
|
||||
/// Verify that lock() transitions from Mutable -> CacheLocked state
|
||||
#[test]
|
||||
fn test_lock_transition() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
worker.set_system_prompt("System");
|
||||
worker.push_message(Message::user("Hello"));
|
||||
worker.push_message(Message::assistant("Hi"));
|
||||
worker.push_item(Item::user_message("Hello"));
|
||||
worker.push_item(Item::assistant_message("Hi"));
|
||||
|
||||
// ロック
|
||||
// Lock
|
||||
let locked_worker = worker.lock();
|
||||
|
||||
// Locked状態でも履歴とシステムプロンプトにアクセス可能
|
||||
// History and system prompt are still accessible in CacheLocked state
|
||||
assert_eq!(locked_worker.get_system_prompt(), Some("System"));
|
||||
assert_eq!(locked_worker.history().len(), 2);
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 2);
|
||||
}
|
||||
|
||||
/// unlock()でLocked -> Mutable状態に遷移することを確認
|
||||
/// Verify that unlock() transitions from CacheLocked -> Mutable state
|
||||
#[test]
|
||||
fn test_unlock_transition() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
worker.push_message(Message::user("Hello"));
|
||||
worker.push_item(Item::user_message("Hello"));
|
||||
let locked_worker = worker.lock();
|
||||
|
||||
// アンロック
|
||||
// Unlock
|
||||
let mut worker = locked_worker.unlock();
|
||||
|
||||
// Mutable状態に戻ったので履歴操作が可能
|
||||
worker.push_message(Message::assistant("Hi"));
|
||||
// History operations are available again in Mutable state
|
||||
worker.push_item(Item::assistant_message("Hi"));
|
||||
worker.clear_history();
|
||||
assert!(worker.history().is_empty());
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ターン実行と状態保持のテスト
|
||||
// Turn Execution and State Preservation Tests
|
||||
// =============================================================================
|
||||
|
||||
/// Mutable状態でターンを実行し、履歴が正しく更新されることを確認
|
||||
/// Verify that history is correctly updated after running a turn in Mutable state
|
||||
#[tokio::test]
|
||||
async fn test_mutable_run_updates_history() {
|
||||
let events = vec![
|
||||
@@ -151,33 +206,27 @@ async fn test_mutable_run_updates_history() {
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// 実行
|
||||
// Execute
|
||||
let result = worker.run("Hi there").await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
// 履歴が更新されている
|
||||
// History is updated
|
||||
let history = worker.history();
|
||||
assert_eq!(history.len(), 2); // user + assistant
|
||||
|
||||
// ユーザーメッセージ
|
||||
assert!(matches!(
|
||||
&history[0].content,
|
||||
MessageContent::Text(t) if t == "Hi there"
|
||||
));
|
||||
// User message
|
||||
assert_eq!(history[0].as_text(), Some("Hi there"));
|
||||
|
||||
// アシスタントメッセージ
|
||||
assert!(matches!(
|
||||
&history[1].content,
|
||||
MessageContent::Text(t) if t == "Hello, I'm an assistant!"
|
||||
));
|
||||
// Assistant message
|
||||
assert_eq!(history[1].as_text(), Some("Hello, I'm an assistant!"));
|
||||
}
|
||||
|
||||
/// Locked状態で複数ターンを実行し、履歴が正しく累積することを確認
|
||||
/// Verify that history accumulates correctly over multiple turns in CacheLocked state
|
||||
#[tokio::test]
|
||||
async fn test_locked_multi_turn_history_accumulation() {
|
||||
// 2回のリクエストに対応するレスポンスを準備
|
||||
// Prepare responses for 2 requests
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
// 1回目のレスポンス
|
||||
// First response
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Nice to meet you!"),
|
||||
@@ -186,7 +235,7 @@ async fn test_locked_multi_turn_history_accumulation() {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
// 2回目のレスポンス
|
||||
// Second response
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "I can help with that."),
|
||||
@@ -199,37 +248,37 @@ async fn test_locked_multi_turn_history_accumulation() {
|
||||
|
||||
let worker = Worker::new(client).system_prompt("You are helpful.");
|
||||
|
||||
// ロック(システムプロンプト設定後)
|
||||
// Lock (after setting system prompt)
|
||||
let mut locked_worker = worker.lock();
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 0); // メッセージはまだない
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 0); // No items yet
|
||||
|
||||
// 1ターン目
|
||||
// Turn 1
|
||||
let result1 = locked_worker.run("Hello!").await;
|
||||
assert!(result1.is_ok());
|
||||
assert_eq!(locked_worker.history().len(), 2); // user + assistant
|
||||
|
||||
// 2ターン目
|
||||
// Turn 2
|
||||
let result2 = locked_worker.run("Can you help me?").await;
|
||||
assert!(result2.is_ok());
|
||||
assert_eq!(locked_worker.history().len(), 4); // 2 * (user + assistant)
|
||||
|
||||
// 履歴の内容を確認
|
||||
// Verify history contents
|
||||
let history = locked_worker.history();
|
||||
|
||||
// 1ターン目のユーザーメッセージ
|
||||
assert!(matches!(&history[0].content, MessageContent::Text(t) if t == "Hello!"));
|
||||
// Turn 1 user message
|
||||
assert_eq!(history[0].as_text(), Some("Hello!"));
|
||||
|
||||
// 1ターン目のアシスタントメッセージ
|
||||
assert!(matches!(&history[1].content, MessageContent::Text(t) if t == "Nice to meet you!"));
|
||||
// Turn 1 assistant message
|
||||
assert_eq!(history[1].as_text(), Some("Nice to meet you!"));
|
||||
|
||||
// 2ターン目のユーザーメッセージ
|
||||
assert!(matches!(&history[2].content, MessageContent::Text(t) if t == "Can you help me?"));
|
||||
// Turn 2 user message
|
||||
assert_eq!(history[2].as_text(), Some("Can you help me?"));
|
||||
|
||||
// 2ターン目のアシスタントメッセージ
|
||||
assert!(matches!(&history[3].content, MessageContent::Text(t) if t == "I can help with that."));
|
||||
// Turn 2 assistant message
|
||||
assert_eq!(history[3].as_text(), Some("I can help with that."));
|
||||
}
|
||||
|
||||
/// locked_prefix_lenがロック時点の履歴長を正しく記録することを確認
|
||||
/// Verify that locked_prefix_len correctly records history length at lock time
|
||||
#[tokio::test]
|
||||
async fn test_locked_prefix_len_tracking() {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
@@ -253,25 +302,25 @@ async fn test_locked_prefix_len_tracking() {
|
||||
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// 事前にメッセージを追加
|
||||
worker.push_message(Message::user("Pre-existing message 1"));
|
||||
worker.push_message(Message::assistant("Pre-existing response 1"));
|
||||
// Add items beforehand
|
||||
worker.push_item(Item::user_message("Pre-existing message 1"));
|
||||
worker.push_item(Item::assistant_message("Pre-existing response 1"));
|
||||
|
||||
assert_eq!(worker.history().len(), 2);
|
||||
|
||||
// ロック
|
||||
// Lock
|
||||
let mut locked_worker = worker.lock();
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 2); // ロック時点で2メッセージ
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 2); // 2 items at lock time
|
||||
|
||||
// ターン実行
|
||||
// Execute turn
|
||||
locked_worker.run("New message").await.unwrap();
|
||||
|
||||
// 履歴は増えるが、locked_prefix_lenは変わらない
|
||||
// History grows but locked_prefix_len remains unchanged
|
||||
assert_eq!(locked_worker.history().len(), 4); // 2 + 2
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 2); // 変わらない
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 2); // Unchanged
|
||||
}
|
||||
|
||||
/// ターンカウントが正しくインクリメントされることを確認
|
||||
/// Verify that turn count is correctly incremented
|
||||
#[tokio::test]
|
||||
async fn test_turn_count_increment() {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
@@ -304,7 +353,7 @@ async fn test_turn_count_increment() {
|
||||
assert_eq!(worker.turn_count(), 2);
|
||||
}
|
||||
|
||||
/// unlock後に履歴を編集し、再度lockできることを確認
|
||||
/// Verify that history can be edited after unlock and re-locked
|
||||
#[tokio::test]
|
||||
async fn test_unlock_edit_relock() {
|
||||
let client = MockLlmClient::with_responses(vec![vec![
|
||||
@@ -317,30 +366,91 @@ async fn test_unlock_edit_relock() {
|
||||
]]);
|
||||
|
||||
let worker = Worker::new(client)
|
||||
.with_message(Message::user("Hello"))
|
||||
.with_message(Message::assistant("Hi"));
|
||||
.with_item(Item::user_message("Hello"))
|
||||
.with_item(Item::assistant_message("Hi"));
|
||||
|
||||
// ロック -> アンロック
|
||||
// Lock -> Unlock
|
||||
let locked = worker.lock();
|
||||
assert_eq!(locked.locked_prefix_len(), 2);
|
||||
|
||||
let mut unlocked = locked.unlock();
|
||||
|
||||
// 履歴を編集
|
||||
// Edit history
|
||||
unlocked.clear_history();
|
||||
unlocked.push_message(Message::user("Fresh start"));
|
||||
unlocked.push_item(Item::user_message("Fresh start"));
|
||||
|
||||
// 再ロック
|
||||
// Re-lock
|
||||
let relocked = unlocked.lock();
|
||||
assert_eq!(relocked.history().len(), 1);
|
||||
assert_eq!(relocked.locked_prefix_len(), 1);
|
||||
}
|
||||
|
||||
/// Verify that tools registered before lock and after unlock remain effective.
|
||||
#[tokio::test]
|
||||
async fn test_lock_unlock_relock_tools_remain_effective() {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
vec![
|
||||
Event::tool_use_start(0, "call_1", "tool_a"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "done-a"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
vec![
|
||||
Event::tool_use_start(0, "call_2", "tool_b"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "done-b"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
]);
|
||||
|
||||
let mut worker = Worker::new(client);
|
||||
let tool_a = CountingTool::new("tool_a");
|
||||
worker
|
||||
.register_tool(tool_a.definition())
|
||||
.expect("register tool_a should succeed");
|
||||
|
||||
let mut locked = worker.lock();
|
||||
locked.run("first").await.expect("first run");
|
||||
assert_eq!(tool_a.call_count(), 1, "tool_a should be called once");
|
||||
|
||||
let mut unlocked = locked.unlock();
|
||||
let tool_b = CountingTool::new("tool_b");
|
||||
unlocked
|
||||
.register_tool(tool_b.definition())
|
||||
.expect("register tool_b after unlock should succeed");
|
||||
|
||||
let mut relocked = unlocked.lock();
|
||||
relocked.run("second").await.expect("second run");
|
||||
|
||||
assert_eq!(tool_a.call_count(), 1, "tool_a should not be called again");
|
||||
assert_eq!(tool_b.call_count(), 1, "tool_b should be called once");
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// システムプロンプト保持のテスト
|
||||
// System Prompt Preservation Tests
|
||||
// =============================================================================
|
||||
|
||||
/// Locked状態でもシステムプロンプトが保持されることを確認
|
||||
/// Verify that system prompt is preserved in CacheLocked state
|
||||
#[test]
|
||||
fn test_system_prompt_preserved_in_locked_state() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
@@ -356,7 +466,7 @@ fn test_system_prompt_preserved_in_locked_state() {
|
||||
);
|
||||
}
|
||||
|
||||
/// unlock -> 再lock でシステムプロンプトを変更できることを確認
|
||||
/// Verify that system prompt can be changed after unlock -> re-lock
|
||||
#[test]
|
||||
fn test_system_prompt_change_after_unlock() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
@@ -1,182 +0,0 @@
|
||||
//! Hook関連の型定義
|
||||
//!
|
||||
//! Worker層でのターン制御・介入に使用される型
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
|
||||
// =============================================================================
|
||||
// Control Flow Types
|
||||
// =============================================================================
|
||||
|
||||
/// Hook処理の制御フロー
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum ControlFlow {
|
||||
/// 処理を続行
|
||||
Continue,
|
||||
/// 現在の処理をスキップ(Tool実行など)
|
||||
Skip,
|
||||
/// 処理を中断
|
||||
Abort(String),
|
||||
}
|
||||
|
||||
/// ターン終了時の判定結果
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum TurnResult {
|
||||
/// ターンを終了
|
||||
Finish,
|
||||
/// メッセージを追加してターン継続(自己修正など)
|
||||
ContinueWithMessages(Vec<crate::Message>),
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Tool Call / Result Types
|
||||
// =============================================================================
|
||||
|
||||
/// ツール呼び出し情報
|
||||
///
|
||||
/// LLMからのToolUseブロックを表現し、Hook処理で改変可能
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolCall {
|
||||
/// ツール呼び出しID(レスポンスとの紐付けに使用)
|
||||
pub id: String,
|
||||
/// ツール名
|
||||
pub name: String,
|
||||
/// 入力引数(JSON)
|
||||
pub input: Value,
|
||||
}
|
||||
|
||||
/// ツール実行結果
|
||||
///
|
||||
/// ツール実行後の結果を表現し、Hook処理で改変可能
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolResult {
|
||||
/// 対応するツール呼び出しID
|
||||
pub tool_use_id: String,
|
||||
/// 結果コンテンツ
|
||||
pub content: String,
|
||||
/// エラーかどうか
|
||||
#[serde(default)]
|
||||
pub is_error: bool,
|
||||
}
|
||||
|
||||
impl ToolResult {
|
||||
/// 成功結果を作成
|
||||
pub fn success(tool_use_id: impl Into<String>, content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
tool_use_id: tool_use_id.into(),
|
||||
content: content.into(),
|
||||
is_error: false,
|
||||
}
|
||||
}
|
||||
|
||||
/// エラー結果を作成
|
||||
pub fn error(tool_use_id: impl Into<String>, content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
tool_use_id: tool_use_id.into(),
|
||||
content: content.into(),
|
||||
is_error: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Hook Error
|
||||
// =============================================================================
|
||||
|
||||
/// Hookエラー
|
||||
#[derive(Debug, Error)]
|
||||
pub enum HookError {
|
||||
/// 処理が中断された
|
||||
#[error("Aborted: {0}")]
|
||||
Aborted(String),
|
||||
/// 内部エラー
|
||||
#[error("Hook error: {0}")]
|
||||
Internal(String),
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// WorkerHook Trait
|
||||
// =============================================================================
|
||||
|
||||
/// ターンの進行・ツール実行に介入するためのトレイト
|
||||
///
|
||||
/// Hookを使うと、メッセージ送信前、ツール実行前後、ターン終了時に
|
||||
/// 処理を挟んだり、実行をキャンセルしたりできます。
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::hook::{ControlFlow, HookError, ToolCall, TurnResult, WorkerHook};
|
||||
/// use worker::Message;
|
||||
///
|
||||
/// struct ValidationHook;
|
||||
///
|
||||
/// #[async_trait::async_trait]
|
||||
/// impl WorkerHook for ValidationHook {
|
||||
/// async fn before_tool_call(&self, call: &mut ToolCall) -> Result<ControlFlow, HookError> {
|
||||
/// // 危険なツールをブロック
|
||||
/// if call.name == "delete_all" {
|
||||
/// return Ok(ControlFlow::Skip);
|
||||
/// }
|
||||
/// Ok(ControlFlow::Continue)
|
||||
/// }
|
||||
///
|
||||
/// async fn on_turn_end(&self, messages: &[Message]) -> Result<TurnResult, HookError> {
|
||||
/// // 条件を満たさなければ追加メッセージで継続
|
||||
/// if messages.len() < 3 {
|
||||
/// return Ok(TurnResult::ContinueWithMessages(vec![
|
||||
/// Message::user("Please elaborate.")
|
||||
/// ]));
|
||||
/// }
|
||||
/// Ok(TurnResult::Finish)
|
||||
/// }
|
||||
/// }
|
||||
/// ```
|
||||
///
|
||||
/// # デフォルト実装
|
||||
///
|
||||
/// すべてのメソッドにはデフォルト実装があり、何も行わず`Continue`を返します。
|
||||
/// 必要なメソッドのみオーバーライドしてください。
|
||||
#[async_trait]
|
||||
pub trait WorkerHook: Send + Sync {
|
||||
/// メッセージ送信前に呼ばれる
|
||||
///
|
||||
/// リクエストに含まれるメッセージリストを参照・改変できます。
|
||||
/// `ControlFlow::Abort`を返すとターンが中断されます。
|
||||
async fn on_message_send(
|
||||
&self,
|
||||
_context: &mut Vec<crate::Message>,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// ツール実行前に呼ばれる
|
||||
///
|
||||
/// ツール呼び出しの引数を書き換えたり、実行をスキップしたりできます。
|
||||
/// `ControlFlow::Skip`を返すとこのツールの実行がスキップされます。
|
||||
async fn before_tool_call(&self, _tool_call: &mut ToolCall) -> Result<ControlFlow, HookError> {
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// ツール実行後に呼ばれる
|
||||
///
|
||||
/// ツールの実行結果を書き換えたり、隠蔽したりできます。
|
||||
async fn after_tool_call(
|
||||
&self,
|
||||
_tool_result: &mut ToolResult,
|
||||
) -> Result<ControlFlow, HookError> {
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// ターン終了時に呼ばれる
|
||||
///
|
||||
/// 生成されたメッセージを検査し、必要なら追加メッセージで継続を指示できます。
|
||||
/// `TurnResult::ContinueWithMessages`を返すと、指定したメッセージを追加して
|
||||
/// 次のターンに進みます。
|
||||
async fn on_turn_end(&self, _messages: &[crate::Message]) -> Result<TurnResult, HookError> {
|
||||
Ok(TurnResult::Finish)
|
||||
}
|
||||
}
|
||||
@@ -1,56 +0,0 @@
|
||||
//! worker - LLMワーカーライブラリ
|
||||
//!
|
||||
//! LLMとの対話を管理するコンポーネントを提供します。
|
||||
//!
|
||||
//! # 主要なコンポーネント
|
||||
//!
|
||||
//! - [`Worker`] - LLMとの対話を管理する中心コンポーネント
|
||||
//! - [`tool::Tool`] - LLMから呼び出し可能なツール
|
||||
//! - [`hook::WorkerHook`] - ターン進行への介入
|
||||
//! - [`subscriber::WorkerSubscriber`] - ストリーミングイベントの購読
|
||||
//!
|
||||
//! # Quick Start
|
||||
//!
|
||||
//! ```ignore
|
||||
//! use worker::{Worker, Message};
|
||||
//!
|
||||
//! // Workerを作成
|
||||
//! let mut worker = Worker::new(client)
|
||||
//! .system_prompt("You are a helpful assistant.");
|
||||
//!
|
||||
//! // ツールを登録(オプション)
|
||||
//! use worker::tool::Tool;
|
||||
//! worker.register_tool(my_tool);
|
||||
//!
|
||||
//! // 対話を実行
|
||||
//! let history = worker.run("Hello!").await?;
|
||||
//! ```
|
||||
//!
|
||||
//! # キャッシュ保護
|
||||
//!
|
||||
//! KVキャッシュのヒット率を最大化するには、[`Worker::lock()`]で
|
||||
//! ロック状態に遷移してから実行してください。
|
||||
//!
|
||||
//! ```ignore
|
||||
//! let mut locked = worker.lock();
|
||||
//! locked.run("user input").await?;
|
||||
//! ```
|
||||
|
||||
pub mod llm_client;
|
||||
pub mod timeline;
|
||||
|
||||
pub mod event;
|
||||
mod handler;
|
||||
pub mod hook;
|
||||
mod message;
|
||||
pub mod state;
|
||||
pub mod subscriber;
|
||||
pub mod tool;
|
||||
mod worker;
|
||||
|
||||
// =============================================================================
|
||||
// トップレベル公開(最も頻繁に使う型)
|
||||
// =============================================================================
|
||||
|
||||
pub use message::{ContentPart, Message, MessageContent, Role};
|
||||
pub use worker::{Worker, WorkerConfig, WorkerError};
|
||||
@@ -1,39 +0,0 @@
|
||||
//! LLMクライアント共通trait定義
|
||||
|
||||
use std::pin::Pin;
|
||||
|
||||
use crate::llm_client::{ClientError, Request, event::Event};
|
||||
use async_trait::async_trait;
|
||||
use futures::Stream;
|
||||
|
||||
/// LLMクライアントのtrait
|
||||
///
|
||||
/// 各プロバイダはこのtraitを実装し、統一されたインターフェースを提供する。
|
||||
#[async_trait]
|
||||
pub trait LlmClient: Send + Sync {
|
||||
/// ストリーミングリクエストを送信し、Eventストリームを返す
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `request` - リクエスト情報
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(Stream)` - イベントストリーム
|
||||
/// * `Err(ClientError)` - エラー
|
||||
async fn stream(
|
||||
&self,
|
||||
request: Request,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<Event, ClientError>> + Send>>, ClientError>;
|
||||
}
|
||||
|
||||
/// `Box<dyn LlmClient>` に対する `LlmClient` の実装
|
||||
///
|
||||
/// これにより、動的ディスパッチを使用するクライアントも `Worker` で利用可能になる。
|
||||
#[async_trait]
|
||||
impl LlmClient for Box<dyn LlmClient> {
|
||||
async fn stream(
|
||||
&self,
|
||||
request: Request,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<Event, ClientError>> + Send>>, ClientError> {
|
||||
(**self).stream(request).await
|
||||
}
|
||||
}
|
||||
@@ -1,195 +0,0 @@
|
||||
//! Anthropic リクエスト生成
|
||||
|
||||
use serde::Serialize;
|
||||
|
||||
use crate::llm_client::{
|
||||
Request,
|
||||
types::{ContentPart, Message, MessageContent, Role, ToolDefinition},
|
||||
};
|
||||
|
||||
use super::AnthropicScheme;
|
||||
|
||||
/// Anthropic APIへのリクエストボディ
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct AnthropicRequest {
|
||||
pub model: String,
|
||||
pub max_tokens: u32,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub system: Option<String>,
|
||||
pub messages: Vec<AnthropicMessage>,
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub tools: Vec<AnthropicTool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub top_p: Option<f32>,
|
||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||
pub stop_sequences: Vec<String>,
|
||||
pub stream: bool,
|
||||
}
|
||||
|
||||
/// Anthropic メッセージ
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct AnthropicMessage {
|
||||
pub role: String,
|
||||
pub content: AnthropicContent,
|
||||
}
|
||||
|
||||
/// Anthropic コンテンツ
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum AnthropicContent {
|
||||
Text(String),
|
||||
Parts(Vec<AnthropicContentPart>),
|
||||
}
|
||||
|
||||
/// Anthropic コンテンツパーツ
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub(crate) enum AnthropicContentPart {
|
||||
#[serde(rename = "text")]
|
||||
Text { text: String },
|
||||
#[serde(rename = "tool_use")]
|
||||
ToolUse {
|
||||
id: String,
|
||||
name: String,
|
||||
input: serde_json::Value,
|
||||
},
|
||||
#[serde(rename = "tool_result")]
|
||||
ToolResult {
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
},
|
||||
}
|
||||
|
||||
/// Anthropic ツール定義
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(crate) struct AnthropicTool {
|
||||
pub name: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
pub input_schema: serde_json::Value,
|
||||
}
|
||||
|
||||
impl AnthropicScheme {
|
||||
/// RequestからAnthropicのリクエストボディを構築
|
||||
pub(crate) fn build_request(&self, model: &str, request: &Request) -> AnthropicRequest {
|
||||
let messages = request
|
||||
.messages
|
||||
.iter()
|
||||
.map(|m| self.convert_message(m))
|
||||
.collect();
|
||||
|
||||
let tools = request.tools.iter().map(|t| self.convert_tool(t)).collect();
|
||||
|
||||
AnthropicRequest {
|
||||
model: model.to_string(),
|
||||
max_tokens: request.config.max_tokens.unwrap_or(4096),
|
||||
system: request.system_prompt.clone(),
|
||||
messages,
|
||||
tools,
|
||||
temperature: request.config.temperature,
|
||||
top_p: request.config.top_p,
|
||||
stop_sequences: request.config.stop_sequences.clone(),
|
||||
stream: true,
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_message(&self, message: &Message) -> AnthropicMessage {
|
||||
let role = match message.role {
|
||||
Role::User => "user",
|
||||
Role::Assistant => "assistant",
|
||||
};
|
||||
|
||||
let content = match &message.content {
|
||||
MessageContent::Text(text) => AnthropicContent::Text(text.clone()),
|
||||
MessageContent::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
} => AnthropicContent::Parts(vec![AnthropicContentPart::ToolResult {
|
||||
tool_use_id: tool_use_id.clone(),
|
||||
content: content.clone(),
|
||||
}]),
|
||||
MessageContent::Parts(parts) => {
|
||||
let converted: Vec<_> = parts
|
||||
.iter()
|
||||
.map(|p| match p {
|
||||
ContentPart::Text { text } => {
|
||||
AnthropicContentPart::Text { text: text.clone() }
|
||||
}
|
||||
ContentPart::ToolUse { id, name, input } => AnthropicContentPart::ToolUse {
|
||||
id: id.clone(),
|
||||
name: name.clone(),
|
||||
input: input.clone(),
|
||||
},
|
||||
ContentPart::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
} => AnthropicContentPart::ToolResult {
|
||||
tool_use_id: tool_use_id.clone(),
|
||||
content: content.clone(),
|
||||
},
|
||||
})
|
||||
.collect();
|
||||
AnthropicContent::Parts(converted)
|
||||
}
|
||||
};
|
||||
|
||||
AnthropicMessage {
|
||||
role: role.to_string(),
|
||||
content,
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_tool(&self, tool: &ToolDefinition) -> AnthropicTool {
|
||||
AnthropicTool {
|
||||
name: tool.name.clone(),
|
||||
description: tool.description.clone(),
|
||||
input_schema: tool.input_schema.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_build_simple_request() {
|
||||
let scheme = AnthropicScheme::new();
|
||||
let request = Request::new()
|
||||
.system("You are a helpful assistant.")
|
||||
.user("Hello!");
|
||||
|
||||
let anthropic_req = scheme.build_request("claude-sonnet-4-20250514", &request);
|
||||
|
||||
assert_eq!(anthropic_req.model, "claude-sonnet-4-20250514");
|
||||
assert_eq!(
|
||||
anthropic_req.system,
|
||||
Some("You are a helpful assistant.".to_string())
|
||||
);
|
||||
assert_eq!(anthropic_req.messages.len(), 1);
|
||||
assert!(anthropic_req.stream);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_request_with_tool() {
|
||||
let scheme = AnthropicScheme::new();
|
||||
let request = Request::new().user("What's the weather?").tool(
|
||||
ToolDefinition::new("get_weather")
|
||||
.description("Get current weather")
|
||||
.input_schema(serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": { "type": "string" }
|
||||
},
|
||||
"required": ["location"]
|
||||
})),
|
||||
);
|
||||
|
||||
let anthropic_req = scheme.build_request("claude-sonnet-4-20250514", &request);
|
||||
|
||||
assert_eq!(anthropic_req.tools.len(), 1);
|
||||
assert_eq!(anthropic_req.tools[0].name, "get_weather");
|
||||
}
|
||||
}
|
||||
@@ -1,198 +0,0 @@
|
||||
//! LLMクライアント共通型定義
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// リクエスト構造体
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct Request {
|
||||
/// システムプロンプト
|
||||
pub system_prompt: Option<String>,
|
||||
/// メッセージ履歴
|
||||
pub messages: Vec<Message>,
|
||||
/// ツール定義
|
||||
pub tools: Vec<ToolDefinition>,
|
||||
/// リクエスト設定
|
||||
pub config: RequestConfig,
|
||||
}
|
||||
|
||||
impl Request {
|
||||
/// 新しいリクエストを作成
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
/// システムプロンプトを設定
|
||||
pub fn system(mut self, prompt: impl Into<String>) -> Self {
|
||||
self.system_prompt = Some(prompt.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// ユーザーメッセージを追加
|
||||
pub fn user(mut self, content: impl Into<String>) -> Self {
|
||||
self.messages.push(Message::user(content));
|
||||
self
|
||||
}
|
||||
|
||||
/// アシスタントメッセージを追加
|
||||
pub fn assistant(mut self, content: impl Into<String>) -> Self {
|
||||
self.messages.push(Message::assistant(content));
|
||||
self
|
||||
}
|
||||
|
||||
/// メッセージを追加
|
||||
pub fn message(mut self, message: Message) -> Self {
|
||||
self.messages.push(message);
|
||||
self
|
||||
}
|
||||
|
||||
/// ツールを追加
|
||||
pub fn tool(mut self, tool: ToolDefinition) -> Self {
|
||||
self.tools.push(tool);
|
||||
self
|
||||
}
|
||||
|
||||
/// 設定を適用
|
||||
pub fn config(mut self, config: RequestConfig) -> Self {
|
||||
self.config = config;
|
||||
self
|
||||
}
|
||||
|
||||
/// max_tokensを設定
|
||||
pub fn max_tokens(mut self, max_tokens: u32) -> Self {
|
||||
self.config.max_tokens = Some(max_tokens);
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// メッセージ
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct Message {
|
||||
/// ロール
|
||||
pub role: Role,
|
||||
/// コンテンツ
|
||||
pub content: MessageContent,
|
||||
}
|
||||
|
||||
impl Message {
|
||||
/// ユーザーメッセージを作成
|
||||
pub fn user(content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
role: Role::User,
|
||||
content: MessageContent::Text(content.into()),
|
||||
}
|
||||
}
|
||||
|
||||
/// アシスタントメッセージを作成
|
||||
pub fn assistant(content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Text(content.into()),
|
||||
}
|
||||
}
|
||||
|
||||
/// ツール結果メッセージを作成
|
||||
pub fn tool_result(tool_use_id: impl Into<String>, content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
role: Role::User,
|
||||
content: MessageContent::ToolResult {
|
||||
tool_use_id: tool_use_id.into(),
|
||||
content: content.into(),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// ロール
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum Role {
|
||||
User,
|
||||
Assistant,
|
||||
}
|
||||
|
||||
/// メッセージコンテンツ
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum MessageContent {
|
||||
/// テキストコンテンツ
|
||||
Text(String),
|
||||
/// ツール結果
|
||||
ToolResult {
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
},
|
||||
/// 複合コンテンツ (テキスト + ツール使用等)
|
||||
Parts(Vec<ContentPart>),
|
||||
}
|
||||
|
||||
/// コンテンツパーツ
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum ContentPart {
|
||||
/// テキスト
|
||||
#[serde(rename = "text")]
|
||||
Text { text: String },
|
||||
/// ツール使用
|
||||
#[serde(rename = "tool_use")]
|
||||
ToolUse {
|
||||
id: String,
|
||||
name: String,
|
||||
input: serde_json::Value,
|
||||
},
|
||||
/// ツール結果
|
||||
#[serde(rename = "tool_result")]
|
||||
ToolResult {
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
},
|
||||
}
|
||||
|
||||
/// ツール定義
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolDefinition {
|
||||
/// ツール名
|
||||
pub name: String,
|
||||
/// 説明
|
||||
pub description: Option<String>,
|
||||
/// 入力スキーマ (JSON Schema)
|
||||
pub input_schema: serde_json::Value,
|
||||
}
|
||||
|
||||
impl ToolDefinition {
|
||||
/// 新しいツール定義を作成
|
||||
pub fn new(name: impl Into<String>) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
description: None,
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
/// 説明を設定
|
||||
pub fn description(mut self, desc: impl Into<String>) -> Self {
|
||||
self.description = Some(desc.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// 入力スキーマを設定
|
||||
pub fn input_schema(mut self, schema: serde_json::Value) -> Self {
|
||||
self.input_schema = schema;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// リクエスト設定
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct RequestConfig {
|
||||
/// 最大トークン数
|
||||
pub max_tokens: Option<u32>,
|
||||
/// Temperature
|
||||
pub temperature: Option<f32>,
|
||||
/// Top P
|
||||
pub top_p: Option<f32>,
|
||||
/// ストップシーケンス
|
||||
pub stop_sequences: Vec<String>,
|
||||
}
|
||||
@@ -1,116 +0,0 @@
|
||||
//! メッセージ型
|
||||
//!
|
||||
//! LLMとの会話で使用されるメッセージ構造。
|
||||
//! [`Message::user`]や[`Message::assistant`]で簡単に作成できます。
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// メッセージのロール
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum Role {
|
||||
/// ユーザー
|
||||
User,
|
||||
/// アシスタント
|
||||
Assistant,
|
||||
}
|
||||
|
||||
/// 会話のメッセージ
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::Message;
|
||||
///
|
||||
/// // ユーザーメッセージ
|
||||
/// let user_msg = Message::user("Hello!");
|
||||
///
|
||||
/// // アシスタントメッセージ
|
||||
/// let assistant_msg = Message::assistant("Hi there!");
|
||||
/// ```
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct Message {
|
||||
/// ロール
|
||||
pub role: Role,
|
||||
/// コンテンツ
|
||||
pub content: MessageContent,
|
||||
}
|
||||
|
||||
/// メッセージコンテンツ
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum MessageContent {
|
||||
/// テキストコンテンツ
|
||||
Text(String),
|
||||
/// ツール結果
|
||||
ToolResult {
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
},
|
||||
/// 複合コンテンツ (テキスト + ツール使用等)
|
||||
Parts(Vec<ContentPart>),
|
||||
}
|
||||
|
||||
/// コンテンツパーツ
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum ContentPart {
|
||||
/// テキスト
|
||||
#[serde(rename = "text")]
|
||||
Text { text: String },
|
||||
/// ツール使用
|
||||
#[serde(rename = "tool_use")]
|
||||
ToolUse {
|
||||
id: String,
|
||||
name: String,
|
||||
input: serde_json::Value,
|
||||
},
|
||||
/// ツール結果
|
||||
#[serde(rename = "tool_result")]
|
||||
ToolResult {
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
},
|
||||
}
|
||||
|
||||
impl Message {
|
||||
/// ユーザーメッセージを作成
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::Message;
|
||||
/// let msg = Message::user("こんにちは");
|
||||
/// ```
|
||||
pub fn user(content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
role: Role::User,
|
||||
content: MessageContent::Text(content.into()),
|
||||
}
|
||||
}
|
||||
|
||||
/// アシスタントメッセージを作成
|
||||
///
|
||||
/// 通常はWorker内部で自動生成されますが、
|
||||
/// 履歴の初期化などで手動作成も可能です。
|
||||
pub fn assistant(content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Text(content.into()),
|
||||
}
|
||||
}
|
||||
|
||||
/// ツール結果メッセージを作成
|
||||
///
|
||||
/// Worker内部でツール実行後に自動生成されます。
|
||||
/// 通常は直接作成する必要はありません。
|
||||
pub fn tool_result(tool_use_id: impl Into<String>, content: impl Into<String>) -> Self {
|
||||
Self {
|
||||
role: Role::User,
|
||||
content: MessageContent::ToolResult {
|
||||
tool_use_id: tool_use_id.into(),
|
||||
content: content.into(),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,60 +0,0 @@
|
||||
//! Worker状態
|
||||
//!
|
||||
//! Type-stateパターンによるキャッシュ保護のための状態マーカー型。
|
||||
//! Workerは`Mutable` → `Locked`の状態遷移を持ちます。
|
||||
|
||||
/// Worker状態を表すマーカートレイト
|
||||
///
|
||||
/// このトレイトはシールされており、外部から実装することはできません。
|
||||
pub trait WorkerState: private::Sealed + Send + Sync + 'static {}
|
||||
|
||||
mod private {
|
||||
pub trait Sealed {}
|
||||
}
|
||||
|
||||
/// 編集可能状態
|
||||
///
|
||||
/// この状態では以下の操作が可能です:
|
||||
/// - システムプロンプトの設定・変更
|
||||
/// - メッセージ履歴の編集(追加、削除、クリア)
|
||||
/// - ツール・Hookの登録
|
||||
///
|
||||
/// `Worker::lock()`により[`Locked`]状態へ遷移できます。
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::Worker;
|
||||
///
|
||||
/// let mut worker = Worker::new(client)
|
||||
/// .system_prompt("You are helpful.");
|
||||
///
|
||||
/// // 履歴を編集可能
|
||||
/// worker.push_message(Message::user("Hello"));
|
||||
/// worker.clear_history();
|
||||
///
|
||||
/// // ロックして保護状態へ
|
||||
/// let locked = worker.lock();
|
||||
/// ```
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
pub struct Mutable;
|
||||
|
||||
impl private::Sealed for Mutable {}
|
||||
impl WorkerState for Mutable {}
|
||||
|
||||
/// ロック状態(キャッシュ保護)
|
||||
///
|
||||
/// この状態では以下の制限があります:
|
||||
/// - システムプロンプトの変更不可
|
||||
/// - 既存メッセージ履歴の変更不可(末尾への追記のみ)
|
||||
///
|
||||
/// LLM APIのKVキャッシュヒットを保証するため、
|
||||
/// 実行時にはこの状態の使用が推奨されます。
|
||||
///
|
||||
/// `Worker::unlock()`により[`Mutable`]状態へ戻せますが、
|
||||
/// キャッシュ保護が解除されることに注意してください。
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
pub struct Locked;
|
||||
|
||||
impl private::Sealed for Locked {}
|
||||
impl WorkerState for Locked {}
|
||||
@@ -1,90 +0,0 @@
|
||||
//! ツール定義
|
||||
//!
|
||||
//! LLMから呼び出し可能なツールを定義するためのトレイト。
|
||||
//! 通常は`#[tool]`マクロを使用して自動実装します。
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
|
||||
/// ツール実行時のエラー
|
||||
#[derive(Debug, Error)]
|
||||
pub enum ToolError {
|
||||
/// 引数が不正
|
||||
#[error("Invalid argument: {0}")]
|
||||
InvalidArgument(String),
|
||||
/// 実行に失敗
|
||||
#[error("Execution failed: {0}")]
|
||||
ExecutionFailed(String),
|
||||
/// 内部エラー
|
||||
#[error("Internal error: {0}")]
|
||||
Internal(String),
|
||||
}
|
||||
|
||||
/// LLMから呼び出し可能なツールを定義するトレイト
|
||||
///
|
||||
/// ツールはLLMが外部リソースにアクセスしたり、
|
||||
/// 計算を実行したりするために使用します。
|
||||
///
|
||||
/// # 実装方法
|
||||
///
|
||||
/// 通常は`#[tool]`マクロを使用して自動実装します:
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::tool;
|
||||
///
|
||||
/// #[tool(description = "Search the web for information")]
|
||||
/// async fn search(query: String) -> String {
|
||||
/// // 検索処理
|
||||
/// format!("Results for: {}", query)
|
||||
/// }
|
||||
/// ```
|
||||
///
|
||||
/// # 手動実装
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::tool::{Tool, ToolError};
|
||||
/// use serde_json::{json, Value};
|
||||
///
|
||||
/// struct MyTool;
|
||||
///
|
||||
/// #[async_trait::async_trait]
|
||||
/// impl Tool for MyTool {
|
||||
/// fn name(&self) -> &str { "my_tool" }
|
||||
/// fn description(&self) -> &str { "My custom tool" }
|
||||
/// fn input_schema(&self) -> Value {
|
||||
/// json!({
|
||||
/// "type": "object",
|
||||
/// "properties": {
|
||||
/// "query": { "type": "string" }
|
||||
/// },
|
||||
/// "required": ["query"]
|
||||
/// })
|
||||
/// }
|
||||
/// async fn execute(&self, input: &str) -> Result<String, ToolError> {
|
||||
/// Ok("result".to_string())
|
||||
/// }
|
||||
/// }
|
||||
/// ```
|
||||
#[async_trait]
|
||||
pub trait Tool: Send + Sync {
|
||||
/// ツール名(LLMが識別に使用)
|
||||
fn name(&self) -> &str;
|
||||
|
||||
/// ツールの説明(LLMへのプロンプトに含まれる)
|
||||
fn description(&self) -> &str;
|
||||
|
||||
/// 引数のJSON Schema
|
||||
///
|
||||
/// LLMはこのスキーマに従って引数を生成します。
|
||||
fn input_schema(&self) -> Value;
|
||||
|
||||
/// ツールを実行する
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `input_json` - LLMが生成したJSON形式の引数
|
||||
///
|
||||
/// # Returns
|
||||
/// 実行結果の文字列。この内容がLLMに返されます。
|
||||
async fn execute(&self, input_json: &str) -> Result<String, ToolError>;
|
||||
}
|
||||
@@ -1,788 +0,0 @@
|
||||
use std::collections::HashMap;
|
||||
use std::marker::PhantomData;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use futures::StreamExt;
|
||||
use tracing::{debug, info, trace, warn};
|
||||
|
||||
use crate::{
|
||||
ContentPart, Message, MessageContent, Role,
|
||||
hook::{ControlFlow, HookError, ToolCall, ToolResult, TurnResult, WorkerHook},
|
||||
llm_client::{ClientError, LlmClient, Request, ToolDefinition},
|
||||
state::{Locked, Mutable, WorkerState},
|
||||
subscriber::{
|
||||
ErrorSubscriberAdapter, StatusSubscriberAdapter, TextBlockSubscriberAdapter,
|
||||
ToolUseBlockSubscriberAdapter, UsageSubscriberAdapter, WorkerSubscriber,
|
||||
},
|
||||
timeline::{TextBlockCollector, Timeline, ToolCallCollector},
|
||||
tool::{Tool, ToolError},
|
||||
};
|
||||
|
||||
// =============================================================================
|
||||
// Worker Error
|
||||
// =============================================================================
|
||||
|
||||
/// Workerエラー
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum WorkerError {
|
||||
/// クライアントエラー
|
||||
#[error("Client error: {0}")]
|
||||
Client(#[from] ClientError),
|
||||
/// ツールエラー
|
||||
#[error("Tool error: {0}")]
|
||||
Tool(#[from] ToolError),
|
||||
/// Hookエラー
|
||||
#[error("Hook error: {0}")]
|
||||
Hook(#[from] HookError),
|
||||
/// 処理が中断された
|
||||
#[error("Aborted: {0}")]
|
||||
Aborted(String),
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Worker Config
|
||||
// =============================================================================
|
||||
|
||||
/// Worker設定
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct WorkerConfig {
|
||||
// 将来の拡張用(現在は空)
|
||||
_private: (),
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ターン制御用コールバック保持
|
||||
// =============================================================================
|
||||
|
||||
/// ターンイベントを通知するためのコールバック (型消去)
|
||||
trait TurnNotifier: Send {
|
||||
fn on_turn_start(&self, turn: usize);
|
||||
fn on_turn_end(&self, turn: usize);
|
||||
}
|
||||
|
||||
struct SubscriberTurnNotifier<S: WorkerSubscriber + 'static> {
|
||||
subscriber: Arc<Mutex<S>>,
|
||||
}
|
||||
|
||||
impl<S: WorkerSubscriber + 'static> TurnNotifier for SubscriberTurnNotifier<S> {
|
||||
fn on_turn_start(&self, turn: usize) {
|
||||
if let Ok(mut s) = self.subscriber.lock() {
|
||||
s.on_turn_start(turn);
|
||||
}
|
||||
}
|
||||
|
||||
fn on_turn_end(&self, turn: usize) {
|
||||
if let Ok(mut s) = self.subscriber.lock() {
|
||||
s.on_turn_end(turn);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Worker
|
||||
// =============================================================================
|
||||
|
||||
/// LLMとの対話を管理する中心コンポーネント
|
||||
///
|
||||
/// ユーザーからの入力を受け取り、LLMにリクエストを送信し、
|
||||
/// ツール呼び出しがあれば自動的に実行してターンを進行させます。
|
||||
///
|
||||
/// # 状態遷移(Type-state)
|
||||
///
|
||||
/// - [`Mutable`]: 初期状態。システムプロンプトや履歴を自由に編集可能。
|
||||
/// - [`Locked`]: キャッシュ保護状態。`lock()`で遷移。前方コンテキストは不変。
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::{Worker, Message};
|
||||
///
|
||||
/// // Workerを作成してツールを登録
|
||||
/// let mut worker = Worker::new(client)
|
||||
/// .system_prompt("You are a helpful assistant.");
|
||||
/// worker.register_tool(my_tool);
|
||||
///
|
||||
/// // 対話を実行
|
||||
/// let history = worker.run("Hello!").await?;
|
||||
/// ```
|
||||
///
|
||||
/// # キャッシュ保護が必要な場合
|
||||
///
|
||||
/// ```ignore
|
||||
/// let mut worker = Worker::new(client)
|
||||
/// .system_prompt("...");
|
||||
///
|
||||
/// // 履歴を設定後、ロックしてキャッシュを保護
|
||||
/// let mut locked = worker.lock();
|
||||
/// locked.run("user input").await?;
|
||||
/// ```
|
||||
pub struct Worker<C: LlmClient, S: WorkerState = Mutable> {
|
||||
/// LLMクライアント
|
||||
client: C,
|
||||
/// イベントタイムライン
|
||||
timeline: Timeline,
|
||||
/// テキストブロックコレクター(Timeline用ハンドラ)
|
||||
text_block_collector: TextBlockCollector,
|
||||
/// ツールコールコレクター(Timeline用ハンドラ)
|
||||
tool_call_collector: ToolCallCollector,
|
||||
/// 登録されたツール
|
||||
tools: HashMap<String, Arc<dyn Tool>>,
|
||||
/// 登録されたHook
|
||||
hooks: Vec<Box<dyn WorkerHook>>,
|
||||
/// システムプロンプト
|
||||
system_prompt: Option<String>,
|
||||
/// メッセージ履歴(Workerが所有)
|
||||
history: Vec<Message>,
|
||||
/// ロック時点での履歴長(Locked状態でのみ意味を持つ)
|
||||
locked_prefix_len: usize,
|
||||
/// ターンカウント
|
||||
turn_count: usize,
|
||||
/// ターン通知用のコールバック
|
||||
turn_notifiers: Vec<Box<dyn TurnNotifier>>,
|
||||
/// 状態マーカー
|
||||
_state: PhantomData<S>,
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// 共通実装(全状態で利用可能)
|
||||
// =============================================================================
|
||||
|
||||
impl<C: LlmClient, S: WorkerState> Worker<C, S> {
|
||||
/// イベント購読者を登録する
|
||||
///
|
||||
/// 登録したSubscriberは、LLMからのストリーミングイベントを
|
||||
/// リアルタイムで受信できます。UIへのストリーム表示などに利用します。
|
||||
///
|
||||
/// # 受信できるイベント
|
||||
///
|
||||
/// - **ブロックイベント**: `on_text_block`, `on_tool_use_block`
|
||||
/// - **メタイベント**: `on_usage`, `on_status`, `on_error`
|
||||
/// - **完了イベント**: `on_text_complete`, `on_tool_call_complete`
|
||||
/// - **ターン制御**: `on_turn_start`, `on_turn_end`
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::{Worker, WorkerSubscriber, TextBlockEvent};
|
||||
///
|
||||
/// struct MyPrinter;
|
||||
/// impl WorkerSubscriber for MyPrinter {
|
||||
/// type TextBlockScope = ();
|
||||
/// type ToolUseBlockScope = ();
|
||||
///
|
||||
/// fn on_text_block(&mut self, _: &mut (), event: &TextBlockEvent) {
|
||||
/// if let TextBlockEvent::Delta(text) = event {
|
||||
/// print!("{}", text);
|
||||
/// }
|
||||
/// }
|
||||
/// }
|
||||
///
|
||||
/// worker.subscribe(MyPrinter);
|
||||
/// ```
|
||||
pub fn subscribe<Sub: WorkerSubscriber + 'static>(&mut self, subscriber: Sub) {
|
||||
let subscriber = Arc::new(Mutex::new(subscriber));
|
||||
|
||||
// TextBlock用ハンドラを登録
|
||||
self.timeline
|
||||
.on_text_block(TextBlockSubscriberAdapter::new(subscriber.clone()));
|
||||
|
||||
// ToolUseBlock用ハンドラを登録
|
||||
self.timeline
|
||||
.on_tool_use_block(ToolUseBlockSubscriberAdapter::new(subscriber.clone()));
|
||||
|
||||
// Meta系ハンドラを登録
|
||||
self.timeline
|
||||
.on_usage(UsageSubscriberAdapter::new(subscriber.clone()));
|
||||
self.timeline
|
||||
.on_status(StatusSubscriberAdapter::new(subscriber.clone()));
|
||||
self.timeline
|
||||
.on_error(ErrorSubscriberAdapter::new(subscriber.clone()));
|
||||
|
||||
// ターン制御用コールバックを登録
|
||||
self.turn_notifiers
|
||||
.push(Box::new(SubscriberTurnNotifier { subscriber }));
|
||||
}
|
||||
|
||||
/// ツールを登録する
|
||||
///
|
||||
/// 登録されたツールはLLMからの呼び出しで自動的に実行されます。
|
||||
/// 同名のツールを登録した場合、後から登録したものが優先されます。
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::Worker;
|
||||
/// use my_tools::SearchTool;
|
||||
///
|
||||
/// worker.register_tool(SearchTool::new());
|
||||
/// ```
|
||||
pub fn register_tool(&mut self, tool: impl Tool + 'static) {
|
||||
let name = tool.name().to_string();
|
||||
self.tools.insert(name, Arc::new(tool));
|
||||
}
|
||||
|
||||
/// 複数のツールを登録
|
||||
pub fn register_tools(&mut self, tools: impl IntoIterator<Item = impl Tool + 'static>) {
|
||||
for tool in tools {
|
||||
self.register_tool(tool);
|
||||
}
|
||||
}
|
||||
|
||||
/// Hookを追加する
|
||||
///
|
||||
/// Hookはターンの進行・ツール実行に介入できます。
|
||||
/// 複数のHookを登録した場合、登録順に実行されます。
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```ignore
|
||||
/// use worker::{Worker, WorkerHook, ControlFlow, ToolCall};
|
||||
///
|
||||
/// struct LoggingHook;
|
||||
///
|
||||
/// #[async_trait::async_trait]
|
||||
/// impl WorkerHook for LoggingHook {
|
||||
/// async fn before_tool_call(&self, call: &mut ToolCall) -> Result<ControlFlow, HookError> {
|
||||
/// println!("Calling tool: {}", call.name);
|
||||
/// Ok(ControlFlow::Continue)
|
||||
/// }
|
||||
/// }
|
||||
///
|
||||
/// worker.add_hook(LoggingHook);
|
||||
/// ```
|
||||
pub fn add_hook(&mut self, hook: impl WorkerHook + 'static) {
|
||||
self.hooks.push(Box::new(hook));
|
||||
}
|
||||
|
||||
/// タイムラインへの可変参照を取得(追加ハンドラ登録用)
|
||||
pub fn timeline_mut(&mut self) -> &mut Timeline {
|
||||
&mut self.timeline
|
||||
}
|
||||
|
||||
/// 履歴への参照を取得
|
||||
pub fn history(&self) -> &[Message] {
|
||||
&self.history
|
||||
}
|
||||
|
||||
/// システムプロンプトへの参照を取得
|
||||
pub fn get_system_prompt(&self) -> Option<&str> {
|
||||
self.system_prompt.as_deref()
|
||||
}
|
||||
|
||||
/// 現在のターンカウントを取得
|
||||
pub fn turn_count(&self) -> usize {
|
||||
self.turn_count
|
||||
}
|
||||
|
||||
/// 登録されたツールからToolDefinitionのリストを生成
|
||||
fn build_tool_definitions(&self) -> Vec<ToolDefinition> {
|
||||
self.tools
|
||||
.values()
|
||||
.map(|tool| {
|
||||
ToolDefinition::new(tool.name())
|
||||
.description(tool.description())
|
||||
.input_schema(tool.input_schema())
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// テキストブロックとツール呼び出しからアシスタントメッセージを構築
|
||||
fn build_assistant_message(
|
||||
&self,
|
||||
text_blocks: &[String],
|
||||
tool_calls: &[ToolCall],
|
||||
) -> Option<Message> {
|
||||
// テキストもツール呼び出しもない場合はNone
|
||||
if text_blocks.is_empty() && tool_calls.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
// テキストのみの場合はシンプルなテキストメッセージ
|
||||
if tool_calls.is_empty() {
|
||||
let text = text_blocks.join("");
|
||||
return Some(Message::assistant(text));
|
||||
}
|
||||
|
||||
// ツール呼び出しがある場合は Parts として構築
|
||||
let mut parts = Vec::new();
|
||||
|
||||
// テキストパーツを追加
|
||||
for text in text_blocks {
|
||||
if !text.is_empty() {
|
||||
parts.push(ContentPart::Text { text: text.clone() });
|
||||
}
|
||||
}
|
||||
|
||||
// ツール呼び出しパーツを追加
|
||||
for call in tool_calls {
|
||||
parts.push(ContentPart::ToolUse {
|
||||
id: call.id.clone(),
|
||||
name: call.name.clone(),
|
||||
input: call.input.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
Some(Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Parts(parts),
|
||||
})
|
||||
}
|
||||
|
||||
/// リクエストを構築
|
||||
fn build_request(&self, tool_definitions: &[ToolDefinition]) -> Request {
|
||||
let mut request = Request::new();
|
||||
|
||||
// システムプロンプトを設定
|
||||
if let Some(ref system) = self.system_prompt {
|
||||
request = request.system(system);
|
||||
}
|
||||
|
||||
// メッセージを追加
|
||||
for msg in &self.history {
|
||||
// Message から llm_client::Message への変換
|
||||
request = request.message(crate::llm_client::Message {
|
||||
role: match msg.role {
|
||||
Role::User => crate::llm_client::Role::User,
|
||||
Role::Assistant => crate::llm_client::Role::Assistant,
|
||||
},
|
||||
content: match &msg.content {
|
||||
MessageContent::Text(t) => crate::llm_client::MessageContent::Text(t.clone()),
|
||||
MessageContent::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
} => crate::llm_client::MessageContent::ToolResult {
|
||||
tool_use_id: tool_use_id.clone(),
|
||||
content: content.clone(),
|
||||
},
|
||||
MessageContent::Parts(parts) => crate::llm_client::MessageContent::Parts(
|
||||
parts
|
||||
.iter()
|
||||
.map(|p| match p {
|
||||
ContentPart::Text { text } => {
|
||||
crate::llm_client::ContentPart::Text { text: text.clone() }
|
||||
}
|
||||
ContentPart::ToolUse { id, name, input } => {
|
||||
crate::llm_client::ContentPart::ToolUse {
|
||||
id: id.clone(),
|
||||
name: name.clone(),
|
||||
input: input.clone(),
|
||||
}
|
||||
}
|
||||
ContentPart::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
} => crate::llm_client::ContentPart::ToolResult {
|
||||
tool_use_id: tool_use_id.clone(),
|
||||
content: content.clone(),
|
||||
},
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
// ツール定義を追加
|
||||
for tool_def in tool_definitions {
|
||||
request = request.tool(tool_def.clone());
|
||||
}
|
||||
|
||||
request
|
||||
}
|
||||
|
||||
/// Hooks: on_message_send
|
||||
async fn run_on_message_send_hooks(&self) -> Result<ControlFlow, WorkerError> {
|
||||
for hook in &self.hooks {
|
||||
// Note: Locked状態でも履歴全体を参照として渡す(変更は不可)
|
||||
// HookのAPIを変更し、immutable参照のみを渡すようにする必要があるかもしれない
|
||||
// 現在は空のVecを渡して回避(要検討)
|
||||
let mut temp_context = self.history.clone();
|
||||
let result = hook.on_message_send(&mut temp_context).await?;
|
||||
match result {
|
||||
ControlFlow::Continue => continue,
|
||||
ControlFlow::Skip => return Ok(ControlFlow::Skip),
|
||||
ControlFlow::Abort(reason) => return Ok(ControlFlow::Abort(reason)),
|
||||
}
|
||||
}
|
||||
Ok(ControlFlow::Continue)
|
||||
}
|
||||
|
||||
/// Hooks: on_turn_end
|
||||
async fn run_on_turn_end_hooks(&self) -> Result<TurnResult, WorkerError> {
|
||||
for hook in &self.hooks {
|
||||
let result = hook.on_turn_end(&self.history).await?;
|
||||
match result {
|
||||
TurnResult::Finish => continue,
|
||||
TurnResult::ContinueWithMessages(msgs) => {
|
||||
return Ok(TurnResult::ContinueWithMessages(msgs));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(TurnResult::Finish)
|
||||
}
|
||||
|
||||
/// ツールを並列実行
|
||||
///
|
||||
/// 全てのツールに対してbefore_tool_callフックを実行後、
|
||||
/// 許可されたツールを並列に実行し、結果にafter_tool_callフックを適用する。
|
||||
async fn execute_tools(
|
||||
&self,
|
||||
tool_calls: Vec<ToolCall>,
|
||||
) -> Result<Vec<ToolResult>, WorkerError> {
|
||||
use futures::future::join_all;
|
||||
|
||||
// Phase 1: before_tool_call フックを適用(スキップ/中断を判定)
|
||||
let mut approved_calls = Vec::new();
|
||||
for mut tool_call in tool_calls {
|
||||
let mut skip = false;
|
||||
for hook in &self.hooks {
|
||||
let result = hook.before_tool_call(&mut tool_call).await?;
|
||||
match result {
|
||||
ControlFlow::Continue => {}
|
||||
ControlFlow::Skip => {
|
||||
skip = true;
|
||||
break;
|
||||
}
|
||||
ControlFlow::Abort(reason) => {
|
||||
return Err(WorkerError::Aborted(reason));
|
||||
}
|
||||
}
|
||||
}
|
||||
if !skip {
|
||||
approved_calls.push(tool_call);
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 2: 許可されたツールを並列実行
|
||||
let futures: Vec<_> = approved_calls
|
||||
.into_iter()
|
||||
.map(|tool_call| {
|
||||
let tools = &self.tools;
|
||||
async move {
|
||||
if let Some(tool) = tools.get(&tool_call.name) {
|
||||
let input_json =
|
||||
serde_json::to_string(&tool_call.input).unwrap_or_default();
|
||||
match tool.execute(&input_json).await {
|
||||
Ok(content) => ToolResult::success(&tool_call.id, content),
|
||||
Err(e) => ToolResult::error(&tool_call.id, e.to_string()),
|
||||
}
|
||||
} else {
|
||||
ToolResult::error(
|
||||
&tool_call.id,
|
||||
format!("Tool '{}' not found", tool_call.name),
|
||||
)
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
let mut results = join_all(futures).await;
|
||||
|
||||
// Phase 3: after_tool_call フックを適用
|
||||
for tool_result in &mut results {
|
||||
for hook in &self.hooks {
|
||||
let result = hook.after_tool_call(tool_result).await?;
|
||||
match result {
|
||||
ControlFlow::Continue => {}
|
||||
ControlFlow::Skip => break,
|
||||
ControlFlow::Abort(reason) => {
|
||||
return Err(WorkerError::Aborted(reason));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(results)
|
||||
}
|
||||
|
||||
/// 内部で使用するターン実行ロジック
|
||||
async fn run_turn_loop(&mut self) -> Result<(), WorkerError> {
|
||||
let tool_definitions = self.build_tool_definitions();
|
||||
|
||||
info!(
|
||||
message_count = self.history.len(),
|
||||
tool_count = tool_definitions.len(),
|
||||
"Starting worker run"
|
||||
);
|
||||
|
||||
loop {
|
||||
// ターン開始を通知
|
||||
let current_turn = self.turn_count;
|
||||
debug!(turn = current_turn, "Turn start");
|
||||
for notifier in &self.turn_notifiers {
|
||||
notifier.on_turn_start(current_turn);
|
||||
}
|
||||
|
||||
// Hook: on_message_send
|
||||
let control = self.run_on_message_send_hooks().await?;
|
||||
if let ControlFlow::Abort(reason) = control {
|
||||
warn!(reason = %reason, "Aborted by hook");
|
||||
// ターン終了を通知(異常終了)
|
||||
for notifier in &self.turn_notifiers {
|
||||
notifier.on_turn_end(current_turn);
|
||||
}
|
||||
return Err(WorkerError::Aborted(reason));
|
||||
}
|
||||
|
||||
// リクエスト構築
|
||||
let request = self.build_request(&tool_definitions);
|
||||
debug!(
|
||||
message_count = request.messages.len(),
|
||||
tool_count = request.tools.len(),
|
||||
has_system = request.system_prompt.is_some(),
|
||||
"Sending request to LLM"
|
||||
);
|
||||
|
||||
// ストリーム処理
|
||||
debug!("Starting stream...");
|
||||
let mut stream = self.client.stream(request).await?;
|
||||
let mut event_count = 0;
|
||||
while let Some(event_result) = stream.next().await {
|
||||
match &event_result {
|
||||
Ok(event) => {
|
||||
trace!(event = ?event, "Received event");
|
||||
event_count += 1;
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(error = %e, "Stream error");
|
||||
}
|
||||
}
|
||||
let event = event_result?;
|
||||
let timeline_event: crate::timeline::event::Event = event.into();
|
||||
self.timeline.dispatch(&timeline_event);
|
||||
}
|
||||
debug!(event_count = event_count, "Stream completed");
|
||||
|
||||
// ターン終了を通知
|
||||
for notifier in &self.turn_notifiers {
|
||||
notifier.on_turn_end(current_turn);
|
||||
}
|
||||
self.turn_count += 1;
|
||||
|
||||
// 収集結果を取得
|
||||
let text_blocks = self.text_block_collector.take_collected();
|
||||
let tool_calls = self.tool_call_collector.take_collected();
|
||||
|
||||
// アシスタントメッセージを履歴に追加
|
||||
let assistant_message = self.build_assistant_message(&text_blocks, &tool_calls);
|
||||
if let Some(msg) = assistant_message {
|
||||
self.history.push(msg);
|
||||
}
|
||||
|
||||
if tool_calls.is_empty() {
|
||||
// ツール呼び出しなし → ターン終了判定
|
||||
let turn_result = self.run_on_turn_end_hooks().await?;
|
||||
match turn_result {
|
||||
TurnResult::Finish => {
|
||||
return Ok(());
|
||||
}
|
||||
TurnResult::ContinueWithMessages(additional) => {
|
||||
self.history.extend(additional);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ツール実行
|
||||
let tool_results = self.execute_tools(tool_calls).await?;
|
||||
|
||||
// ツール結果を履歴に追加
|
||||
for result in tool_results {
|
||||
self.history
|
||||
.push(Message::tool_result(&result.tool_use_id, &result.content));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Mutable状態専用の実装
|
||||
// =============================================================================
|
||||
|
||||
impl<C: LlmClient> Worker<C, Mutable> {
|
||||
/// 新しいWorkerを作成(Mutable状態)
|
||||
pub fn new(client: C) -> Self {
|
||||
let text_block_collector = TextBlockCollector::new();
|
||||
let tool_call_collector = ToolCallCollector::new();
|
||||
let mut timeline = Timeline::new();
|
||||
|
||||
// コレクターをTimelineに登録
|
||||
timeline.on_text_block(text_block_collector.clone());
|
||||
timeline.on_tool_use_block(tool_call_collector.clone());
|
||||
|
||||
Self {
|
||||
client,
|
||||
timeline,
|
||||
text_block_collector,
|
||||
tool_call_collector,
|
||||
tools: HashMap::new(),
|
||||
hooks: Vec::new(),
|
||||
system_prompt: None,
|
||||
history: Vec::new(),
|
||||
locked_prefix_len: 0,
|
||||
turn_count: 0,
|
||||
turn_notifiers: Vec::new(),
|
||||
_state: PhantomData,
|
||||
}
|
||||
}
|
||||
|
||||
/// システムプロンプトを設定(ビルダーパターン)
|
||||
pub fn system_prompt(mut self, prompt: impl Into<String>) -> Self {
|
||||
self.system_prompt = Some(prompt.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// システムプロンプトを設定(可変参照版)
|
||||
pub fn set_system_prompt(&mut self, prompt: impl Into<String>) {
|
||||
self.system_prompt = Some(prompt.into());
|
||||
}
|
||||
|
||||
/// 履歴への可変参照を取得
|
||||
///
|
||||
/// Mutable状態でのみ利用可能。履歴を自由に編集できる。
|
||||
pub fn history_mut(&mut self) -> &mut Vec<Message> {
|
||||
&mut self.history
|
||||
}
|
||||
|
||||
/// 履歴を設定
|
||||
pub fn set_history(&mut self, messages: Vec<Message>) {
|
||||
self.history = messages;
|
||||
}
|
||||
|
||||
/// 履歴にメッセージを追加(ビルダーパターン)
|
||||
pub fn with_message(mut self, message: Message) -> Self {
|
||||
self.history.push(message);
|
||||
self
|
||||
}
|
||||
|
||||
/// 履歴にメッセージを追加
|
||||
pub fn push_message(&mut self, message: Message) {
|
||||
self.history.push(message);
|
||||
}
|
||||
|
||||
/// 複数のメッセージを履歴に追加(ビルダーパターン)
|
||||
pub fn with_messages(mut self, messages: impl IntoIterator<Item = Message>) -> Self {
|
||||
self.history.extend(messages);
|
||||
self
|
||||
}
|
||||
|
||||
/// 複数のメッセージを履歴に追加
|
||||
pub fn extend_history(&mut self, messages: impl IntoIterator<Item = Message>) {
|
||||
self.history.extend(messages);
|
||||
}
|
||||
|
||||
/// 履歴をクリア
|
||||
pub fn clear_history(&mut self) {
|
||||
self.history.clear();
|
||||
}
|
||||
|
||||
/// 設定を適用(将来の拡張用)
|
||||
#[allow(dead_code)]
|
||||
pub fn config(self, _config: WorkerConfig) -> Self {
|
||||
self
|
||||
}
|
||||
|
||||
/// ロックしてLocked状態へ遷移
|
||||
///
|
||||
/// この操作により、現在のシステムプロンプトと履歴が「確定済みプレフィックス」として
|
||||
/// 固定される。以降は履歴への追記のみが可能となり、キャッシュヒットが保証される。
|
||||
pub fn lock(self) -> Worker<C, Locked> {
|
||||
let locked_prefix_len = self.history.len();
|
||||
Worker {
|
||||
client: self.client,
|
||||
timeline: self.timeline,
|
||||
text_block_collector: self.text_block_collector,
|
||||
tool_call_collector: self.tool_call_collector,
|
||||
tools: self.tools,
|
||||
hooks: self.hooks,
|
||||
system_prompt: self.system_prompt,
|
||||
history: self.history,
|
||||
locked_prefix_len,
|
||||
turn_count: self.turn_count,
|
||||
turn_notifiers: self.turn_notifiers,
|
||||
_state: PhantomData,
|
||||
}
|
||||
}
|
||||
|
||||
/// ターンを実行(Mutable状態)
|
||||
///
|
||||
/// 新しいユーザーメッセージを履歴に追加し、LLMにリクエストを送信する。
|
||||
/// ツール呼び出しがある場合は自動的にループする。
|
||||
///
|
||||
/// 注意: この関数は履歴を変更するため、キャッシュ保護が必要な場合は
|
||||
/// `lock()` を呼んでからLocked状態で `run` を使用すること。
|
||||
pub async fn run(&mut self, user_input: impl Into<String>) -> Result<&[Message], WorkerError> {
|
||||
self.history.push(Message::user(user_input));
|
||||
self.run_turn_loop().await?;
|
||||
Ok(&self.history)
|
||||
}
|
||||
|
||||
/// 複数メッセージでターンを実行(Mutable状態)
|
||||
///
|
||||
/// 指定されたメッセージを履歴に追加してから実行する。
|
||||
pub async fn run_with_messages(
|
||||
&mut self,
|
||||
messages: Vec<Message>,
|
||||
) -> Result<&[Message], WorkerError> {
|
||||
self.history.extend(messages);
|
||||
self.run_turn_loop().await?;
|
||||
Ok(&self.history)
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Locked状態専用の実装
|
||||
// =============================================================================
|
||||
|
||||
impl<C: LlmClient> Worker<C, Locked> {
|
||||
/// ターンを実行(Locked状態)
|
||||
///
|
||||
/// 新しいユーザーメッセージを履歴の末尾に追加し、LLMにリクエストを送信する。
|
||||
/// ロック時点より前の履歴(プレフィックス)は不変であるため、キャッシュヒットが保証される。
|
||||
pub async fn run(&mut self, user_input: impl Into<String>) -> Result<&[Message], WorkerError> {
|
||||
self.history.push(Message::user(user_input));
|
||||
self.run_turn_loop().await?;
|
||||
Ok(&self.history)
|
||||
}
|
||||
|
||||
/// 複数メッセージでターンを実行(Locked状態)
|
||||
pub async fn run_with_messages(
|
||||
&mut self,
|
||||
messages: Vec<Message>,
|
||||
) -> Result<&[Message], WorkerError> {
|
||||
self.history.extend(messages);
|
||||
self.run_turn_loop().await?;
|
||||
Ok(&self.history)
|
||||
}
|
||||
|
||||
/// ロック時点のプレフィックス長を取得
|
||||
pub fn locked_prefix_len(&self) -> usize {
|
||||
self.locked_prefix_len
|
||||
}
|
||||
|
||||
/// ロックを解除してMutable状態へ戻す
|
||||
///
|
||||
/// 注意: この操作を行うと、以降のリクエストでキャッシュがヒットしなくなる可能性がある。
|
||||
/// 履歴を編集する必要がある場合にのみ使用すること。
|
||||
pub fn unlock(self) -> Worker<C, Mutable> {
|
||||
Worker {
|
||||
client: self.client,
|
||||
timeline: self.timeline,
|
||||
text_block_collector: self.text_block_collector,
|
||||
tool_call_collector: self.tool_call_collector,
|
||||
tools: self.tools,
|
||||
hooks: self.hooks,
|
||||
system_prompt: self.system_prompt,
|
||||
history: self.history,
|
||||
locked_prefix_len: 0,
|
||||
turn_count: self.turn_count,
|
||||
turn_notifiers: self.turn_notifiers,
|
||||
_state: PhantomData,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
// 基本的なテストのみ。LlmClientを使ったテストは統合テストで行う。
|
||||
}
|
||||
Reference in New Issue
Block a user