Jeremiah Lowin commited on
Commit
8604c71
·
1 Parent(s): a277f25

Support enabled/disabled prompts

Browse files
src/fastmcp/prompts/prompt.py CHANGED
@@ -96,6 +96,7 @@ class Prompt(FastMCPComponent, ABC):
96
  name: str | None = None,
97
  description: str | None = None,
98
  tags: set[str] | None = None,
 
99
  ) -> FunctionPrompt:
100
  """Create a Prompt from a function.
101
 
@@ -106,7 +107,7 @@ class Prompt(FastMCPComponent, ABC):
106
  - A sequence of any of the above
107
  """
108
  return FunctionPrompt.from_function(
109
- fn=fn, name=name, description=description, tags=tags
110
  )
111
 
112
  @abstractmethod
@@ -130,6 +131,7 @@ class FunctionPrompt(Prompt):
130
  name: str | None = None,
131
  description: str | None = None,
132
  tags: set[str] | None = None,
 
133
  ) -> FunctionPrompt:
134
  """Create a Prompt from a function.
135
 
@@ -195,6 +197,7 @@ class FunctionPrompt(Prompt):
195
  description=description,
196
  arguments=arguments,
197
  tags=tags or set(),
 
198
  fn=fn,
199
  )
200
 
 
96
  name: str | None = None,
97
  description: str | None = None,
98
  tags: set[str] | None = None,
99
+ enabled: bool | None = None,
100
  ) -> FunctionPrompt:
101
  """Create a Prompt from a function.
102
 
 
107
  - A sequence of any of the above
108
  """
109
  return FunctionPrompt.from_function(
110
+ fn=fn, name=name, description=description, tags=tags, enabled=enabled
111
  )
112
 
113
  @abstractmethod
 
131
  name: str | None = None,
132
  description: str | None = None,
133
  tags: set[str] | None = None,
134
+ enabled: bool | None = None,
135
  ) -> FunctionPrompt:
136
  """Create a Prompt from a function.
137
 
 
197
  description=description,
198
  arguments=arguments,
199
  tags=tags or set(),
200
+ enabled=enabled if enabled is not None else True,
201
  fn=fn,
202
  )
203
 
src/fastmcp/prompts/prompt_manager.py CHANGED
@@ -40,9 +40,11 @@ class PromptManager:
40
 
41
  self.duplicate_behavior = duplicate_behavior
42
 
43
- def get_prompt(self, key: str) -> Prompt | None:
44
  """Get prompt by key."""
45
- return self._prompts.get(key)
 
 
46
 
47
  def get_prompts(self) -> dict[str, Prompt]:
48
  """Get all registered prompts, indexed by registered key."""
 
40
 
41
  self.duplicate_behavior = duplicate_behavior
42
 
43
+ def get_prompt(self, key: str) -> Prompt:
44
  """Get prompt by key."""
45
+ if key in self._prompts:
46
+ return self._prompts[key]
47
+ raise NotFoundError(f"Unknown prompt: {key}")
48
 
49
  def get_prompts(self) -> dict[str, Prompt]:
50
  """Get all registered prompts, indexed by registered key."""
src/fastmcp/server/server.py CHANGED
@@ -940,6 +940,7 @@ class FastMCP(Generic[LifespanResultT]):
940
  name: str | None = None,
941
  description: str | None = None,
942
  tags: set[str] | None = None,
 
943
  ) -> FunctionPrompt: ...
944
 
945
  @overload
@@ -950,6 +951,7 @@ class FastMCP(Generic[LifespanResultT]):
950
  name: str | None = None,
951
  description: str | None = None,
952
  tags: set[str] | None = None,
 
953
  ) -> Callable[[AnyFunction], FunctionPrompt]: ...
954
 
955
  def prompt(
@@ -959,6 +961,7 @@ class FastMCP(Generic[LifespanResultT]):
959
  name: str | None = None,
960
  description: str | None = None,
961
  tags: set[str] | None = None,
 
962
  ) -> Callable[[AnyFunction], FunctionPrompt] | FunctionPrompt:
963
  """Decorator to register a prompt.
964
 
@@ -1050,6 +1053,7 @@ class FastMCP(Generic[LifespanResultT]):
1050
  name=prompt_name,
1051
  description=description,
1052
  tags=tags,
 
1053
  )
1054
  self.add_prompt(prompt)
1055
 
@@ -1077,6 +1081,7 @@ class FastMCP(Generic[LifespanResultT]):
1077
  name=prompt_name,
1078
  description=description,
1079
  tags=tags,
 
1080
  )
1081
 
1082
  async def run_stdio_async(self) -> None:
 
940
  name: str | None = None,
941
  description: str | None = None,
942
  tags: set[str] | None = None,
943
+ enabled: bool | None = None,
944
  ) -> FunctionPrompt: ...
945
 
946
  @overload
 
951
  name: str | None = None,
952
  description: str | None = None,
953
  tags: set[str] | None = None,
954
+ enabled: bool | None = None,
955
  ) -> Callable[[AnyFunction], FunctionPrompt]: ...
956
 
957
  def prompt(
 
961
  name: str | None = None,
962
  description: str | None = None,
963
  tags: set[str] | None = None,
964
+ enabled: bool | None = None,
965
  ) -> Callable[[AnyFunction], FunctionPrompt] | FunctionPrompt:
966
  """Decorator to register a prompt.
967
 
 
1053
  name=prompt_name,
1054
  description=description,
1055
  tags=tags,
1056
+ enabled=enabled,
1057
  )
1058
  self.add_prompt(prompt)
1059
 
 
1081
  name=prompt_name,
1082
  description=description,
1083
  tags=tags,
1084
+ enabled=enabled,
1085
  )
1086
 
1087
  async def run_stdio_async(self) -> None:
tests/server/test_proxy.py CHANGED
@@ -178,7 +178,9 @@ class TestResources:
178
  assert json.loads(result[0].text) == USERS # type: ignore[attr-defined]
179
 
180
  async def test_read_resource_returns_none_if_not_found(self, proxy_server):
181
- with pytest.raises(McpError, match="Unknown resource: resource://nonexistent"):
 
 
182
  async with Client(proxy_server) as client:
183
  await client.read_resource("resource://nonexistent")
184
 
 
178
  assert json.loads(result[0].text) == USERS # type: ignore[attr-defined]
179
 
180
  async def test_read_resource_returns_none_if_not_found(self, proxy_server):
181
+ with pytest.raises(
182
+ McpError, match="Unknown resource: 'resource://nonexistent'"
183
+ ):
184
  async with Client(proxy_server) as client:
185
  await client.read_resource("resource://nonexistent")
186
 
tests/server/test_server_interactions.py CHANGED
@@ -1220,6 +1220,93 @@ class TestPrompts:
1220
  assert prompt.tags == {"example", "test-tag"}
1221
 
1222
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1223
  class TestPromptContext:
1224
  async def test_prompt_context(self):
1225
  mcp = FastMCP()
 
1220
  assert prompt.tags == {"example", "test-tag"}
1221
 
1222
 
1223
+ class TestPromptEnabled:
1224
+ async def test_toggle_enabled(self):
1225
+ mcp = FastMCP()
1226
+
1227
+ @mcp.prompt
1228
+ def sample_prompt() -> str:
1229
+ return "Hello, world!"
1230
+
1231
+ assert sample_prompt.enabled
1232
+
1233
+ prompt = await mcp.get_prompt("sample_prompt")
1234
+ assert prompt.enabled
1235
+
1236
+ prompt.disable()
1237
+
1238
+ assert not prompt.enabled
1239
+ assert not sample_prompt.enabled
1240
+
1241
+ prompt.enable()
1242
+ assert prompt.enabled
1243
+ assert sample_prompt.enabled
1244
+
1245
+ async def test_prompt_disabled_in_decorator(self):
1246
+ mcp = FastMCP()
1247
+
1248
+ @mcp.prompt(enabled=False)
1249
+ def sample_prompt() -> str:
1250
+ return "Hello, world!"
1251
+
1252
+ async with Client(mcp) as client:
1253
+ prompts = await client.list_prompts()
1254
+ assert len(prompts) == 0
1255
+
1256
+ async def test_prompt_toggle_enabled(self):
1257
+ mcp = FastMCP()
1258
+
1259
+ @mcp.prompt(enabled=False)
1260
+ def sample_prompt() -> str:
1261
+ return "Hello, world!"
1262
+
1263
+ sample_prompt.enable()
1264
+
1265
+ async with Client(mcp) as client:
1266
+ prompts = await client.list_prompts()
1267
+ assert len(prompts) == 1
1268
+
1269
+ async def test_prompt_toggle_disabled(self):
1270
+ mcp = FastMCP()
1271
+
1272
+ @mcp.prompt
1273
+ def sample_prompt() -> str:
1274
+ return "Hello, world!"
1275
+
1276
+ sample_prompt.disable()
1277
+
1278
+ async with Client(mcp) as client:
1279
+ prompts = await client.list_prompts()
1280
+ assert len(prompts) == 0
1281
+
1282
+ async def test_get_prompt_and_disable(self):
1283
+ mcp = FastMCP()
1284
+
1285
+ @mcp.prompt
1286
+ def sample_prompt() -> str:
1287
+ return "Hello, world!"
1288
+
1289
+ prompt = await mcp.get_prompt("sample_prompt")
1290
+ assert prompt.enabled
1291
+
1292
+ sample_prompt.disable()
1293
+
1294
+ async with Client(mcp) as client:
1295
+ result = await client.list_prompts()
1296
+ assert len(result) == 0
1297
+
1298
+ async def test_cant_get_disabled_prompt(self):
1299
+ mcp = FastMCP()
1300
+
1301
+ @mcp.prompt(enabled=False)
1302
+ def sample_prompt() -> str:
1303
+ return "Hello, world!"
1304
+
1305
+ with pytest.raises(McpError, match="Unknown prompt"):
1306
+ async with Client(mcp) as client:
1307
+ await client.get_prompt("sample_prompt")
1308
+
1309
+
1310
  class TestPromptContext:
1311
  async def test_prompt_context(self):
1312
  mcp = FastMCP()